release: bump version to v0.75.0

This commit is contained in:
coso
2026-02-28 22:42:28 +08:00
parent 672a94536a
commit 5cd3eda653
36 changed files with 7130 additions and 778 deletions
+181 -53
View File
@@ -5,22 +5,34 @@
use std::collections::HashMap;
use std::path::{Path, PathBuf};
use std::sync::Arc;
use std::time::Duration;
use dashmap::DashMap;
use tokio::sync::RwLock;
use tokio::time::timeout;
use super::loader::PluginLoader;
use super::task::{
PluginQueueStats, PluginTaskPolicy, PluginTaskRecord, PluginTaskState, PluginTaskTracker,
};
use super::types::{
HookResult, PluginConfig, PluginContext, PluginError, PluginInfo, PluginInstance, PluginStatus,
};
use crate::DynEmitter;
/// 插件管理器配置
#[derive(Debug, Clone)]
pub struct PluginManagerConfig {
/// 默认超时时间 (毫秒)
pub default_timeout_ms: u64,
/// 默认重试次数
pub default_max_retries: u32,
/// 默认重试退避基数 (毫秒)
pub default_retry_backoff_ms: u64,
/// 插件级并发上限
pub default_max_concurrency_per_plugin: usize,
/// 插件级队列长度上限
pub default_queue_limit_per_plugin: usize,
/// 任务记录保留数量
pub task_retention_limit: usize,
/// 是否启用插件系统
pub enabled: bool,
/// 最大并发插件数
@@ -31,6 +43,11 @@ impl Default for PluginManagerConfig {
fn default() -> Self {
Self {
default_timeout_ms: 5000,
default_max_retries: 2,
default_retry_backoff_ms: 300,
default_max_concurrency_per_plugin: 4,
default_queue_limit_per_plugin: 100,
task_retention_limit: 2000,
enabled: true,
max_plugins: 50,
}
@@ -47,6 +64,8 @@ pub struct PluginManager {
configs: DashMap<String, PluginConfig>,
/// 管理器配置
config: PluginManagerConfig,
/// 插件任务治理与跟踪
task_tracker: PluginTaskTracker,
}
impl PluginManager {
@@ -56,6 +75,7 @@ impl PluginManager {
loader: PluginLoader::new(plugins_dir),
plugins: DashMap::new(),
configs: DashMap::new(),
task_tracker: PluginTaskTracker::new(config.task_retention_limit),
config,
}
}
@@ -248,6 +268,46 @@ impl PluginManager {
infos
}
/// 设置插件任务事件发射器
pub async fn set_task_emitter(&self, emitter: DynEmitter) {
self.task_tracker.set_emitter(emitter).await;
}
/// 列出插件任务
pub fn list_tasks(
&self,
plugin_id: Option<&str>,
state: Option<PluginTaskState>,
limit: usize,
) -> Vec<PluginTaskRecord> {
self.task_tracker.list_tasks(plugin_id, state, limit)
}
/// 获取插件任务详情
pub fn get_task(&self, task_id: &str) -> Option<PluginTaskRecord> {
self.task_tracker.get_task(task_id)
}
/// 取消插件任务
pub fn cancel_task(&self, task_id: &str) -> bool {
self.task_tracker.cancel_task(task_id)
}
/// 获取插件队列统计
pub fn get_queue_stats(&self, plugin_id: Option<&str>) -> Vec<PluginQueueStats> {
self.task_tracker.queue_stats(plugin_id)
}
fn build_policy(&self, timeout_ms: u64) -> PluginTaskPolicy {
PluginTaskPolicy {
timeout_ms,
max_retries: self.config.default_max_retries,
retry_backoff_ms: self.config.default_retry_backoff_ms,
max_concurrency_per_plugin: self.config.default_max_concurrency_per_plugin,
queue_limit_per_plugin: self.config.default_queue_limit_per_plugin,
}
}
/// 执行请求前钩子 (带隔离)
pub async fn run_on_request(
&self,
@@ -267,24 +327,41 @@ impl PluginManager {
}
let timeout_ms = instance.config.timeout_ms;
let policy = self.build_policy(timeout_ms);
let plugin = instance.plugin.clone();
let plugin_name = plugin.name().to_string();
let base_ctx = ctx.clone();
let base_request = request.clone();
// 带超时执行
let result = match timeout(
Duration::from_millis(timeout_ms),
plugin.on_request(ctx, request),
)
.await
let result = match self
.task_tracker
.execute(&plugin_name, "on_request", policy, move |_attempt| {
let plugin = plugin.clone();
let mut attempt_ctx = base_ctx.clone();
let mut attempt_request = base_request.clone();
async move {
let hook_result = plugin
.on_request(&mut attempt_ctx, &mut attempt_request)
.await?;
Ok((hook_result, attempt_ctx, attempt_request))
}
})
.await
{
Ok(Ok(result)) => result,
Ok(Err(e)) => {
tracing::warn!("插件 {} on_request 执行失败: {}", plugin_name, e);
HookResult::failure(e.to_string(), timeout_ms)
Ok((hook_result, next_ctx, next_request)) => {
*ctx = next_ctx;
*request = next_request;
hook_result
}
Err(_) => {
tracing::warn!("插件 {} on_request 执行超时", plugin_name);
HookResult::failure(format!("执行超时 ({timeout_ms}ms)"), timeout_ms)
Err(failure) => {
tracing::warn!(
"插件 {} on_request 执行失败: {} (state={:?}, attempts={})",
plugin_name,
failure.message,
failure.state,
failure.attempts
);
HookResult::failure(failure.message, timeout_ms)
}
};
@@ -321,24 +398,41 @@ impl PluginManager {
}
let timeout_ms = instance.config.timeout_ms;
let policy = self.build_policy(timeout_ms);
let plugin = instance.plugin.clone();
let plugin_name = plugin.name().to_string();
let base_ctx = ctx.clone();
let base_response = response.clone();
// 带超时执行
let result = match timeout(
Duration::from_millis(timeout_ms),
plugin.on_response(ctx, response),
)
.await
let result = match self
.task_tracker
.execute(&plugin_name, "on_response", policy, move |_attempt| {
let plugin = plugin.clone();
let mut attempt_ctx = base_ctx.clone();
let mut attempt_response = base_response.clone();
async move {
let hook_result = plugin
.on_response(&mut attempt_ctx, &mut attempt_response)
.await?;
Ok((hook_result, attempt_ctx, attempt_response))
}
})
.await
{
Ok(Ok(result)) => result,
Ok(Err(e)) => {
tracing::warn!("插件 {} on_response 执行失败: {}", plugin_name, e);
HookResult::failure(e.to_string(), timeout_ms)
Ok((hook_result, next_ctx, next_response)) => {
*ctx = next_ctx;
*response = next_response;
hook_result
}
Err(_) => {
tracing::warn!("插件 {} on_response 执行超时", plugin_name);
HookResult::failure(format!("执行超时 ({timeout_ms}ms)"), timeout_ms)
Err(failure) => {
tracing::warn!(
"插件 {} on_response 执行失败: {} (state={:?}, attempts={})",
plugin_name,
failure.message,
failure.state,
failure.attempts
);
HookResult::failure(failure.message, timeout_ms)
}
};
@@ -371,24 +465,38 @@ impl PluginManager {
}
let timeout_ms = instance.config.timeout_ms;
let policy = self.build_policy(timeout_ms);
let plugin = instance.plugin.clone();
let plugin_name = plugin.name().to_string();
let base_ctx = ctx.clone();
let error_text = error.to_string();
// 带超时执行
let result = match timeout(
Duration::from_millis(timeout_ms),
plugin.on_error(ctx, error),
)
.await
let result = match self
.task_tracker
.execute(&plugin_name, "on_error", policy, move |_attempt| {
let plugin = plugin.clone();
let mut attempt_ctx = base_ctx.clone();
let error_text = error_text.clone();
async move {
let hook_result = plugin.on_error(&mut attempt_ctx, &error_text).await?;
Ok((hook_result, attempt_ctx))
}
})
.await
{
Ok(Ok(result)) => result,
Ok(Err(e)) => {
tracing::warn!("插件 {} on_error 执行失败: {}", plugin_name, e);
HookResult::failure(e.to_string(), timeout_ms)
Ok((hook_result, next_ctx)) => {
*ctx = next_ctx;
hook_result
}
Err(_) => {
tracing::warn!("插件 {} on_error 执行超时", plugin_name);
HookResult::failure(format!("执行超时 ({timeout_ms}ms)"), timeout_ms)
Err(failure) => {
tracing::warn!(
"插件 {} on_error 执行失败: {} (state={:?}, attempts={})",
plugin_name,
failure.message,
failure.state,
failure.attempts
);
HookResult::failure(failure.message, timeout_ms)
}
};
@@ -452,9 +560,16 @@ impl PluginManager {
.get(plugin_id)
.ok_or_else(|| PluginError::NotFound(plugin_id.to_string()))?;
// TODO: 检查插件是否实现了 PluginUI trait
// 目前返回空列表
Ok(Vec::new())
let policy = self.build_policy(self.config.default_timeout_ms);
self.task_tracker
.execute(plugin_id, "get_plugin_surfaces", policy, |_attempt| async {
Ok::<_, PluginError>(Vec::new())
})
.await
.map_err(|failure| PluginError::ExecutionError {
plugin_name: plugin_id.to_string(),
message: failure.message,
})
}
/// 处理插件 UI 操作
@@ -468,16 +583,29 @@ impl PluginManager {
.get(plugin_id)
.ok_or_else(|| PluginError::NotFound(plugin_id.to_string()))?;
// TODO: 将操作转发给插件的 handle_action 方法
// 目前返回空列表
tracing::debug!(
"收到插件 {} 的 UI 操作: {} (surface: {})",
plugin_id,
action.name,
action.surface_id
);
let action_name = action.name.clone();
let surface_id = action.surface_id.clone();
let policy = self.build_policy(self.config.default_timeout_ms);
Ok(Vec::new())
self.task_tracker
.execute(plugin_id, "handle_plugin_action", policy, move |_attempt| {
let action_name = action_name.clone();
let surface_id = surface_id.clone();
async move {
tracing::debug!(
"收到插件 {} 的 UI 操作: {} (surface: {})",
plugin_id,
action_name,
surface_id
);
Ok::<_, PluginError>(Vec::new())
}
})
.await
.map_err(|failure| PluginError::ExecutionError {
plugin_name: plugin_id.to_string(),
message: failure.message,
})
}
}
+5
View File
@@ -14,6 +14,7 @@ pub mod examples;
pub mod installer;
mod loader;
mod manager;
mod task;
mod types;
pub mod ui_builder;
pub mod ui_trait;
@@ -22,6 +23,10 @@ pub mod ui_types;
pub use binary_downloader::BinaryDownloader;
pub use loader::PluginLoader;
pub use manager::PluginManager;
pub use task::{
PluginQueueStats, PluginTaskError, PluginTaskEventPayload, PluginTaskFailure, PluginTaskPolicy,
PluginTaskRecord, PluginTaskState, PluginTaskTracker,
};
pub use types::{
BinaryComponentStatus, BinaryManifest, HookResult, PlatformBinaries, Plugin, PluginConfig,
PluginContext, PluginError, PluginInfo, PluginManifest, PluginState, PluginStatus, PluginType,
+855
View File
@@ -0,0 +1,855 @@
//! 插件任务执行治理模型
//!
//! 提供统一的任务状态、重试、超时、并发和队列治理能力。
use chrono::{DateTime, Utc};
use dashmap::DashMap;
use serde::{Deserialize, Serialize};
use std::future::Future;
use std::str::FromStr;
use std::sync::atomic::{AtomicBool, AtomicU64, AtomicUsize, Ordering};
use std::sync::Arc;
use std::time::{Duration, Instant};
use tokio::sync::{RwLock, Semaphore};
use tokio::time::{sleep, timeout};
use uuid::Uuid;
use crate::event_emit::DynEmitter;
use super::types::PluginError;
/// 插件任务状态
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum PluginTaskState {
Queued,
Running,
Retrying,
Succeeded,
Failed,
Cancelled,
TimedOut,
}
impl PluginTaskState {
pub fn is_terminal(self) -> bool {
matches!(
self,
PluginTaskState::Succeeded
| PluginTaskState::Failed
| PluginTaskState::Cancelled
| PluginTaskState::TimedOut
)
}
}
impl FromStr for PluginTaskState {
type Err = String;
fn from_str(s: &str) -> Result<Self, Self::Err> {
match s {
"queued" => Ok(Self::Queued),
"running" => Ok(Self::Running),
"retrying" => Ok(Self::Retrying),
"succeeded" => Ok(Self::Succeeded),
"failed" => Ok(Self::Failed),
"cancelled" => Ok(Self::Cancelled),
"timed_out" => Ok(Self::TimedOut),
_ => Err(format!("未知任务状态: {s}")),
}
}
}
/// 任务错误详情
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct PluginTaskError {
pub code: Option<String>,
pub message: String,
pub retryable: bool,
}
/// 插件任务执行策略
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct PluginTaskPolicy {
pub timeout_ms: u64,
pub max_retries: u32,
pub retry_backoff_ms: u64,
pub max_concurrency_per_plugin: usize,
pub queue_limit_per_plugin: usize,
}
impl Default for PluginTaskPolicy {
fn default() -> Self {
Self {
timeout_ms: 30_000,
max_retries: 2,
retry_backoff_ms: 300,
max_concurrency_per_plugin: 4,
queue_limit_per_plugin: 100,
}
}
}
/// 插件任务记录
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct PluginTaskRecord {
pub task_id: String,
pub plugin_id: String,
pub operation: String,
pub state: PluginTaskState,
pub attempt: u32,
pub max_retries: u32,
pub started_at: DateTime<Utc>,
pub ended_at: Option<DateTime<Utc>>,
pub duration_ms: Option<u64>,
pub error: Option<PluginTaskError>,
}
impl PluginTaskRecord {
fn new(task_id: String, plugin_id: String, operation: String, max_retries: u32) -> Self {
Self {
task_id,
plugin_id,
operation,
state: PluginTaskState::Queued,
attempt: 0,
max_retries,
started_at: Utc::now(),
ended_at: None,
duration_ms: None,
error: None,
}
}
fn finish_with_success(&mut self, attempt: u32, started: Instant) {
self.state = PluginTaskState::Succeeded;
self.attempt = attempt;
self.ended_at = Some(Utc::now());
self.duration_ms = Some(started.elapsed().as_millis() as u64);
self.error = None;
}
fn finish_with_failure(
&mut self,
state: PluginTaskState,
attempt: u32,
started: Instant,
error: PluginTaskError,
) {
self.state = state;
self.attempt = attempt;
self.ended_at = Some(Utc::now());
self.duration_ms = Some(started.elapsed().as_millis() as u64);
self.error = Some(error);
}
}
/// 前端消费的任务事件载荷
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct PluginTaskEventPayload {
pub plugin_id: String,
pub task_id: String,
pub operation: String,
pub state: PluginTaskState,
pub attempt: u32,
pub timestamp: String,
pub error: Option<PluginTaskError>,
}
impl PluginTaskEventPayload {
fn from_record(record: &PluginTaskRecord) -> Self {
Self {
plugin_id: record.plugin_id.clone(),
task_id: record.task_id.clone(),
operation: record.operation.clone(),
state: record.state,
attempt: record.attempt,
timestamp: Utc::now().to_rfc3339(),
error: record.error.clone(),
}
}
}
/// 任务失败返回
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct PluginTaskFailure {
pub task_id: String,
pub state: PluginTaskState,
pub attempts: u32,
pub message: String,
pub retryable: bool,
}
impl PluginTaskFailure {
fn new(
task_id: String,
state: PluginTaskState,
attempts: u32,
message: String,
retryable: bool,
) -> Self {
Self {
task_id,
state,
attempts,
message,
retryable,
}
}
}
/// 插件队列统计信息
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct PluginQueueStats {
pub plugin_id: String,
pub running: usize,
pub waiting: usize,
pub rejected: u64,
pub completed: u64,
pub failed: u64,
pub cancelled: u64,
pub timed_out: u64,
}
#[derive(Default)]
struct QueueMetrics {
running: AtomicUsize,
waiting: AtomicUsize,
rejected: AtomicU64,
completed: AtomicU64,
failed: AtomicU64,
cancelled: AtomicU64,
timed_out: AtomicU64,
}
impl QueueMetrics {
fn snapshot(&self, plugin_id: String) -> PluginQueueStats {
PluginQueueStats {
plugin_id,
running: self.running.load(Ordering::SeqCst),
waiting: self.waiting.load(Ordering::SeqCst),
rejected: self.rejected.load(Ordering::SeqCst),
completed: self.completed.load(Ordering::SeqCst),
failed: self.failed.load(Ordering::SeqCst),
cancelled: self.cancelled.load(Ordering::SeqCst),
timed_out: self.timed_out.load(Ordering::SeqCst),
}
}
}
struct RunningGuard {
metrics: Arc<QueueMetrics>,
}
impl RunningGuard {
fn new(metrics: Arc<QueueMetrics>) -> Self {
Self { metrics }
}
}
impl Drop for RunningGuard {
fn drop(&mut self) {
self.metrics.running.fetch_sub(1, Ordering::SeqCst);
}
}
/// 插件任务跟踪器
pub struct PluginTaskTracker {
tasks: DashMap<String, PluginTaskRecord>,
semaphores: DashMap<String, Arc<Semaphore>>,
queue_metrics: DashMap<String, Arc<QueueMetrics>>,
cancel_flags: DashMap<String, Arc<AtomicBool>>,
retention_limit: usize,
emitter: Arc<RwLock<Option<DynEmitter>>>,
}
impl Default for PluginTaskTracker {
fn default() -> Self {
Self::new(2_000)
}
}
impl PluginTaskTracker {
pub fn new(retention_limit: usize) -> Self {
Self {
tasks: DashMap::new(),
semaphores: DashMap::new(),
queue_metrics: DashMap::new(),
cancel_flags: DashMap::new(),
retention_limit: retention_limit.max(100),
emitter: Arc::new(RwLock::new(None)),
}
}
pub async fn set_emitter(&self, emitter: DynEmitter) {
let mut guard = self.emitter.write().await;
*guard = Some(emitter);
}
pub async fn clear_emitter(&self) {
let mut guard = self.emitter.write().await;
*guard = None;
}
pub fn get_task(&self, task_id: &str) -> Option<PluginTaskRecord> {
self.tasks.get(task_id).map(|entry| entry.value().clone())
}
pub fn list_tasks(
&self,
plugin_id: Option<&str>,
state: Option<PluginTaskState>,
limit: usize,
) -> Vec<PluginTaskRecord> {
let mut records: Vec<PluginTaskRecord> = self
.tasks
.iter()
.filter_map(|entry| {
let record = entry.value();
if let Some(plugin_id_filter) = plugin_id {
if record.plugin_id != plugin_id_filter {
return None;
}
}
if let Some(state_filter) = state {
if record.state != state_filter {
return None;
}
}
Some(record.clone())
})
.collect();
records.sort_by(|a, b| b.started_at.cmp(&a.started_at));
records.truncate(limit.max(1));
records
}
pub fn cancel_task(&self, task_id: &str) -> bool {
let Some(flag) = self.cancel_flags.get(task_id) else {
return false;
};
flag.store(true, Ordering::SeqCst);
true
}
pub fn queue_stats(&self, plugin_id: Option<&str>) -> Vec<PluginQueueStats> {
let mut items = Vec::new();
for entry in &self.queue_metrics {
if let Some(plugin_filter) = plugin_id {
if entry.key() != plugin_filter {
continue;
}
}
items.push(entry.value().snapshot(entry.key().clone()));
}
items.sort_by(|a, b| a.plugin_id.cmp(&b.plugin_id));
items
}
pub async fn execute<T, F, Fut>(
&self,
plugin_id: &str,
operation: &str,
mut policy: PluginTaskPolicy,
mut operation_fn: F,
) -> Result<T, PluginTaskFailure>
where
T: Send + 'static,
F: FnMut(u32) -> Fut + Send,
Fut: Future<Output = Result<T, PluginError>> + Send,
{
if policy.max_concurrency_per_plugin == 0 {
policy.max_concurrency_per_plugin = 1;
}
if policy.queue_limit_per_plugin == 0 {
policy.queue_limit_per_plugin = 1;
}
if policy.timeout_ms == 0 {
policy.timeout_ms = 1;
}
let task_id = Uuid::new_v4().to_string();
let total_started = Instant::now();
let mut record = PluginTaskRecord::new(
task_id.clone(),
plugin_id.to_string(),
operation.to_string(),
policy.max_retries,
);
self.upsert_task(record.clone());
self.emit_task_event(&record).await;
let cancel_flag = Arc::new(AtomicBool::new(false));
self.cancel_flags
.insert(task_id.clone(), Arc::clone(&cancel_flag));
let semaphore = self
.semaphores
.entry(plugin_id.to_string())
.or_insert_with(|| Arc::new(Semaphore::new(policy.max_concurrency_per_plugin)))
.clone();
let metrics = self
.queue_metrics
.entry(plugin_id.to_string())
.or_insert_with(|| Arc::new(QueueMetrics::default()))
.clone();
let waiting_now = metrics.waiting.fetch_add(1, Ordering::SeqCst) + 1;
if waiting_now > policy.queue_limit_per_plugin {
metrics.waiting.fetch_sub(1, Ordering::SeqCst);
metrics.rejected.fetch_add(1, Ordering::SeqCst);
let error = PluginTaskError {
code: Some("QUEUE_LIMIT_EXCEEDED".to_string()),
message: format!(
"插件 {plugin_id} 队列已满 (limit={})",
policy.queue_limit_per_plugin
),
retryable: false,
};
record.finish_with_failure(PluginTaskState::Failed, 0, total_started, error.clone());
self.upsert_task(record.clone());
self.cancel_flags.remove(&task_id);
self.emit_task_event(&record).await;
return Err(PluginTaskFailure::new(
task_id,
PluginTaskState::Failed,
0,
error.message,
false,
));
}
let permit = match semaphore.acquire_owned().await {
Ok(permit) => permit,
Err(err) => {
metrics.waiting.fetch_sub(1, Ordering::SeqCst);
let error = PluginTaskError {
code: Some("SEMAPHORE_CLOSED".to_string()),
message: format!("无法获取插件执行许可: {err}"),
retryable: true,
};
record.finish_with_failure(
PluginTaskState::Failed,
0,
total_started,
error.clone(),
);
self.upsert_task(record.clone());
self.cancel_flags.remove(&task_id);
self.emit_task_event(&record).await;
return Err(PluginTaskFailure::new(
task_id,
PluginTaskState::Failed,
0,
error.message,
true,
));
}
};
metrics.waiting.fetch_sub(1, Ordering::SeqCst);
metrics.running.fetch_add(1, Ordering::SeqCst);
let running_guard = RunningGuard::new(Arc::clone(&metrics));
if cancel_flag.load(Ordering::SeqCst) {
let error = PluginTaskError {
code: Some("TASK_CANCELLED".to_string()),
message: "任务已取消".to_string(),
retryable: false,
};
record.finish_with_failure(PluginTaskState::Cancelled, 0, total_started, error.clone());
metrics.cancelled.fetch_add(1, Ordering::SeqCst);
self.upsert_task(record.clone());
self.cancel_flags.remove(&task_id);
self.emit_task_event(&record).await;
drop(permit);
drop(running_guard);
return Err(PluginTaskFailure::new(
task_id,
PluginTaskState::Cancelled,
0,
error.message,
false,
));
}
let mut attempt: u32 = 0;
loop {
attempt += 1;
record.state = if attempt == 1 {
PluginTaskState::Running
} else {
PluginTaskState::Retrying
};
record.attempt = attempt;
record.error = None;
self.upsert_task(record.clone());
self.emit_task_event(&record).await;
let timed_result = timeout(
Duration::from_millis(policy.timeout_ms),
operation_fn(attempt),
)
.await;
match timed_result {
Ok(Ok(value)) => {
record.finish_with_success(attempt, total_started);
metrics.completed.fetch_add(1, Ordering::SeqCst);
self.upsert_task(record.clone());
self.cancel_flags.remove(&task_id);
self.emit_task_event(&record).await;
drop(permit);
drop(running_guard);
return Ok(value);
}
Ok(Err(err)) => {
let retryable = is_retryable_error(&err);
let can_retry = retryable
&& attempt <= policy.max_retries
&& !cancel_flag.load(Ordering::SeqCst);
if can_retry {
let backoff = backoff_duration(policy.retry_backoff_ms, attempt);
sleep(backoff).await;
continue;
}
let state = if cancel_flag.load(Ordering::SeqCst) {
PluginTaskState::Cancelled
} else {
PluginTaskState::Failed
};
let error = PluginTaskError {
code: classify_error_code(&err),
message: err.to_string(),
retryable,
};
record.finish_with_failure(state, attempt, total_started, error.clone());
match state {
PluginTaskState::Cancelled => {
metrics.cancelled.fetch_add(1, Ordering::SeqCst);
}
PluginTaskState::Failed => {
metrics.failed.fetch_add(1, Ordering::SeqCst);
}
_ => {}
}
self.upsert_task(record.clone());
self.cancel_flags.remove(&task_id);
self.emit_task_event(&record).await;
drop(permit);
drop(running_guard);
return Err(PluginTaskFailure::new(
task_id,
state,
attempt,
error.message,
retryable,
));
}
Err(_) => {
let can_retry =
attempt <= policy.max_retries && !cancel_flag.load(Ordering::SeqCst);
if can_retry {
let backoff = backoff_duration(policy.retry_backoff_ms, attempt);
sleep(backoff).await;
continue;
}
let state = if cancel_flag.load(Ordering::SeqCst) {
PluginTaskState::Cancelled
} else {
PluginTaskState::TimedOut
};
let error = PluginTaskError {
code: Some(if state == PluginTaskState::TimedOut {
"TIMEOUT".to_string()
} else {
"TASK_CANCELLED".to_string()
}),
message: if state == PluginTaskState::TimedOut {
format!("执行超时: {}ms", policy.timeout_ms)
} else {
"任务已取消".to_string()
},
retryable: state == PluginTaskState::TimedOut,
};
record.finish_with_failure(state, attempt, total_started, error.clone());
match state {
PluginTaskState::TimedOut => {
metrics.timed_out.fetch_add(1, Ordering::SeqCst);
}
PluginTaskState::Cancelled => {
metrics.cancelled.fetch_add(1, Ordering::SeqCst);
}
_ => {}
}
self.upsert_task(record.clone());
self.cancel_flags.remove(&task_id);
self.emit_task_event(&record).await;
drop(permit);
drop(running_guard);
return Err(PluginTaskFailure::new(
task_id,
state,
attempt,
error.message,
state == PluginTaskState::TimedOut,
));
}
}
}
}
fn upsert_task(&self, record: PluginTaskRecord) {
self.tasks.insert(record.task_id.clone(), record);
self.trim_retention();
}
fn trim_retention(&self) {
if self.tasks.len() <= self.retention_limit {
return;
}
while self.tasks.len() > self.retention_limit {
let oldest_id = self
.tasks
.iter()
.min_by_key(|entry| entry.value().started_at)
.map(|entry| entry.key().clone());
let Some(oldest_id) = oldest_id else {
break;
};
self.tasks.remove(&oldest_id);
self.cancel_flags.remove(&oldest_id);
}
}
async fn emit_task_event(&self, record: &PluginTaskRecord) {
let payload = PluginTaskEventPayload::from_record(record);
let Ok(value) = serde_json::to_value(payload) else {
return;
};
let emitter = self.emitter.read().await.clone();
if let Some(emitter) = emitter {
let _ = emitter.emit_event("plugin-task-event", &value);
}
}
}
fn backoff_duration(base_ms: u64, attempt: u32) -> Duration {
let factor = 2_u64.saturating_pow(attempt.saturating_sub(1));
Duration::from_millis(base_ms.max(1).saturating_mul(factor))
}
fn classify_error_code(err: &PluginError) -> Option<String> {
match err {
PluginError::Timeout { .. } => Some("TIMEOUT".to_string()),
PluginError::Disabled(_) => Some("PLUGIN_DISABLED".to_string()),
PluginError::NotFound(_) => Some("PLUGIN_NOT_FOUND".to_string()),
PluginError::ConfigError(_) => Some("CONFIG_ERROR".to_string()),
PluginError::LoadError(_) => Some("LOAD_ERROR".to_string()),
PluginError::InitError(_) => Some("INIT_ERROR".to_string()),
PluginError::ExecutionError { message, .. } => {
if message.contains("401") || message.contains("403") {
Some("AUTH_ERROR".to_string())
} else if message.contains("429") {
Some("RATE_LIMIT".to_string())
} else if message.contains("500")
|| message.contains("502")
|| message.contains("503")
|| message.contains("504")
{
Some("UPSTREAM_5XX".to_string())
} else {
Some("EXECUTION_ERROR".to_string())
}
}
_ => Some("UNKNOWN".to_string()),
}
}
fn is_retryable_error(err: &PluginError) -> bool {
match err {
PluginError::Timeout { .. } => true,
PluginError::ExecutionError { message, .. } => {
let lower = message.to_lowercase();
message.contains("429")
|| message.contains("500")
|| message.contains("502")
|| message.contains("503")
|| message.contains("504")
|| lower.contains("timeout")
|| lower.contains("temporar")
|| lower.contains("connection")
|| lower.contains("network")
}
_ => false,
}
}
#[cfg(test)]
mod tests {
use super::*;
use tokio::sync::Mutex;
#[tokio::test]
async fn test_execute_success_and_record_terminal_state() {
let tracker = PluginTaskTracker::new(100);
let policy = PluginTaskPolicy::default();
let result = tracker
.execute("demo-plugin", "on_request", policy, |_attempt| async move {
Ok::<_, PluginError>("ok".to_string())
})
.await
.expect("执行应成功");
assert_eq!(result, "ok");
let tasks = tracker.list_tasks(Some("demo-plugin"), None, 10);
assert_eq!(tasks.len(), 1);
assert_eq!(tasks[0].state, PluginTaskState::Succeeded);
assert_eq!(tasks[0].attempt, 1);
}
#[tokio::test]
async fn test_retry_then_success() {
let tracker = PluginTaskTracker::new(100);
let policy = PluginTaskPolicy {
max_retries: 2,
retry_backoff_ms: 1,
..PluginTaskPolicy::default()
};
let counter = Arc::new(Mutex::new(0_u32));
let result = tracker
.execute("retry-plugin", "on_response", policy, {
let counter = Arc::clone(&counter);
move |_attempt| {
let counter = Arc::clone(&counter);
async move {
let mut lock = counter.lock().await;
*lock += 1;
if *lock < 2 {
Err(PluginError::ExecutionError {
plugin_name: "retry-plugin".to_string(),
message: "503 upstream unavailable".to_string(),
})
} else {
Ok::<_, PluginError>("recovered".to_string())
}
}
}
})
.await
.expect("应在重试后成功");
assert_eq!(result, "recovered");
let tasks = tracker.list_tasks(Some("retry-plugin"), None, 10);
assert_eq!(tasks[0].state, PluginTaskState::Succeeded);
assert_eq!(tasks[0].attempt, 2);
}
#[tokio::test]
async fn test_timeout_to_terminal_state() {
let tracker = PluginTaskTracker::new(100);
let policy = PluginTaskPolicy {
timeout_ms: 30,
max_retries: 0,
..PluginTaskPolicy::default()
};
let result = tracker
.execute(
"timeout-plugin",
"on_error",
policy,
|_attempt| async move {
sleep(Duration::from_millis(80)).await;
Ok::<_, PluginError>("late".to_string())
},
)
.await;
assert!(result.is_err());
let err = result.expect_err("应超时失败");
assert_eq!(err.state, PluginTaskState::TimedOut);
let tasks = tracker.list_tasks(Some("timeout-plugin"), None, 10);
assert_eq!(tasks[0].state, PluginTaskState::TimedOut);
}
#[tokio::test]
async fn test_queue_limit_rejection() {
let tracker = Arc::new(PluginTaskTracker::new(100));
let policy = PluginTaskPolicy {
max_concurrency_per_plugin: 1,
queue_limit_per_plugin: 1,
timeout_ms: 500,
max_retries: 0,
..PluginTaskPolicy::default()
};
let tracker_a = Arc::clone(&tracker);
let policy_a = policy.clone();
let t1 = tokio::spawn(async move {
tracker_a
.execute(
"queue-plugin",
"on_request",
policy_a,
|_attempt| async move {
sleep(Duration::from_millis(150)).await;
Ok::<_, PluginError>("t1".to_string())
},
)
.await
});
sleep(Duration::from_millis(20)).await;
let tracker_b = Arc::clone(&tracker);
let policy_b = policy.clone();
let t2 = tokio::spawn(async move {
tracker_b
.execute(
"queue-plugin",
"on_request",
policy_b,
|_attempt| async move {
sleep(Duration::from_millis(80)).await;
Ok::<_, PluginError>("t2".to_string())
},
)
.await
});
sleep(Duration::from_millis(20)).await;
let t3 = tracker
.execute(
"queue-plugin",
"on_request",
policy,
|_attempt| async move { Ok::<_, PluginError>("t3".to_string()) },
)
.await;
let r1 = t1.await.expect("join t1");
let r2 = t2.await.expect("join t2");
assert!(r1.is_ok());
assert!(r2.is_ok());
assert!(t3.is_err());
let stats = tracker.queue_stats(Some("queue-plugin"));
assert_eq!(stats.len(), 1);
assert!(stats[0].rejected >= 1);
}
}
@@ -266,6 +266,7 @@ mod property_tests {
default_timeout_ms: 1000,
enabled: true,
max_plugins: 10,
..PluginManagerConfig::default()
};
let manager = PluginManager::new(temp_dir.path().to_path_buf(), config);
@@ -297,6 +298,7 @@ mod property_tests {
default_timeout_ms: 1000,
enabled: false, // 禁用插件系统
max_plugins: 10,
..PluginManagerConfig::default()
};
let manager = PluginManager::new(temp_dir.path().to_path_buf(), config);
@@ -744,7 +744,10 @@ impl CodexProvider {
}
// 3. OAuth 刷新流程(标准流程)
let refresh_token = self.credentials.refresh_token.as_ref().unwrap();
let refresh_token =
self.credentials.refresh_token.as_ref().ok_or_else(|| {
create_config_error("OAuth 刷新令牌不可用 (refresh_token is None)")
})?;
tracing::info!("[CODEX] 正在刷新 access token");
@@ -2351,7 +2354,7 @@ pub async fn start_codex_oauth_server_and_get_url() -> Result<
let uuid = Uuid::new_v4().to_string();
let timestamp = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.unwrap_or_default()
.as_secs();
let filename = format!("codex_{}_{}.json", &uuid[..8], timestamp);
let creds_file_path = creds_dir.join(&filename);
+12 -12
View File
@@ -5,7 +5,7 @@
use super::batch::{BatchTask, BatchTaskStatus};
use super::template::TaskTemplate;
use anyhow::{Context, Result};
use proxycast_core::database::DbConnection;
use proxycast_core::database::{lock_db, DbConnection};
use rusqlite::{params, OptionalExtension};
use uuid::Uuid;
@@ -15,7 +15,7 @@ pub struct BatchTaskDao;
impl BatchTaskDao {
/// 初始化数据库表
pub fn init_tables(db: &DbConnection) -> Result<()> {
let conn = db.lock().unwrap();
let conn = lock_db(db).map_err(|e| anyhow::anyhow!(e))?;
// 创建批量任务表
conn.execute(
@@ -69,7 +69,7 @@ impl BatchTaskDao {
/// 保存批量任务
pub fn save(db: &DbConnection, batch_task: &BatchTask) -> Result<()> {
let conn = db.lock().unwrap();
let conn = lock_db(db).map_err(|e| anyhow::anyhow!(e))?;
let options_json = serde_json::to_string(&batch_task.options)?;
let tasks_json = serde_json::to_string(&batch_task.tasks)?;
@@ -104,7 +104,7 @@ impl BatchTaskDao {
/// 根据 ID 查询批量任务
pub fn get_by_id(db: &DbConnection, id: &Uuid) -> Result<Option<BatchTask>> {
let conn = db.lock().unwrap();
let conn = lock_db(db).map_err(|e| anyhow::anyhow!(e))?;
let mut stmt = conn.prepare(
"SELECT id, name, template_id, status, options_json, tasks_json, results_json,
@@ -185,7 +185,7 @@ impl BatchTaskDao {
/// 查询所有批量任务
pub fn list_all(db: &DbConnection, limit: usize) -> Result<Vec<BatchTask>> {
let conn = db.lock().unwrap();
let conn = lock_db(db).map_err(|e| anyhow::anyhow!(e))?;
let mut stmt = conn.prepare(
"SELECT id, name, template_id, status, options_json, tasks_json, results_json,
@@ -268,7 +268,7 @@ impl BatchTaskDao {
/// 删除批量任务
pub fn delete(db: &DbConnection, id: &Uuid) -> Result<bool> {
let conn = db.lock().unwrap();
let conn = lock_db(db).map_err(|e| anyhow::anyhow!(e))?;
let affected = conn.execute(
"DELETE FROM batch_tasks WHERE id = ?1",
@@ -280,7 +280,7 @@ impl BatchTaskDao {
/// 更新批量任务状态
pub fn update_status(db: &DbConnection, id: &Uuid, status: BatchTaskStatus) -> Result<()> {
let conn = db.lock().unwrap();
let conn = lock_db(db).map_err(|e| anyhow::anyhow!(e))?;
conn.execute(
"UPDATE batch_tasks SET status = ?1 WHERE id = ?2",
@@ -301,7 +301,7 @@ impl BatchTaskDao {
started_at: Option<chrono::DateTime<chrono::Utc>>,
completed_at: Option<chrono::DateTime<chrono::Utc>>,
) -> Result<()> {
let conn = db.lock().unwrap();
let conn = lock_db(db).map_err(|e| anyhow::anyhow!(e))?;
let results_json = if results.is_empty() {
None
@@ -330,7 +330,7 @@ pub struct TemplateDao;
impl TemplateDao {
/// 保存模板
pub fn save(db: &DbConnection, template: &TaskTemplate) -> Result<()> {
let conn = db.lock().unwrap();
let conn = lock_db(db).map_err(|e| anyhow::anyhow!(e))?;
conn.execute(
"INSERT OR REPLACE INTO batch_templates
@@ -357,7 +357,7 @@ impl TemplateDao {
/// 根据 ID 查询模板
pub fn get_by_id(db: &DbConnection, id: &Uuid) -> Result<Option<TaskTemplate>> {
let conn = db.lock().unwrap();
let conn = lock_db(db).map_err(|e| anyhow::anyhow!(e))?;
let mut stmt = conn.prepare(
"SELECT id, name, description, model, system_prompt, user_message_template,
@@ -391,7 +391,7 @@ impl TemplateDao {
/// 查询所有模板
pub fn list_all(db: &DbConnection) -> Result<Vec<TaskTemplate>> {
let conn = db.lock().unwrap();
let conn = lock_db(db).map_err(|e| anyhow::anyhow!(e))?;
let mut stmt = conn.prepare(
"SELECT id, name, description, model, system_prompt, user_message_template,
@@ -429,7 +429,7 @@ impl TemplateDao {
/// 删除模板
pub fn delete(db: &DbConnection, id: &Uuid) -> Result<bool> {
let conn = db.lock().unwrap();
let conn = lock_db(db).map_err(|e| anyhow::anyhow!(e))?;
let affected = conn.execute(
"DELETE FROM batch_templates WHERE id = ?1",