diff --git a/src-tauri/Cargo.lock b/src-tauri/Cargo.lock index dbafe9a70..86399fc0a 100644 --- a/src-tauri/Cargo.lock +++ b/src-tauri/Cargo.lock @@ -6673,6 +6673,7 @@ dependencies = [ "proxycast-core", "proxycast-credential", "proxycast-infra", + "proxycast-processor", "proxycast-providers", "proxycast-services", "proxycast-terminal", @@ -6813,6 +6814,25 @@ dependencies = [ "uuid", ] +[[package]] +name = "proxycast-processor" +version = "0.60.0" +dependencies = [ + "async-trait", + "parking_lot", + "proptest", + "proxycast-core", + "proxycast-infra", + "proxycast-services", + "serde", + "serde_json", + "subtle", + "thiserror 1.0.69", + "tokio", + "tracing", + "uuid", +] + [[package]] name = "proxycast-providers" version = "0.60.0" diff --git a/src-tauri/Cargo.toml b/src-tauri/Cargo.toml index 8e84cc26b..e7cb811f1 100644 --- a/src-tauri/Cargo.toml +++ b/src-tauri/Cargo.toml @@ -19,6 +19,7 @@ proxycast-services = { path = "crates/services" } proxycast-terminal = { path = "crates/terminal" } proxycast-credential = { path = "crates/credential" } proxycast-websocket = { path = "crates/websocket" } +proxycast-processor = { path = "crates/processor" } voice-core = { path = "crates/voice-core" } # 序列化 @@ -197,6 +198,7 @@ proxycast-services.workspace = true proxycast-terminal.workspace = true proxycast-credential.workspace = true proxycast-websocket.workspace = true +proxycast-processor.workspace = true voice-core.workspace = true # Tauri diff --git a/src-tauri/crates/processor/Cargo.toml b/src-tauri/crates/processor/Cargo.toml new file mode 100644 index 000000000..c2c8c3f6a --- /dev/null +++ b/src-tauri/crates/processor/Cargo.toml @@ -0,0 +1,23 @@ +[package] +name = "proxycast-processor" +version.workspace = true +edition.workspace = true +authors.workspace = true + +[dependencies] +proxycast-core.workspace = true +proxycast-infra.workspace = true +proxycast-services.workspace = true + +serde.workspace = true +serde_json.workspace = true +tokio.workspace = true +async-trait.workspace = true +thiserror.workspace = true +tracing.workspace = true +parking_lot.workspace = true +subtle.workspace = true +uuid.workspace = true + +[dev-dependencies] +proptest.workspace = true diff --git a/src-tauri/crates/processor/src/lib.rs b/src-tauri/crates/processor/src/lib.rs new file mode 100644 index 000000000..abcbbf9c1 --- /dev/null +++ b/src-tauri/crates/processor/src/lib.rs @@ -0,0 +1,14 @@ +//! 请求处理器 crate +//! +//! 提供统一的请求处理管道,集成路由、容错、监控、插件等功能模块。 +//! +//! ## 模块结构 +//! +//! - `steps` - 管道步骤(认证、注入、路由、插件、Provider、遥测) + +pub mod processor; +pub mod steps; + +pub use processor::RequestProcessor; +pub use proxycast_core::processor::RequestContext; +pub use steps::*; diff --git a/src-tauri/crates/processor/src/processor.rs b/src-tauri/crates/processor/src/processor.rs new file mode 100644 index 000000000..6ff643241 --- /dev/null +++ b/src-tauri/crates/processor/src/processor.rs @@ -0,0 +1,190 @@ +//! 请求处理器实现 +//! +//! 提供统一的请求处理管道,集成路由、容错、监控、插件等功能模块。 +//! +//! # 架构 +//! +//! 请求处理流程: +//! 1. 认证 (AuthStep) +//! 2. 参数注入 (InjectionStep) +//! 3. 路由解析 (RoutingStep) +//! 4. 插件前置钩子 (PluginPreStep) +//! 5. Provider 调用 (ProviderStep) - 包含重试和故障转移 +//! 6. 插件后置钩子 (PluginPostStep) +//! 7. 统计记录 (TelemetryStep) + +pub use proxycast_core::processor::RequestContext; + +use parking_lot::RwLock as ParkingLotRwLock; +use proxycast_core::plugin::PluginManager; +use proxycast_core::router::{ModelMapper, Router}; +use proxycast_core::ProviderType; +use proxycast_infra::{ + Failover, Injector, Retrier, StatsAggregator, TimeoutController, TokenTracker, +}; +use proxycast_services::provider_pool_service::ProviderPoolService; +use std::sync::Arc; +use tokio::sync::RwLock; + +/// 统一的请求处理器 +/// +/// 集成所有功能模块,提供完整的请求处理管道 +pub struct RequestProcessor { + /// 路由器 + pub router: Arc>, + /// 模型映射器 + pub mapper: Arc>, + /// 参数注入器 + pub injector: Arc>, + /// 重试器 + pub retrier: Arc, + /// 故障转移器 + pub failover: Arc, + /// 超时控制器 + pub timeout: Arc, + /// 插件管理器 + pub plugins: Arc, + /// 统计聚合器(使用 parking_lot::RwLock 以支持与 TelemetryState 共享) + pub stats: Arc>, + /// Token 追踪器(使用 parking_lot::RwLock 以支持与 TelemetryState 共享) + pub tokens: Arc>, + /// 凭证池服务 + pub pool_service: Arc, + /// 热重载协调锁(避免配置更新期间请求读取不一致的配置) + pub reload_lock: Arc>, +} + +impl RequestProcessor { + /// 创建新的请求处理器 + pub fn new( + router: Arc>, + mapper: Arc>, + injector: Arc>, + retrier: Arc, + failover: Arc, + timeout: Arc, + plugins: Arc, + stats: Arc>, + tokens: Arc>, + pool_service: Arc, + ) -> Self { + Self { + router, + mapper, + injector, + retrier, + failover, + timeout, + plugins, + stats, + tokens, + pool_service, + reload_lock: Arc::new(RwLock::new(())), + } + } + + /// 使用默认配置创建请求处理器 + pub fn with_defaults(pool_service: Arc) -> Self { + Self { + router: Arc::new(RwLock::new(Self::create_router_with_defaults())), + mapper: Arc::new(RwLock::new(ModelMapper::new())), + injector: Arc::new(RwLock::new(Injector::new())), + retrier: Arc::new(Retrier::with_defaults()), + failover: Arc::new(Failover::with_defaults()), + timeout: Arc::new(TimeoutController::with_defaults()), + plugins: Arc::new(PluginManager::with_defaults()), + stats: Arc::new(ParkingLotRwLock::new(StatsAggregator::with_defaults())), + tokens: Arc::new(ParkingLotRwLock::new(TokenTracker::with_defaults())), + pool_service, + reload_lock: Arc::new(RwLock::new(())), + } + } + + /// 创建带默认路由规则的路由器 + /// + /// 注意:不再添加硬编码的路由规则,让用户设置的默认 Provider 生效 + fn create_router_with_defaults() -> Router { + let router = Router::new_empty(); + tracing::info!("[ROUTER] 初始化空路由器,等待从配置加载默认 Provider"); + router + } + + /// 使用共享的统计和 Token 追踪器创建请求处理器 + pub fn with_shared_telemetry( + pool_service: Arc, + stats: Arc>, + tokens: Arc>, + ) -> Self { + Self { + router: Arc::new(RwLock::new(Self::create_router_with_defaults())), + mapper: Arc::new(RwLock::new(ModelMapper::new())), + injector: Arc::new(RwLock::new(Injector::new())), + retrier: Arc::new(Retrier::with_defaults()), + failover: Arc::new(Failover::with_defaults()), + timeout: Arc::new(TimeoutController::with_defaults()), + plugins: Arc::new(PluginManager::with_defaults()), + stats, + tokens, + pool_service, + reload_lock: Arc::new(RwLock::new(())), + } + } + + /// 解析模型别名 + pub async fn resolve_model(&self, model: &str) -> String { + let mapper = self.mapper.read().await; + mapper.resolve(model) + } + + /// 解析模型别名并更新请求上下文 + pub async fn resolve_model_for_context(&self, ctx: &mut RequestContext) -> String { + let resolved = self.resolve_model(&ctx.original_model).await; + ctx.set_resolved_model(resolved.clone()); + + tracing::debug!( + "[MAPPER] request_id={} original_model={} resolved_model={}", + ctx.request_id, + ctx.original_model, + resolved + ); + + resolved + } + + /// 根据模型选择 Provider + pub async fn route_model(&self, model: &str) -> (Option, bool) { + let router = self.router.read().await; + let result = router.route(model); + (result.provider, result.is_default) + } + + /// 根据模型选择 Provider 并更新请求上下文 + pub async fn route_for_context(&self, ctx: &mut RequestContext) -> Option { + let (provider, is_default) = self.route_model(&ctx.resolved_model).await; + + if let Some(p) = provider { + ctx.set_provider(p); + tracing::info!( + "[ROUTE] request_id={} model={} provider={} is_default={}", + ctx.request_id, + ctx.resolved_model, + p, + is_default + ); + } else { + tracing::warn!( + "[ROUTE] request_id={} model={} 未设置默认 Provider", + ctx.request_id, + ctx.resolved_model + ); + } + + provider + } + + /// 执行完整的路由解析流程(模型别名解析 + Provider 选择) + pub async fn resolve_and_route(&self, ctx: &mut RequestContext) -> Option { + self.resolve_model_for_context(ctx).await; + self.route_for_context(ctx).await + } +} diff --git a/src-tauri/src/processor/steps/auth.rs b/src-tauri/crates/processor/src/steps/auth.rs similarity index 74% rename from src-tauri/src/processor/steps/auth.rs rename to src-tauri/crates/processor/src/steps/auth.rs index cd4d54c38..b6b32d79a 100644 --- a/src-tauri/src/processor/steps/auth.rs +++ b/src-tauri/crates/processor/src/steps/auth.rs @@ -1,26 +1,19 @@ //! 认证步骤 -//! -//! 验证请求的 API Key #![allow(dead_code)] use super::traits::{PipelineStep, StepError}; -use crate::processor::RequestContext; use async_trait::async_trait; +use proxycast_core::processor::RequestContext; use subtle::ConstantTimeEq; -/// 认证步骤 -/// -/// 验证请求中的 API Key 是否有效 +/// 认证步骤 - 验证请求中的 API Key pub struct AuthStep { - /// 期望的 API Key expected_key: String, - /// 是否启用 enabled: bool, } impl AuthStep { - /// 创建新的认证步骤 pub fn new(expected_key: String) -> Self { Self { expected_key, @@ -28,13 +21,11 @@ impl AuthStep { } } - /// 设置是否启用 pub fn with_enabled(mut self, enabled: bool) -> Self { self.enabled = enabled; self } - /// 验证 API Key pub fn verify(&self, provided_key: Option<&str>) -> Result<(), StepError> { match provided_key { Some(key) if key.as_bytes().ct_eq(self.expected_key.as_bytes()).into() => Ok(()), @@ -51,19 +42,16 @@ impl PipelineStep for AuthStep { ctx: &mut RequestContext, _payload: &mut serde_json::Value, ) -> Result<(), StepError> { - // 从元数据中获取 API Key let api_key = ctx .get_metadata("api_key") .and_then(|v| v.as_str()) .map(|s| s.to_string()); - self.verify(api_key.as_deref()) } fn name(&self) -> &str { "auth" } - fn is_enabled(&self) -> bool { self.enabled } @@ -82,17 +70,16 @@ mod tests { #[test] fn test_auth_step_verify_invalid_key() { let step = AuthStep::new("test-key".to_string()); - let result = step.verify(Some("wrong-key")); - assert!(result.is_err()); - assert!(matches!(result.unwrap_err(), StepError::Auth(_))); + assert!(matches!( + step.verify(Some("wrong-key")), + Err(StepError::Auth(_)) + )); } #[test] fn test_auth_step_verify_no_key() { let step = AuthStep::new("test-key".to_string()); - let result = step.verify(None); - assert!(result.is_err()); - assert!(matches!(result.unwrap_err(), StepError::Auth(_))); + assert!(matches!(step.verify(None), Err(StepError::Auth(_)))); } #[tokio::test] @@ -101,8 +88,6 @@ mod tests { let mut ctx = RequestContext::new("model".to_string()); ctx.set_metadata("api_key", serde_json::json!("test-key")); let mut payload = serde_json::json!({}); - - let result = step.execute(&mut ctx, &mut payload).await; - assert!(result.is_ok()); + assert!(step.execute(&mut ctx, &mut payload).await.is_ok()); } } diff --git a/src-tauri/src/processor/steps/injection.rs b/src-tauri/crates/processor/src/steps/injection.rs similarity index 82% rename from src-tauri/src/processor/steps/injection.rs rename to src-tauri/crates/processor/src/steps/injection.rs index 0ab80242a..6887b37ae 100644 --- a/src-tauri/src/processor/steps/injection.rs +++ b/src-tauri/crates/processor/src/steps/injection.rs @@ -1,28 +1,21 @@ //! 参数注入步骤 -//! -//! 根据配置的规则注入请求参数 #![allow(dead_code)] use super::traits::{PipelineStep, StepError}; -use crate::injection::Injector; -use crate::processor::RequestContext; use async_trait::async_trait; +use proxycast_core::processor::RequestContext; +use proxycast_infra::Injector; use std::sync::Arc; use tokio::sync::RwLock; /// 参数注入步骤 -/// -/// 根据模型匹配规则注入请求参数 pub struct InjectionStep { - /// 注入器 injector: Arc>, - /// 是否启用 enabled: Arc>, } impl InjectionStep { - /// 创建新的注入步骤 pub fn new(injector: Arc>) -> Self { Self { injector, @@ -30,12 +23,10 @@ impl InjectionStep { } } - /// 设置是否启用 pub fn with_enabled(self, enabled: Arc>) -> Self { Self { enabled, ..self } } - /// 检查是否启用 pub async fn is_injection_enabled(&self) -> bool { *self.enabled.read().await } @@ -62,8 +53,6 @@ impl PipelineStep for InjectionStep { result.applied_rules, result.injected_params ); - - // 记录注入信息到元数据 ctx.set_metadata( "injection_result", serde_json::json!({ @@ -84,7 +73,7 @@ impl PipelineStep for InjectionStep { #[cfg(test)] mod tests { use super::*; - use crate::injection::InjectionRule; + use proxycast_infra::InjectionRule; #[tokio::test] async fn test_injection_step_execute() { @@ -94,13 +83,10 @@ mod tests { "claude-*", serde_json::json!({"temperature": 0.7}), )); - let step = InjectionStep::new(Arc::new(RwLock::new(injector))); let mut ctx = RequestContext::new("claude-sonnet-4-5".to_string()); let mut payload = serde_json::json!({"model": "claude-sonnet-4-5"}); - - let result = step.execute(&mut ctx, &mut payload).await; - assert!(result.is_ok()); + assert!(step.execute(&mut ctx, &mut payload).await.is_ok()); assert_eq!(payload["temperature"], 0.7); } @@ -112,15 +98,11 @@ mod tests { "claude-*", serde_json::json!({"temperature": 0.7}), )); - let step = InjectionStep::new(Arc::new(RwLock::new(injector))) .with_enabled(Arc::new(RwLock::new(false))); let mut ctx = RequestContext::new("claude-sonnet-4-5".to_string()); let mut payload = serde_json::json!({"model": "claude-sonnet-4-5"}); - - let result = step.execute(&mut ctx, &mut payload).await; - assert!(result.is_ok()); - // 参数不应该被注入 + assert!(step.execute(&mut ctx, &mut payload).await.is_ok()); assert!(payload.get("temperature").is_none()); } } diff --git a/src-tauri/crates/processor/src/steps/mod.rs b/src-tauri/crates/processor/src/steps/mod.rs new file mode 100644 index 000000000..47b119d8a --- /dev/null +++ b/src-tauri/crates/processor/src/steps/mod.rs @@ -0,0 +1,25 @@ +//! 管道步骤模块 +//! +//! 定义请求处理管道中的各个步骤 + +mod auth; +mod injection; +mod plugin; +mod provider; +mod routing; +mod telemetry; +mod traits; + +#[allow(unused_imports)] +pub use auth::AuthStep; +#[allow(unused_imports)] +pub use injection::InjectionStep; +#[allow(unused_imports)] +pub use plugin::{PluginPostStep, PluginPreStep}; +pub use provider::{ProviderCallError, ProviderCallResult, ProviderStep}; +#[allow(unused_imports)] +pub use routing::RoutingStep; +#[allow(unused_imports)] +pub use telemetry::TelemetryStep; +#[allow(unused_imports)] +pub use traits::{PipelineStep, StepError}; diff --git a/src-tauri/src/processor/steps/plugin.rs b/src-tauri/crates/processor/src/steps/plugin.rs similarity index 74% rename from src-tauri/src/processor/steps/plugin.rs rename to src-tauri/crates/processor/src/steps/plugin.rs index 4bbbdb6f0..d4fd54a06 100644 --- a/src-tauri/src/processor/steps/plugin.rs +++ b/src-tauri/crates/processor/src/steps/plugin.rs @@ -1,26 +1,20 @@ //! 插件钩子步骤 -//! -//! 执行插件的前置和后置钩子 #![allow(dead_code)] use super::traits::{PipelineStep, StepError}; -use crate::plugin::PluginManager; -use crate::processor::RequestContext; -use crate::ProviderType; use async_trait::async_trait; +use proxycast_core::plugin::PluginManager; +use proxycast_core::processor::RequestContext; +use proxycast_core::ProviderType; use std::sync::Arc; /// 插件前置钩子步骤 -/// -/// 在 Provider 调用前执行所有启用插件的 on_request 钩子 pub struct PluginPreStep { - /// 插件管理器 plugins: Arc, } impl PluginPreStep { - /// 创建新的插件前置步骤 pub fn new(plugins: Arc) -> Self { Self { plugins } } @@ -33,36 +27,26 @@ impl PipelineStep for PluginPreStep { ctx: &mut RequestContext, payload: &mut serde_json::Value, ) -> Result<(), StepError> { - // 初始化插件上下文 let provider = ctx.provider.unwrap_or(ProviderType::Kiro); ctx.init_plugin_context(provider); - // 获取插件上下文的可变引用 if let Some(plugin_ctx) = ctx.plugin_context_mut() { let results = self.plugins.run_on_request(plugin_ctx, payload).await; - - // 检查是否有失败的钩子 for result in &results { if !result.success { tracing::warn!("[PLUGIN] on_request hook failed: {:?}", result.error); - // 插件失败不阻止请求继续,只记录警告 } } - - // 记录插件执行结果到元数据 ctx.set_metadata( "plugin_pre_results", serde_json::json!(results .iter() .map(|r| serde_json::json!({ - "success": r.success, - "modified": r.modified, - "duration_ms": r.duration_ms + "success": r.success, "modified": r.modified, "duration_ms": r.duration_ms })) .collect::>()), ); } - Ok(()) } @@ -72,24 +56,18 @@ impl PipelineStep for PluginPreStep { } /// 插件后置钩子步骤 -/// -/// 在 Provider 调用后执行所有启用插件的 on_response 钩子 pub struct PluginPostStep { - /// 插件管理器 plugins: Arc, } impl PluginPostStep { - /// 创建新的插件后置步骤 pub fn new(plugins: Arc) -> Self { Self { plugins } } - /// 执行错误钩子 pub async fn run_on_error(&self, ctx: &mut RequestContext, error: &str) { if let Some(plugin_ctx) = ctx.plugin_context_mut() { let results = self.plugins.run_on_error(plugin_ctx, error).await; - for result in &results { if !result.success { tracing::warn!("[PLUGIN] on_error hook failed: {:?}", result.error); @@ -108,28 +86,21 @@ impl PipelineStep for PluginPostStep { ) -> Result<(), StepError> { if let Some(plugin_ctx) = ctx.plugin_context_mut() { let results = self.plugins.run_on_response(plugin_ctx, payload).await; - - // 检查是否有失败的钩子 for result in &results { if !result.success { tracing::warn!("[PLUGIN] on_response hook failed: {:?}", result.error); } } - - // 记录插件执行结果到元数据 ctx.set_metadata( "plugin_post_results", serde_json::json!(results .iter() .map(|r| serde_json::json!({ - "success": r.success, - "modified": r.modified, - "duration_ms": r.duration_ms + "success": r.success, "modified": r.modified, "duration_ms": r.duration_ms })) .collect::>()), ); } - Ok(()) } @@ -146,13 +117,10 @@ mod tests { async fn test_plugin_pre_step_execute() { let plugins = Arc::new(PluginManager::with_defaults()); let step = PluginPreStep::new(plugins); - let mut ctx = RequestContext::new("model".to_string()); ctx.set_provider(ProviderType::Kiro); let mut payload = serde_json::json!({"model": "model"}); - - let result = step.execute(&mut ctx, &mut payload).await; - assert!(result.is_ok()); + assert!(step.execute(&mut ctx, &mut payload).await.is_ok()); assert!(ctx.plugin_ctx.is_some()); } @@ -160,13 +128,10 @@ mod tests { async fn test_plugin_post_step_execute() { let plugins = Arc::new(PluginManager::with_defaults()); let step = PluginPostStep::new(plugins); - let mut ctx = RequestContext::new("model".to_string()); ctx.set_provider(ProviderType::Kiro); ctx.init_plugin_context(ProviderType::Kiro); let mut payload = serde_json::json!({"response": "test"}); - - let result = step.execute(&mut ctx, &mut payload).await; - assert!(result.is_ok()); + assert!(step.execute(&mut ctx, &mut payload).await.is_ok()); } } diff --git a/src-tauri/src/processor/steps/provider.rs b/src-tauri/crates/processor/src/steps/provider.rs similarity index 77% rename from src-tauri/src/processor/steps/provider.rs rename to src-tauri/crates/processor/src/steps/provider.rs index b12f73e1a..4f8a525d8 100644 --- a/src-tauri/src/processor/steps/provider.rs +++ b/src-tauri/crates/processor/src/steps/provider.rs @@ -5,45 +5,36 @@ #![allow(dead_code)] use super::traits::{PipelineStep, StepError}; -use crate::processor::RequestContext; -use crate::resilience::{ - Failover, FailoverConfig, FailoverManager, Retrier, RetryConfig, TimeoutConfig, - TimeoutController, TimeoutError, -}; -use crate::services::provider_pool_service::ProviderPoolService; -use crate::ProviderType; use async_trait::async_trait; +use proxycast_core::processor::RequestContext; +use proxycast_core::ProviderType; +use proxycast_infra::resilience::{FailoverManager, TimeoutError}; +use proxycast_infra::{ + Failover, FailoverConfig, Retrier, RetryConfig, TimeoutConfig, TimeoutController, +}; +use proxycast_services::provider_pool_service::ProviderPoolService; use std::future::Future; use std::sync::Arc; /// Provider 调用结果 #[derive(Debug, Clone)] pub struct ProviderCallResult { - /// 响应内容 pub response: serde_json::Value, - /// HTTP 状态码 pub status_code: u16, - /// 延迟(毫秒) pub latency_ms: u64, - /// 使用的凭证 ID pub credential_id: Option, } /// Provider 调用错误 #[derive(Debug, Clone)] pub struct ProviderCallError { - /// 错误消息 pub message: String, - /// HTTP 状态码(如果有) pub status_code: Option, - /// 是否可重试 pub retryable: bool, - /// 是否应触发故障转移 pub should_failover: bool, } impl ProviderCallError { - /// 创建可重试错误 pub fn retryable(message: impl Into, status_code: Option) -> Self { Self { message: message.into(), @@ -53,7 +44,6 @@ impl ProviderCallError { } } - /// 创建需要故障转移的错误 pub fn failover(message: impl Into, status_code: Option) -> Self { Self { message: message.into(), @@ -63,7 +53,6 @@ impl ProviderCallError { } } - /// 创建不可恢复错误 pub fn fatal(message: impl Into, status_code: Option) -> Self { Self { message: message.into(), @@ -73,28 +62,20 @@ impl ProviderCallError { } } - /// 检查是否为配额超限错误 pub fn is_quota_exceeded(&self) -> bool { Failover::is_quota_exceeded(self.status_code, &self.message) } } /// Provider 调用步骤 -/// -/// 包含重试、故障转移和超时控制的 Provider 调用 pub struct ProviderStep { - /// 重试器 retrier: Arc, - /// 故障转移器 failover: Arc, - /// 超时控制器 timeout: Arc, - /// 凭证池服务 pool_service: Arc, } impl ProviderStep { - /// 创建新的 Provider 步骤 pub fn new( retrier: Arc, failover: Arc, @@ -109,7 +90,6 @@ impl ProviderStep { } } - /// 使用默认配置创建 pub fn with_defaults(pool_service: Arc) -> Self { Self { retrier: Arc::new(Retrier::with_defaults()), @@ -119,7 +99,6 @@ impl ProviderStep { } } - /// 使用自定义配置创建 pub fn with_config( retry_config: RetryConfig, failover_config: FailoverConfig, @@ -134,36 +113,20 @@ impl ProviderStep { } } - /// 获取重试器 pub fn retrier(&self) -> &Retrier { &self.retrier } - - /// 获取故障转移器 pub fn failover(&self) -> &Failover { &self.failover } - - /// 获取超时控制器 pub fn timeout(&self) -> &TimeoutController { &self.timeout } - - /// 获取凭证池服务 pub fn pool_service(&self) -> &ProviderPoolService { &self.pool_service } /// 带重试执行 Provider 调用 - /// - /// 使用 Retrier 包装 Provider 调用,自动处理可重试错误 - /// - /// # Arguments - /// * `ctx` - 请求上下文 - /// * `operation` - Provider 调用操作 - /// - /// # Returns - /// 成功返回调用结果,失败返回错误 pub async fn execute_with_retry( &self, ctx: &mut RequestContext, @@ -178,13 +141,10 @@ impl ProviderStep { loop { attempts += 1; - match operation().await { Ok(result) => return Ok(result), Err(err) => { - // 增加重试计数 ctx.increment_retry(); - tracing::warn!( "[RETRY] request_id={} attempt={}/{} error={} status={:?} retryable={}", ctx.request_id, @@ -194,19 +154,13 @@ impl ProviderStep { err.status_code, err.retryable ); - - // 如果不可重试,立即返回 if !err.retryable { return Err(err); } - - // 检查状态码是否可重试 let should_retry = err .status_code .is_none_or(|code| self.retrier.config().is_retryable(code)); - let should_failover = err.should_failover || err.is_quota_exceeded(); - if !should_retry || attempts > max_retries { return Err(ProviderCallError { message: err.message, @@ -215,8 +169,6 @@ impl ProviderStep { should_failover, }); } - - // 等待退避时间 let delay = self.retrier.backoff_delay(attempts - 1); tokio::time::sleep(delay).await; } @@ -225,15 +177,6 @@ impl ProviderStep { } /// 带超时执行 Provider 调用 - /// - /// 使用 TimeoutController 包装 Provider 调用,自动处理超时 - /// - /// # Arguments - /// * `ctx` - 请求上下文 - /// * `operation` - Provider 调用操作 - /// - /// # Returns - /// 成功返回调用结果,失败返回错误 pub async fn execute_with_timeout( &self, ctx: &RequestContext, @@ -243,7 +186,6 @@ impl ProviderStep { F: Future>, { let timeout_result = self.timeout.execute_with_timeout(operation).await; - match timeout_result { Ok(call_result) => call_result, Err(timeout_err) => { @@ -252,14 +194,12 @@ impl ProviderStep { TimeoutError::StreamIdleTimeout { timeout_ms, .. } => *timeout_ms, TimeoutError::Cancelled => 0, }; - tracing::warn!( "[TIMEOUT] request_id={} error={} timeout_ms={}", ctx.request_id, timeout_err, timeout_ms ); - Err(ProviderCallError { message: timeout_err.to_string(), status_code: Some(408), @@ -270,17 +210,7 @@ impl ProviderStep { } } - /// 带故障转移执行 Provider 调用 - /// - /// 使用 Failover 处理 Provider 失败,自动切换到其他 Provider - /// - /// # Arguments - /// * `ctx` - 请求上下文 - /// * `error` - Provider 调用错误 - /// * `available_providers` - 可用的 Provider 列表 - /// - /// # Returns - /// 如果可以故障转移,返回新的 Provider;否则返回 None + /// 带故障转移处理 Provider 失败 pub fn handle_failover( &self, ctx: &RequestContext, @@ -288,14 +218,12 @@ impl ProviderStep { available_providers: &[ProviderType], ) -> Option { let current_provider = ctx.provider?; - let result = self.failover.handle_failure( current_provider, error.status_code, &error.message, available_providers, ); - if result.switched { tracing::info!( "[FAILOVER] request_id={} from={} to={:?} reason={:?}", @@ -317,16 +245,6 @@ impl ProviderStep { } /// 带重试、超时和故障转移执行完整的 Provider 调用 - /// - /// 这是主要的调用入口,集成了所有容错机制 - /// - /// # Arguments - /// * `ctx` - 请求上下文 - /// * `operation` - Provider 调用操作工厂 - /// * `available_providers` - 可用的 Provider 列表 - /// - /// # Returns - /// 成功返回调用结果,失败返回 StepError pub async fn execute_with_resilience( &self, ctx: &mut RequestContext, @@ -344,9 +262,8 @@ impl ProviderStep { let max_retries = self.retrier.config().max_retries; 'failover: loop { - // 更新上下文中的 Provider ctx.set_provider(current_provider); - ctx.retry_count = 0; // 重置重试计数 + ctx.retry_count = 0; tracing::info!( "[PROVIDER] request_id={} provider={} model={} failover_attempt={}", @@ -356,68 +273,48 @@ impl ProviderStep { failover_attempts ); - // 重试循环 let mut retry_attempts = 0u32; - let result: Result = loop { - retry_attempts += 1; - - // 带超时执行调用 - let call_result = self - .execute_with_timeout(ctx, operation_factory(current_provider)) - .await; - - match call_result { - Ok(result) => break Ok(result), - Err(err) => { - ctx.increment_retry(); - - tracing::warn!( + let result: Result = + loop { + retry_attempts += 1; + let call_result = self + .execute_with_timeout(ctx, operation_factory(current_provider)) + .await; + match call_result { + Ok(result) => break Ok(result), + Err(err) => { + ctx.increment_retry(); + tracing::warn!( "[RETRY] request_id={} attempt={}/{} error={} status={:?} retryable={}", - ctx.request_id, - retry_attempts, - max_retries + 1, - err.message, - err.status_code, - err.retryable + ctx.request_id, retry_attempts, max_retries + 1, + err.message, err.status_code, err.retryable ); - - // 如果不可重试,立即返回错误 - if !err.retryable { - break Err(err); + if !err.retryable { + break Err(err); + } + let should_retry = err + .status_code + .is_none_or(|code| self.retrier.config().is_retryable(code)); + let should_failover = err.should_failover || err.is_quota_exceeded(); + if !should_retry || retry_attempts > max_retries { + break Err(ProviderCallError { + message: err.message, + status_code: err.status_code, + retryable: false, + should_failover, + }); + } + let delay = self.retrier.backoff_delay(retry_attempts - 1); + tokio::time::sleep(delay).await; } - - // 检查状态码是否可重试 - let should_retry = err - .status_code - .is_none_or(|code| self.retrier.config().is_retryable(code)); - - let should_failover = err.should_failover || err.is_quota_exceeded(); - - if !should_retry || retry_attempts > max_retries { - break Err(ProviderCallError { - message: err.message, - status_code: err.status_code, - retryable: false, - should_failover, - }); - } - - // 等待退避时间 - let delay = self.retrier.backoff_delay(retry_attempts - 1); - tokio::time::sleep(delay).await; } - } - }; + }; match result { - Ok(call_result) => { - return Ok(call_result); - } + Ok(call_result) => return Ok(call_result), Err(err) => { - // 检查是否应该故障转移 if err.should_failover || err.is_quota_exceeded() { failover_attempts += 1; - if failover_attempts >= max_failover_attempts { tracing::error!( "[PROVIDER] request_id={} all_providers_failed attempts={}", @@ -429,15 +326,12 @@ impl ProviderStep { err.message ))); } - - // 尝试故障转移 let failover_result = failover_manager.handle_failure_and_switch( current_provider, err.status_code, &err.message, available_providers, ); - if let Some(new_provider) = failover_result.new_provider { tracing::info!( "[FAILOVER] request_id={} from={} to={} reason={:?}", @@ -450,20 +344,16 @@ impl ProviderStep { continue 'failover; } } - - // 无法恢复,返回错误 return Err(StepError::Provider(err.message)); } } } } - /// 检查错误是否为配额超限 pub fn is_quota_exceeded_error(&self, error: &ProviderCallError) -> bool { error.is_quota_exceeded() } - /// 检查状态码是否可重试 pub fn is_retryable_status(&self, status_code: u16) -> bool { self.retrier.config().is_retryable(status_code) } @@ -476,10 +366,6 @@ impl PipelineStep for ProviderStep { ctx: &mut RequestContext, _payload: &mut serde_json::Value, ) -> Result<(), StepError> { - // 注意:实际的 Provider 调用逻辑在 server.rs 中实现 - // 这里的 execute 方法主要用于管道步骤的统一接口 - // 实际调用应使用 execute_with_resilience 方法 - tracing::info!( "[PROVIDER] request_id={} provider={:?} model={} retry_count={}", ctx.request_id, @@ -487,8 +373,6 @@ impl PipelineStep for ProviderStep { ctx.resolved_model, ctx.retry_count ); - - // 占位实现 - 实际调用通过 execute_with_resilience 进行 Ok(()) } @@ -506,7 +390,6 @@ mod tests { async fn test_provider_step_new() { let pool_service = Arc::new(ProviderPoolService::new()); let step = ProviderStep::with_defaults(pool_service); - assert_eq!(step.name(), "provider"); assert!(step.is_enabled()); } @@ -515,14 +398,10 @@ mod tests { async fn test_provider_step_execute() { let pool_service = Arc::new(ProviderPoolService::new()); let step = ProviderStep::with_defaults(pool_service); - let mut ctx = RequestContext::new("claude-sonnet-4-5".to_string()); let mut payload = serde_json::json!({"model": "claude-sonnet-4-5"}); - - let result = step.execute(&mut ctx, &mut payload).await; - assert!(result.is_ok()); + assert!(step.execute(&mut ctx, &mut payload).await.is_ok()); } - #[tokio::test] async fn test_provider_step_with_config() { let pool_service = Arc::new(ProviderPoolService::new()); @@ -586,7 +465,6 @@ mod tests { let pool_service = Arc::new(ProviderPoolService::new()); let step = ProviderStep::with_defaults(pool_service); - // 可重试状态码 assert!(step.is_retryable_status(408)); assert!(step.is_retryable_status(429)); assert!(step.is_retryable_status(500)); @@ -594,7 +472,6 @@ mod tests { assert!(step.is_retryable_status(503)); assert!(step.is_retryable_status(504)); - // 不可重试状态码 assert!(!step.is_retryable_status(200)); assert!(!step.is_retryable_status(400)); assert!(!step.is_retryable_status(401)); @@ -620,8 +497,7 @@ mod tests { .await; assert!(result.is_ok()); - let call_result = result.unwrap(); - assert_eq!(call_result.status_code, 200); + assert_eq!(result.unwrap().status_code, 200); } #[tokio::test] @@ -638,7 +514,7 @@ mod tests { assert!(result.is_err()); let err = result.unwrap_err(); - assert_eq!(err.status_code, Some(401)); // 保留原始状态码 + assert_eq!(err.status_code, Some(401)); assert!(!err.retryable); } @@ -657,7 +533,6 @@ mod tests { ]; let new_provider = step.handle_failover(&ctx, &error, &available); - assert!(new_provider.is_some()); assert_eq!(new_provider.unwrap(), ProviderType::Gemini); } @@ -670,10 +545,9 @@ mod tests { ctx.set_provider(ProviderType::Kiro); let error = ProviderCallError::failover("Rate limit exceeded", Some(429)); - let available = vec![ProviderType::Kiro]; // 只有一个 Provider + let available = vec![ProviderType::Kiro]; let new_provider = step.handle_failover(&ctx, &error, &available); - assert!(new_provider.is_none()); } @@ -706,7 +580,7 @@ mod tests { #[tokio::test] async fn test_execute_with_timeout_timeout() { let pool_service = Arc::new(ProviderPoolService::new()); - let timeout_config = TimeoutConfig::new(50, 0); // 50ms 超时 + let timeout_config = TimeoutConfig::new(50, 0); let step = ProviderStep::with_config( RetryConfig::default(), FailoverConfig::default(), diff --git a/src-tauri/src/processor/steps/routing.rs b/src-tauri/crates/processor/src/steps/routing.rs similarity index 71% rename from src-tauri/src/processor/steps/routing.rs rename to src-tauri/crates/processor/src/steps/routing.rs index d290207ec..0b738da04 100644 --- a/src-tauri/src/processor/steps/routing.rs +++ b/src-tauri/crates/processor/src/steps/routing.rs @@ -1,31 +1,23 @@ //! 路由解析步骤 -//! -//! 解析模型别名并选择 Provider #![allow(dead_code)] use super::traits::{PipelineStep, StepError}; -use crate::processor::RequestContext; -use crate::router::{ModelMapper, Router}; -use crate::ProviderType; use async_trait::async_trait; +use proxycast_core::processor::RequestContext; +use proxycast_core::router::{ModelMapper, Router}; +use proxycast_core::ProviderType; use std::sync::Arc; use tokio::sync::RwLock; /// 路由解析步骤 -/// -/// 解析模型别名并根据路由规则选择 Provider pub struct RoutingStep { - /// 路由器 router: Arc>, - /// 模型映射器 mapper: Arc>, - /// 默认 Provider default_provider: Arc>, } impl RoutingStep { - /// 创建新的路由步骤 pub fn new( router: Arc>, mapper: Arc>, @@ -38,20 +30,14 @@ impl RoutingStep { } } - /// 解析模型别名 pub async fn resolve_model(&self, model: &str) -> String { let mapper = self.mapper.read().await; mapper.resolve(model) } - /// 根据模型选择 Provider pub async fn select_provider(&self, model: &str) -> Result { let router = self.router.read().await; - - // 使用路由规则(如果没有匹配的规则,会返回默认 Provider) let result = router.route(model); - - // 如果没有设置默认 Provider,返回错误 result.provider.ok_or_else(|| { StepError::Routing("未设置默认 Provider,请先在设置中选择一个默认 Provider".to_string()) }) @@ -65,16 +51,13 @@ impl PipelineStep for RoutingStep { ctx: &mut RequestContext, payload: &mut serde_json::Value, ) -> Result<(), StepError> { - // 解析模型别名 let resolved_model = self.resolve_model(&ctx.original_model).await; ctx.set_resolved_model(resolved_model.clone()); - // 更新 payload 中的模型名 if let Some(obj) = payload.as_object_mut() { obj.insert("model".to_string(), serde_json::json!(resolved_model)); } - // 选择 Provider let provider = self.select_provider(&ctx.resolved_model).await?; ctx.set_provider(provider); @@ -102,58 +85,41 @@ mod tests { async fn test_routing_step_resolve_model() { let mut mapper = ModelMapper::new(); mapper.add_alias("gpt-4", "claude-sonnet-4-5"); - let step = RoutingStep::new( Arc::new(RwLock::new(Router::new(ProviderType::Kiro))), Arc::new(RwLock::new(mapper)), Arc::new(RwLock::new("kiro".to_string())), ); - - // 别名应解析为实际模型 - let resolved = step.resolve_model("gpt-4").await; - assert_eq!(resolved, "claude-sonnet-4-5"); - - // 非别名应返回原值 - let resolved = step.resolve_model("unknown-model").await; - assert_eq!(resolved, "unknown-model"); + assert_eq!(step.resolve_model("gpt-4").await, "claude-sonnet-4-5"); + assert_eq!(step.resolve_model("unknown-model").await, "unknown-model"); } #[tokio::test] async fn test_routing_step_select_provider() { let router = Router::new(ProviderType::Kiro); - let step = RoutingStep::new( Arc::new(RwLock::new(router)), Arc::new(RwLock::new(ModelMapper::new())), Arc::new(RwLock::new("kiro".to_string())), ); - - // 所有模型都使用默认 Provider - let provider = step.select_provider("gemini-2.5-flash").await; - assert!(provider.is_ok()); - assert_eq!(provider.unwrap(), ProviderType::Kiro); - - let provider = step.select_provider("claude-sonnet-4-5").await; - assert!(provider.is_ok()); - assert_eq!(provider.unwrap(), ProviderType::Kiro); + assert_eq!( + step.select_provider("gemini-2.5-flash").await.unwrap(), + ProviderType::Kiro + ); } #[tokio::test] async fn test_routing_step_execute() { let mut mapper = ModelMapper::new(); mapper.add_alias("gpt-4", "claude-sonnet-4-5"); - let step = RoutingStep::new( Arc::new(RwLock::new(Router::new(ProviderType::Kiro))), Arc::new(RwLock::new(mapper)), Arc::new(RwLock::new("kiro".to_string())), ); - let mut ctx = RequestContext::new("gpt-4".to_string()); let mut payload = serde_json::json!({"model": "gpt-4"}); - - let result = step.execute(&mut ctx, &mut payload).await; - assert!(result.is_ok()); + assert!(step.execute(&mut ctx, &mut payload).await.is_ok()); assert_eq!(ctx.resolved_model, "claude-sonnet-4-5"); assert_eq!(ctx.provider, Some(ProviderType::Kiro)); assert_eq!(payload["model"], "claude-sonnet-4-5"); diff --git a/src-tauri/src/processor/steps/telemetry.rs b/src-tauri/crates/processor/src/steps/telemetry.rs similarity index 68% rename from src-tauri/src/processor/steps/telemetry.rs rename to src-tauri/crates/processor/src/steps/telemetry.rs index e3d7daa83..328cc8472 100644 --- a/src-tauri/src/processor/steps/telemetry.rs +++ b/src-tauri/crates/processor/src/steps/telemetry.rs @@ -1,39 +1,29 @@ //! 统计记录步骤 -//! -//! 记录请求统计和 Token 使用 #![allow(dead_code)] use super::traits::{PipelineStep, StepError}; -use crate::processor::RequestContext; -use crate::telemetry::{ - RequestLog, RequestStatus, StatsAggregator, TokenSource, TokenTracker, TokenUsageRecord, -}; -use crate::ProviderType; use async_trait::async_trait; use parking_lot::RwLock; +use proxycast_core::processor::RequestContext; +use proxycast_core::ProviderType; +use proxycast_infra::{ + RequestLog, RequestStatus, StatsAggregator, TokenSource, TokenTracker, TokenUsageRecord, +}; use std::sync::Arc; /// 统计记录步骤 -/// -/// 记录请求统计和 Token 使用信息 pub struct TelemetryStep { - /// 统计聚合器(使用 parking_lot::RwLock 以支持与 TelemetryState 共享) stats: Arc>, - /// Token 追踪器(使用 parking_lot::RwLock 以支持与 TelemetryState 共享) tokens: Arc>, } impl TelemetryStep { - /// 创建新的统计记录步骤 pub fn new(stats: Arc>, tokens: Arc>) -> Self { Self { stats, tokens } } /// 记录请求日志 - /// - /// 请求完成后记录统计,按 Provider 和模型分组 - /// _需求: 4.1_ pub fn record_request( &self, ctx: &RequestContext, @@ -47,8 +37,6 @@ impl TelemetryStep { ctx.resolved_model.clone(), ctx.is_stream, ); - - // 设置状态和持续时间 match status { RequestStatus::Success => log.mark_success(ctx.elapsed_ms(), 200), RequestStatus::Failed => { @@ -60,24 +48,15 @@ impl TelemetryStep { log.duration_ms = ctx.elapsed_ms(); } } - - // 设置凭证 ID if let Some(cred_id) = &ctx.credential_id { log.set_credential_id(cred_id.clone()); } - - // 设置重试次数 log.retry_count = ctx.retry_count; - - // 使用 parking_lot::RwLock 的同步写锁 let stats = self.stats.write(); stats.record(log); } /// 记录 Token 使用 - /// - /// 从响应提取 Token 数,无 Token 时使用估算 - /// _需求: 4.2, 4.3_ pub fn record_tokens( &self, ctx: &RequestContext, @@ -86,8 +65,6 @@ impl TelemetryStep { source: TokenSource, ) { let provider = ctx.provider.unwrap_or(ProviderType::Kiro); - - // 只有当至少有一个 Token 值时才记录 if input_tokens.is_some() || output_tokens.is_some() { let record = TokenUsageRecord::new( uuid::Uuid::new_v4().to_string(), @@ -98,47 +75,38 @@ impl TelemetryStep { source, ) .with_request_id(ctx.request_id.clone()); - - // 使用 parking_lot::RwLock 的同步写锁 let tokens = self.tokens.write(); tokens.record(record); } } /// 从响应中提取并记录 Token 使用 - /// - /// 支持 OpenAI 和 Anthropic 两种响应格式 pub fn record_tokens_from_response(&self, ctx: &RequestContext, response: &serde_json::Value) { - // 尝试从 OpenAI 格式响应中提取 Token if let Some(usage) = response.get("usage") { - let input_tokens = usage + // OpenAI 格式 + let input = usage .get("prompt_tokens") .and_then(|v| v.as_u64()) .map(|v| v as u32); - let output_tokens = usage + let output = usage .get("completion_tokens") .and_then(|v| v.as_u64()) .map(|v| v as u32); - - if input_tokens.is_some() || output_tokens.is_some() { - self.record_tokens(ctx, input_tokens, output_tokens, TokenSource::Actual); + if input.is_some() || output.is_some() { + self.record_tokens(ctx, input, output, TokenSource::Actual); return; } - } - - // 尝试从 Anthropic 格式响应中提取 Token - if let Some(usage) = response.get("usage") { - let input_tokens = usage + // Anthropic 格式 + let input = usage .get("input_tokens") .and_then(|v| v.as_u64()) .map(|v| v as u32); - let output_tokens = usage + let output = usage .get("output_tokens") .and_then(|v| v.as_u64()) .map(|v| v as u32); - - if input_tokens.is_some() || output_tokens.is_some() { - self.record_tokens(ctx, input_tokens, output_tokens, TokenSource::Actual); + if input.is_some() || output.is_some() { + self.record_tokens(ctx, input, output, TokenSource::Actual); } } } @@ -151,12 +119,8 @@ impl PipelineStep for TelemetryStep { ctx: &mut RequestContext, payload: &mut serde_json::Value, ) -> Result<(), StepError> { - // 记录成功的请求(同步方法,使用 parking_lot::RwLock) self.record_request(ctx, RequestStatus::Success, None); - - // 从响应中提取并记录 Token(同步方法) self.record_tokens_from_response(ctx, payload); - tracing::info!( "[TELEMETRY] request_id={} provider={:?} model={} duration_ms={}", ctx.request_id, @@ -164,7 +128,6 @@ impl PipelineStep for TelemetryStep { ctx.resolved_model, ctx.elapsed_ms() ); - Ok(()) } @@ -182,12 +145,9 @@ mod tests { let stats = Arc::new(RwLock::new(StatsAggregator::with_defaults())); let tokens = Arc::new(RwLock::new(TokenTracker::with_defaults())); let step = TelemetryStep::new(stats.clone(), tokens); - let ctx = RequestContext::new("claude-sonnet-4-5".to_string()); step.record_request(&ctx, RequestStatus::Success, None); - - let stats_guard = stats.read(); - assert_eq!(stats_guard.len(), 1); + assert_eq!(stats.read().len(), 1); } #[test] @@ -195,12 +155,9 @@ mod tests { let stats = Arc::new(RwLock::new(StatsAggregator::with_defaults())); let tokens = Arc::new(RwLock::new(TokenTracker::with_defaults())); let step = TelemetryStep::new(stats, tokens.clone()); - let ctx = RequestContext::new("claude-sonnet-4-5".to_string()); step.record_tokens(&ctx, Some(100), Some(50), TokenSource::Actual); - - let tokens_guard = tokens.read(); - assert_eq!(tokens_guard.len(), 1); + assert_eq!(tokens.read().len(), 1); } #[test] @@ -208,19 +165,11 @@ mod tests { let stats = Arc::new(RwLock::new(StatsAggregator::with_defaults())); let tokens = Arc::new(RwLock::new(TokenTracker::with_defaults())); let step = TelemetryStep::new(stats, tokens.clone()); - let ctx = RequestContext::new("claude-sonnet-4-5".to_string()); - let response = serde_json::json!({ - "usage": { - "prompt_tokens": 100, - "completion_tokens": 50 - } - }); - + let response = + serde_json::json!({"usage": {"prompt_tokens": 100, "completion_tokens": 50}}); step.record_tokens_from_response(&ctx, &response); - - let tokens_guard = tokens.read(); - assert_eq!(tokens_guard.len(), 1); + assert_eq!(tokens.read().len(), 1); } #[tokio::test] @@ -228,22 +177,11 @@ mod tests { let stats = Arc::new(RwLock::new(StatsAggregator::with_defaults())); let tokens = Arc::new(RwLock::new(TokenTracker::with_defaults())); let step = TelemetryStep::new(stats.clone(), tokens.clone()); - let mut ctx = RequestContext::new("claude-sonnet-4-5".to_string()); - let mut payload = serde_json::json!({ - "usage": { - "prompt_tokens": 100, - "completion_tokens": 50 - } - }); - - let result = step.execute(&mut ctx, &mut payload).await; - assert!(result.is_ok()); - - let stats_guard = stats.read(); - assert_eq!(stats_guard.len(), 1); - - let tokens_guard = tokens.read(); - assert_eq!(tokens_guard.len(), 1); + let mut payload = + serde_json::json!({"usage": {"prompt_tokens": 100, "completion_tokens": 50}}); + assert!(step.execute(&mut ctx, &mut payload).await.is_ok()); + assert_eq!(stats.read().len(), 1); + assert_eq!(tokens.read().len(), 1); } } diff --git a/src-tauri/src/processor/steps/traits.rs b/src-tauri/crates/processor/src/steps/traits.rs similarity index 71% rename from src-tauri/src/processor/steps/traits.rs rename to src-tauri/crates/processor/src/steps/traits.rs index 1f3ac444f..a387405e4 100644 --- a/src-tauri/src/processor/steps/traits.rs +++ b/src-tauri/crates/processor/src/steps/traits.rs @@ -1,48 +1,31 @@ //! 管道步骤 trait 定义 -//! -//! 定义所有管道步骤必须实现的接口 #![allow(dead_code)] -use crate::processor::RequestContext; use async_trait::async_trait; +use proxycast_core::processor::RequestContext; use thiserror::Error; /// 步骤错误 #[derive(Error, Debug, Clone)] pub enum StepError { - /// 认证错误 #[error("认证错误: {0}")] Auth(String), - - /// 路由错误 #[error("路由错误: {0}")] Routing(String), - - /// 注入错误 #[error("注入错误: {0}")] Injection(String), - - /// Provider 错误 #[error("Provider 错误: {0}")] Provider(String), - - /// 插件错误 #[error("插件错误: {plugin_name} - {message}")] Plugin { plugin_name: String, message: String, }, - - /// 遥测错误 #[error("遥测错误: {0}")] Telemetry(String), - - /// 超时错误 #[error("超时: {timeout_ms}ms")] Timeout { timeout_ms: u64 }, - - /// 内部错误 #[error("内部错误: {0}")] Internal(String), } @@ -64,28 +47,16 @@ impl StepError { } /// 管道步骤 trait -/// -/// 所有管道步骤必须实现此 trait #[async_trait] pub trait PipelineStep: Send + Sync { - /// 执行步骤 - /// - /// # Arguments - /// * `ctx` - 请求上下文 - /// * `payload` - 请求/响应负载 - /// - /// # Returns - /// 成功返回 `Ok(())`,失败返回 `Err(StepError)` async fn execute( &self, ctx: &mut RequestContext, payload: &mut serde_json::Value, ) -> Result<(), StepError>; - /// 获取步骤名称 fn name(&self) -> &str; - /// 检查步骤是否启用 fn is_enabled(&self) -> bool { true } diff --git a/src-tauri/src/processor/steps/mod.rs b/src-tauri/src/processor/steps/mod.rs index 047765f84..d9334d1e3 100644 --- a/src-tauri/src/processor/steps/mod.rs +++ b/src-tauri/src/processor/steps/mod.rs @@ -1,27 +1,4 @@ -//! 管道步骤模块 -//! -//! 定义请求处理管道中的各个步骤 +//! 管道步骤模块(re-export from proxycast-processor) -mod auth; -mod injection; -mod plugin; -mod provider; -mod routing; -mod telemetry; -mod traits; - -// 这些类型目前未在外部使用,但保留以供将来扩展 #[allow(unused_imports)] -pub use auth::AuthStep; -#[allow(unused_imports)] -pub use injection::InjectionStep; -#[allow(unused_imports)] -pub use plugin::{PluginPostStep, PluginPreStep}; -#[allow(unused_imports)] -pub use provider::ProviderStep; -#[allow(unused_imports)] -pub use routing::RoutingStep; -#[allow(unused_imports)] -pub use telemetry::TelemetryStep; -#[allow(unused_imports)] -pub use traits::PipelineStep; +pub use proxycast_processor::steps::*;