mirror of
https://github.com/aiclientproxy/proxycast.git
synced 2026-09-24 23:10:56 +08:00
refactor: 迁移 processor steps 到独立 proxycast-processor crate
- 创建 crates/processor/ 独立 crate(~2207 行) - 迁移 7 个 step 模块: auth, injection, plugin, provider, routing, telemetry, traits - 迁移 RequestProcessor 基础结构(processor.rs) - 主 crate processor/steps/mod.rs 改为 re-export 层 - RequestProcessor 保留在主 crate(依赖 Tauri 相关类型) - 30 个 crate 单元测试 + 8 个主 crate 测试全部通过
This commit is contained in:
Generated
+20
@@ -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"
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
@@ -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::*;
|
||||
@@ -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<RwLock<Router>>,
|
||||
/// 模型映射器
|
||||
pub mapper: Arc<RwLock<ModelMapper>>,
|
||||
/// 参数注入器
|
||||
pub injector: Arc<RwLock<Injector>>,
|
||||
/// 重试器
|
||||
pub retrier: Arc<Retrier>,
|
||||
/// 故障转移器
|
||||
pub failover: Arc<Failover>,
|
||||
/// 超时控制器
|
||||
pub timeout: Arc<TimeoutController>,
|
||||
/// 插件管理器
|
||||
pub plugins: Arc<PluginManager>,
|
||||
/// 统计聚合器(使用 parking_lot::RwLock 以支持与 TelemetryState 共享)
|
||||
pub stats: Arc<ParkingLotRwLock<StatsAggregator>>,
|
||||
/// Token 追踪器(使用 parking_lot::RwLock 以支持与 TelemetryState 共享)
|
||||
pub tokens: Arc<ParkingLotRwLock<TokenTracker>>,
|
||||
/// 凭证池服务
|
||||
pub pool_service: Arc<ProviderPoolService>,
|
||||
/// 热重载协调锁(避免配置更新期间请求读取不一致的配置)
|
||||
pub reload_lock: Arc<RwLock<()>>,
|
||||
}
|
||||
|
||||
impl RequestProcessor {
|
||||
/// 创建新的请求处理器
|
||||
pub fn new(
|
||||
router: Arc<RwLock<Router>>,
|
||||
mapper: Arc<RwLock<ModelMapper>>,
|
||||
injector: Arc<RwLock<Injector>>,
|
||||
retrier: Arc<Retrier>,
|
||||
failover: Arc<Failover>,
|
||||
timeout: Arc<TimeoutController>,
|
||||
plugins: Arc<PluginManager>,
|
||||
stats: Arc<ParkingLotRwLock<StatsAggregator>>,
|
||||
tokens: Arc<ParkingLotRwLock<TokenTracker>>,
|
||||
pool_service: Arc<ProviderPoolService>,
|
||||
) -> 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<ProviderPoolService>) -> 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<ProviderPoolService>,
|
||||
stats: Arc<ParkingLotRwLock<StatsAggregator>>,
|
||||
tokens: Arc<ParkingLotRwLock<TokenTracker>>,
|
||||
) -> 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<ProviderType>, 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<ProviderType> {
|
||||
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<ProviderType> {
|
||||
self.resolve_model_for_context(ctx).await;
|
||||
self.route_for_context(ctx).await
|
||||
}
|
||||
}
|
||||
+8
-23
@@ -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());
|
||||
}
|
||||
}
|
||||
+5
-23
@@ -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<RwLock<Injector>>,
|
||||
/// 是否启用
|
||||
enabled: Arc<RwLock<bool>>,
|
||||
}
|
||||
|
||||
impl InjectionStep {
|
||||
/// 创建新的注入步骤
|
||||
pub fn new(injector: Arc<RwLock<Injector>>) -> Self {
|
||||
Self {
|
||||
injector,
|
||||
@@ -30,12 +23,10 @@ impl InjectionStep {
|
||||
}
|
||||
}
|
||||
|
||||
/// 设置是否启用
|
||||
pub fn with_enabled(self, enabled: Arc<RwLock<bool>>) -> 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());
|
||||
}
|
||||
}
|
||||
@@ -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};
|
||||
+7
-42
@@ -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<PluginManager>,
|
||||
}
|
||||
|
||||
impl PluginPreStep {
|
||||
/// 创建新的插件前置步骤
|
||||
pub fn new(plugins: Arc<PluginManager>) -> 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::<Vec<_>>()),
|
||||
);
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
@@ -72,24 +56,18 @@ impl PipelineStep for PluginPreStep {
|
||||
}
|
||||
|
||||
/// 插件后置钩子步骤
|
||||
///
|
||||
/// 在 Provider 调用后执行所有启用插件的 on_response 钩子
|
||||
pub struct PluginPostStep {
|
||||
/// 插件管理器
|
||||
plugins: Arc<PluginManager>,
|
||||
}
|
||||
|
||||
impl PluginPostStep {
|
||||
/// 创建新的插件后置步骤
|
||||
pub fn new(plugins: Arc<PluginManager>) -> 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::<Vec<_>>()),
|
||||
);
|
||||
}
|
||||
|
||||
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());
|
||||
}
|
||||
}
|
||||
+46
-172
@@ -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<String>,
|
||||
}
|
||||
|
||||
/// Provider 调用错误
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct ProviderCallError {
|
||||
/// 错误消息
|
||||
pub message: String,
|
||||
/// HTTP 状态码(如果有)
|
||||
pub status_code: Option<u16>,
|
||||
/// 是否可重试
|
||||
pub retryable: bool,
|
||||
/// 是否应触发故障转移
|
||||
pub should_failover: bool,
|
||||
}
|
||||
|
||||
impl ProviderCallError {
|
||||
/// 创建可重试错误
|
||||
pub fn retryable(message: impl Into<String>, status_code: Option<u16>) -> Self {
|
||||
Self {
|
||||
message: message.into(),
|
||||
@@ -53,7 +44,6 @@ impl ProviderCallError {
|
||||
}
|
||||
}
|
||||
|
||||
/// 创建需要故障转移的错误
|
||||
pub fn failover(message: impl Into<String>, status_code: Option<u16>) -> Self {
|
||||
Self {
|
||||
message: message.into(),
|
||||
@@ -63,7 +53,6 @@ impl ProviderCallError {
|
||||
}
|
||||
}
|
||||
|
||||
/// 创建不可恢复错误
|
||||
pub fn fatal(message: impl Into<String>, status_code: Option<u16>) -> 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<Retrier>,
|
||||
/// 故障转移器
|
||||
failover: Arc<Failover>,
|
||||
/// 超时控制器
|
||||
timeout: Arc<TimeoutController>,
|
||||
/// 凭证池服务
|
||||
pool_service: Arc<ProviderPoolService>,
|
||||
}
|
||||
|
||||
impl ProviderStep {
|
||||
/// 创建新的 Provider 步骤
|
||||
pub fn new(
|
||||
retrier: Arc<Retrier>,
|
||||
failover: Arc<Failover>,
|
||||
@@ -109,7 +90,6 @@ impl ProviderStep {
|
||||
}
|
||||
}
|
||||
|
||||
/// 使用默认配置创建
|
||||
pub fn with_defaults(pool_service: Arc<ProviderPoolService>) -> 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<F, Fut>(
|
||||
&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<F>(
|
||||
&self,
|
||||
ctx: &RequestContext,
|
||||
@@ -243,7 +186,6 @@ impl ProviderStep {
|
||||
F: Future<Output = Result<ProviderCallResult, ProviderCallError>>,
|
||||
{
|
||||
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<ProviderType> {
|
||||
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<F, Fut>(
|
||||
&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<ProviderCallResult, ProviderCallError> = 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<ProviderCallResult, ProviderCallError> =
|
||||
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(),
|
||||
+10
-44
@@ -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<RwLock<Router>>,
|
||||
/// 模型映射器
|
||||
mapper: Arc<RwLock<ModelMapper>>,
|
||||
/// 默认 Provider
|
||||
default_provider: Arc<RwLock<String>>,
|
||||
}
|
||||
|
||||
impl RoutingStep {
|
||||
/// 创建新的路由步骤
|
||||
pub fn new(
|
||||
router: Arc<RwLock<Router>>,
|
||||
mapper: Arc<RwLock<ModelMapper>>,
|
||||
@@ -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<ProviderType, StepError> {
|
||||
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");
|
||||
+25
-87
@@ -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<RwLock<StatsAggregator>>,
|
||||
/// Token 追踪器(使用 parking_lot::RwLock 以支持与 TelemetryState 共享)
|
||||
tokens: Arc<RwLock<TokenTracker>>,
|
||||
}
|
||||
|
||||
impl TelemetryStep {
|
||||
/// 创建新的统计记录步骤
|
||||
pub fn new(stats: Arc<RwLock<StatsAggregator>>, tokens: Arc<RwLock<TokenTracker>>) -> 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);
|
||||
}
|
||||
}
|
||||
+1
-30
@@ -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
|
||||
}
|
||||
@@ -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::*;
|
||||
|
||||
Reference in New Issue
Block a user