mirror of
https://github.com/aiclientproxy/proxycast.git
synced 2026-09-24 23:10:56 +08:00
release: bump version to v0.75.0
This commit is contained in:
@@ -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,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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",
|
||||
|
||||
Reference in New Issue
Block a user