mirror of
https://wget.la/https://github.com/leookun/cursor-byok
synced 2026-10-06 05:12:03 +08:00
refactor: replace Button with TruncatedButton in CursorModelCards and PluginManagementPage
- Updated the UI components in CursorModelCards and PluginManagementPage to use TruncatedButton for better text handling and display. - Adjusted styles in CursorSettings and PluginManagementPage to ensure proper button layout and responsiveness. - Added new ActionMenu component for handling additional actions in PluginManagementPage. - Enhanced localization files to include new strings for the ActionMenu and TruncatedButton components.
This commit is contained in:
+5
-1
@@ -39,7 +39,11 @@ impl App {
|
||||
let assets = PromptAssets::embedded()?;
|
||||
let compiler = PromptCompiler::new(assets);
|
||||
let plugin_runtime = PluginRuntime::managed()?;
|
||||
let plugins = PluginRegistry::managed(store.clone(), plugin_runtime.clone())?;
|
||||
let plugins = PluginRegistry::managed(
|
||||
store.clone(),
|
||||
plugin_runtime.clone(),
|
||||
config.app_version.clone(),
|
||||
)?;
|
||||
let provider = std::sync::Arc::new(ProviderRouter::new(
|
||||
store.clone(),
|
||||
plugins.clone(),
|
||||
|
||||
@@ -58,6 +58,8 @@ pub struct Config {
|
||||
pub provider_stream_idle_timeout: Duration,
|
||||
pub console: Option<ConsoleSource>,
|
||||
pub use_persisted_ports: bool,
|
||||
/// 面向用户的应用版本;桌面壳会覆盖为自身版本,用于插件 minAppVersion 门控。
|
||||
pub app_version: String,
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
@@ -109,6 +111,7 @@ impl Config {
|
||||
provider_stream_idle_timeout: DEFAULT_PROVIDER_STREAM_IDLE_TIMEOUT,
|
||||
console,
|
||||
use_persisted_ports: false,
|
||||
app_version: env!("CARGO_PKG_VERSION").into(),
|
||||
})
|
||||
}
|
||||
|
||||
@@ -122,6 +125,7 @@ impl Config {
|
||||
provider_stream_idle_timeout: DEFAULT_PROVIDER_STREAM_IDLE_TIMEOUT,
|
||||
console: None,
|
||||
use_persisted_ports: true,
|
||||
app_version: env!("CARGO_PKG_VERSION").into(),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
+108
-10
@@ -1,10 +1,10 @@
|
||||
//! Materializes built-in plugins bundled in the binary into the managed dir.
|
||||
use std::path::PathBuf;
|
||||
//! Pre-installs bundled built-in plugins into the user's installed directory.
|
||||
use std::path::Path;
|
||||
|
||||
use super::definition::write_if_changed;
|
||||
use crate::{config, Result};
|
||||
use crate::Result;
|
||||
|
||||
/// 随二进制打包的内置插件文件;发布构建没有源码目录,靠这里落盘。
|
||||
/// 随二进制打包的内置插件文件;发布构建没有源码目录,靠这里预装。
|
||||
const CODEX_AUTH: &[(&str, &str)] = &[
|
||||
(
|
||||
"plugin.json",
|
||||
@@ -57,14 +57,42 @@ const CODEX_AUTH: &[(&str, &str)] = &[
|
||||
),
|
||||
];
|
||||
|
||||
/// 把内置插件写入受管目录并返回该目录,作为插件目录的扫描根之一。
|
||||
pub(super) fn materialize() -> Result<PathBuf> {
|
||||
let root = config::managed_data_dir()?.join("plugins/build-in");
|
||||
write_plugin(&root.join("codex-auth"), CODEX_AUTH)?;
|
||||
Ok(root)
|
||||
const PLUGINS: &[(&str, &[(&str, &str)])] = &[("codex-auth", CODEX_AUTH)];
|
||||
|
||||
/// 把内置插件预装到 installed 目录。manifest 的 version 是缓存键:
|
||||
/// 版本一致时零写盘;版本变化时整目录同步并清理旧版本残留文件。
|
||||
pub(super) fn install(installed: &Path) -> Result<()> {
|
||||
for (name, files) in PLUGINS {
|
||||
let directory = installed.join(name);
|
||||
if disk_version(&directory) == Some(embedded_version(files)?) {
|
||||
continue;
|
||||
}
|
||||
write_plugin(&directory, files)?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn write_plugin(directory: &std::path::Path, files: &[(&str, &str)]) -> Result<()> {
|
||||
fn embedded_version(files: &[(&str, &str)]) -> Result<String> {
|
||||
let manifest = files
|
||||
.iter()
|
||||
.find(|(name, _)| *name == "plugin.json")
|
||||
.map(|(_, content)| *content)
|
||||
.expect("built-in plugin bundles plugin.json");
|
||||
let value: serde_json::Value = serde_json::from_str(manifest)?;
|
||||
value
|
||||
.get("version")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.map(str::to_owned)
|
||||
.ok_or_else(|| crate::Error::Config("built-in plugin manifest requires version".into()))
|
||||
}
|
||||
|
||||
fn disk_version(directory: &Path) -> Option<String> {
|
||||
let manifest = std::fs::read_to_string(directory.join("plugin.json")).ok()?;
|
||||
let value: serde_json::Value = serde_json::from_str(&manifest).ok()?;
|
||||
Some(value.get("version")?.as_str()?.to_owned())
|
||||
}
|
||||
|
||||
fn write_plugin(directory: &Path, files: &[(&str, &str)]) -> Result<()> {
|
||||
for (relative, content) in files {
|
||||
let path = directory.join(relative);
|
||||
let parent = path.parent().expect("plugin file path has a parent");
|
||||
@@ -81,5 +109,75 @@ fn write_plugin(directory: &std::path::Path, files: &[(&str, &str)]) -> Result<(
|
||||
std::fs::set_permissions(&path, std::fs::Permissions::from_mode(0o600))?;
|
||||
}
|
||||
}
|
||||
prune_unknown_files(directory, directory, files)?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 删除插件目录中不在嵌入清单里的文件与空目录(旧版本残留)。
|
||||
fn prune_unknown_files(root: &Path, directory: &Path, files: &[(&str, &str)]) -> Result<()> {
|
||||
for entry in std::fs::read_dir(directory)? {
|
||||
let entry = entry?;
|
||||
let path = entry.path();
|
||||
if entry.file_type()?.is_dir() {
|
||||
prune_unknown_files(root, &path, files)?;
|
||||
if std::fs::read_dir(&path)?.next().is_none() {
|
||||
std::fs::remove_dir(&path)?;
|
||||
}
|
||||
continue;
|
||||
}
|
||||
let known = files
|
||||
.iter()
|
||||
.any(|(relative, _)| root.join(relative) == path);
|
||||
if !known {
|
||||
std::fs::remove_file(&path)?;
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
fn embedded_main() -> &'static str {
|
||||
CODEX_AUTH
|
||||
.iter()
|
||||
.find(|(name, _)| *name == "main.ts")
|
||||
.unwrap()
|
||||
.1
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn install_is_version_gated_and_syncs_on_version_change() {
|
||||
let root = tempfile::tempdir().unwrap();
|
||||
let plugin = root.path().join("codex-auth");
|
||||
|
||||
install(root.path()).unwrap();
|
||||
assert_eq!(
|
||||
std::fs::read_to_string(plugin.join("main.ts")).unwrap(),
|
||||
embedded_main()
|
||||
);
|
||||
|
||||
// 版本一致:本地改动与额外文件保持原样,不发生任何写盘。
|
||||
std::fs::write(plugin.join("main.ts"), "edited").unwrap();
|
||||
std::fs::write(plugin.join("stale.ts"), "extra").unwrap();
|
||||
install(root.path()).unwrap();
|
||||
assert_eq!(
|
||||
std::fs::read_to_string(plugin.join("main.ts")).unwrap(),
|
||||
"edited"
|
||||
);
|
||||
assert!(plugin.join("stale.ts").exists());
|
||||
|
||||
// 版本变化:整目录同步回嵌入内容并清理残留。
|
||||
let manifest = std::fs::read_to_string(plugin.join("plugin.json")).unwrap();
|
||||
let mut value: serde_json::Value = serde_json::from_str(&manifest).unwrap();
|
||||
value["version"] = serde_json::Value::String("0.0.1".into());
|
||||
std::fs::write(plugin.join("plugin.json"), value.to_string()).unwrap();
|
||||
install(root.path()).unwrap();
|
||||
assert_eq!(
|
||||
std::fs::read_to_string(plugin.join("main.ts")).unwrap(),
|
||||
embedded_main()
|
||||
);
|
||||
assert!(!plugin.join("stale.ts").exists());
|
||||
}
|
||||
}
|
||||
|
||||
@@ -21,6 +21,7 @@ const MAX_ICON_BYTES: u64 = 1024 * 1024;
|
||||
pub struct PluginCatalog {
|
||||
roots: Vec<PathBuf>,
|
||||
definition_loader: PluginDefinitionLoader,
|
||||
app_version: String,
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
@@ -33,7 +34,7 @@ pub(crate) struct PluginEntry {
|
||||
}
|
||||
|
||||
impl PluginCatalog {
|
||||
pub fn managed() -> Result<Self> {
|
||||
pub fn managed(app_version: String) -> Result<Self> {
|
||||
let installed = config::managed_data_dir()?.join("plugins/installed");
|
||||
fs::create_dir_all(&installed)?;
|
||||
#[cfg(unix)]
|
||||
@@ -41,15 +42,21 @@ impl PluginCatalog {
|
||||
use std::os::unix::fs::PermissionsExt;
|
||||
fs::set_permissions(&installed, fs::Permissions::from_mode(0o700))?;
|
||||
}
|
||||
// 扫描顺序即优先级:用户安装目录 > 源码内置目录(仅 debug,便于热改)
|
||||
// > 随二进制打包后落盘的内置目录;同 ID 时靠前的覆盖靠后的。
|
||||
let mut roots = vec![installed];
|
||||
// 内置插件按版本预装进 installed;版本一致时不写盘。
|
||||
super::builtin::install(&installed)?;
|
||||
// 扫描顺序即优先级:debug 下源码目录优先,保证内置插件热改生效;
|
||||
// 发布构建只有 installed 一个根。
|
||||
#[cfg(debug_assertions)]
|
||||
roots.push(PathBuf::from(env!("CARGO_MANIFEST_DIR")).join("plugins/build-in"));
|
||||
roots.push(super::builtin::materialize()?);
|
||||
let roots = vec![
|
||||
PathBuf::from(env!("CARGO_MANIFEST_DIR")).join("plugins/build-in"),
|
||||
installed,
|
||||
];
|
||||
#[cfg(not(debug_assertions))]
|
||||
let roots = vec![installed];
|
||||
Ok(Self {
|
||||
roots,
|
||||
definition_loader: PluginDefinitionLoader::managed()?,
|
||||
app_version,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -69,7 +76,14 @@ impl PluginCatalog {
|
||||
};
|
||||
directories.sort();
|
||||
for directory in directories {
|
||||
match load_plugin(&directory, &self.definition_loader, executable).await {
|
||||
match load_plugin(
|
||||
&directory,
|
||||
&self.definition_loader,
|
||||
executable,
|
||||
&self.app_version,
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(entry) => {
|
||||
if plugins.contains_key(&entry.manifest.id) {
|
||||
tracing::warn!(plugin = %entry.manifest.id, path = %directory.display(), "ignoring duplicate plugin");
|
||||
@@ -98,6 +112,7 @@ impl PluginCatalog {
|
||||
let manifest: PluginManifest =
|
||||
serde_json::from_slice(&fs::read(directory.join(MANIFEST_FILE_NAME))?)?;
|
||||
manifest.validate(&directory)?;
|
||||
require_app_version(&manifest, &self.app_version)?;
|
||||
let icon = icon_data_url(&directory, &manifest.icon)?;
|
||||
Ok((manifest, icon))
|
||||
})();
|
||||
@@ -126,14 +141,30 @@ fn child_directories(root: &Path) -> Result<Vec<PathBuf>> {
|
||||
Ok(directories)
|
||||
}
|
||||
|
||||
/// 应用过旧时拒绝加载,让插件的 minAppVersion 声明生效。
|
||||
fn require_app_version(manifest: &PluginManifest, app_version: &str) -> Result<()> {
|
||||
let Some(minimum) = &manifest.min_app_version else {
|
||||
return Ok(());
|
||||
};
|
||||
if super::manifest::version_at_least(app_version, minimum) {
|
||||
return Ok(());
|
||||
}
|
||||
Err(Error::Config(format!(
|
||||
"plugin '{}' requires app version {minimum} or newer (current {app_version})",
|
||||
manifest.id
|
||||
)))
|
||||
}
|
||||
|
||||
async fn load_plugin(
|
||||
directory: &Path,
|
||||
loader: &PluginDefinitionLoader,
|
||||
executable: &Path,
|
||||
app_version: &str,
|
||||
) -> Result<PluginEntry> {
|
||||
let manifest: PluginManifest =
|
||||
serde_json::from_slice(&fs::read(directory.join(MANIFEST_FILE_NAME))?)?;
|
||||
manifest.validate(directory)?;
|
||||
require_app_version(&manifest, app_version)?;
|
||||
let icon = icon_data_url(directory, &manifest.icon)?;
|
||||
let entry = directory.join(&manifest.entry).canonicalize()?;
|
||||
let definition = loader.load(executable, directory, &entry).await?;
|
||||
@@ -279,6 +310,7 @@ mod tests {
|
||||
let catalog = PluginCatalog {
|
||||
roots: vec![root],
|
||||
definition_loader: PluginDefinitionLoader::for_test(sdk.path()).unwrap(),
|
||||
app_version: env!("CARGO_PKG_VERSION").into(),
|
||||
};
|
||||
assert!(!catalog.manifests().is_empty());
|
||||
}
|
||||
|
||||
@@ -71,6 +71,7 @@ pub const OAUTH2_ADD_METHOD: &str = "oauth2.0";
|
||||
pub struct PluginDescriptor {
|
||||
pub id: String,
|
||||
pub name: String,
|
||||
pub version: String,
|
||||
pub author: Option<String>,
|
||||
pub icon: String,
|
||||
pub providers: Vec<PluginProviderDescriptor>,
|
||||
|
||||
@@ -14,8 +14,13 @@ pub struct PluginManifest {
|
||||
pub api_version: u32,
|
||||
pub id: String,
|
||||
pub name: String,
|
||||
/// 插件自身版本;内置插件预装时以它为缓存键决定是否重新落盘。
|
||||
pub version: String,
|
||||
#[serde(default)]
|
||||
pub author: Option<String>,
|
||||
/// 插件要求的最低应用版本;应用过旧时插件被忽略。
|
||||
#[serde(default)]
|
||||
pub min_app_version: Option<String>,
|
||||
pub icon: String,
|
||||
pub entry: String,
|
||||
#[serde(default)]
|
||||
@@ -39,6 +44,12 @@ impl PluginManifest {
|
||||
}
|
||||
validate_id(&self.id, "plugin id")?;
|
||||
required(&self.name, "plugin name")?;
|
||||
parse_version(&self.version)
|
||||
.ok_or_else(|| Error::Config(format!("invalid plugin version: {}", self.version)))?;
|
||||
if let Some(minimum) = &self.min_app_version {
|
||||
parse_version(minimum)
|
||||
.ok_or_else(|| Error::Config(format!("invalid plugin minAppVersion: {minimum}")))?;
|
||||
}
|
||||
validate_entry_path(directory, &self.entry)?;
|
||||
validate_asset_path(directory, &self.icon)?;
|
||||
let mut hosts = HashSet::new();
|
||||
@@ -55,6 +66,23 @@ impl PluginManifest {
|
||||
}
|
||||
}
|
||||
|
||||
/// 解析 semver 的核心三段(忽略预发布/构建后缀),格式非法返回 None。
|
||||
pub(super) fn parse_version(value: &str) -> Option<(u64, u64, u64)> {
|
||||
let core = value.split(['-', '+']).next()?;
|
||||
let mut parts = core.split('.');
|
||||
let major = parts.next()?.parse().ok()?;
|
||||
let minor = parts.next()?.parse().ok()?;
|
||||
let patch = parts.next()?.parse().ok()?;
|
||||
parts.next().is_none().then_some((major, minor, patch))
|
||||
}
|
||||
|
||||
pub(super) fn version_at_least(actual: &str, minimum: &str) -> bool {
|
||||
match (parse_version(actual), parse_version(minimum)) {
|
||||
(Some(actual), Some(minimum)) => actual >= minimum,
|
||||
_ => false,
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn validate_id(value: &str, label: &str) -> Result<()> {
|
||||
static ID: std::sync::OnceLock<Regex> = std::sync::OnceLock::new();
|
||||
let expression = ID.get_or_init(|| Regex::new(r"^[a-z0-9]+(?:[._-][a-z0-9]+)*$").unwrap());
|
||||
@@ -172,4 +200,14 @@ mod tests {
|
||||
assert!(validate_network_host("example.com:443").is_err());
|
||||
assert!(validate_network_host("example.com").is_ok());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn compares_semver_cores_and_ignores_prerelease_suffixes() {
|
||||
assert_eq!(parse_version("0.1.5-beta.1"), Some((0, 1, 5)));
|
||||
assert_eq!(parse_version("1.2"), None);
|
||||
assert!(version_at_least("0.1.5-beta.1", "0.1.5"));
|
||||
assert!(version_at_least("0.2.0", "0.1.9"));
|
||||
assert!(!version_at_least("0.1.4", "0.1.5"));
|
||||
assert!(!version_at_least("bogus", "0.1.0"));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -97,13 +97,13 @@ pub struct PluginInvocationPlan {
|
||||
}
|
||||
|
||||
impl PluginRegistry {
|
||||
pub fn managed(store: Store, runtime: PluginRuntime) -> Result<Self> {
|
||||
pub fn managed(store: Store, runtime: PluginRuntime, app_version: String) -> Result<Self> {
|
||||
let data = PluginDataStore::managed()?;
|
||||
Ok(Self {
|
||||
inner: Arc::new(RegistryInner {
|
||||
store,
|
||||
runtime,
|
||||
catalog: PluginCatalog::managed()?,
|
||||
catalog: PluginCatalog::managed(app_version)?,
|
||||
state: PluginStateStore::new(data),
|
||||
entries: RwLock::new(None),
|
||||
workers: Mutex::new(HashMap::new()),
|
||||
@@ -122,6 +122,7 @@ impl PluginRegistry {
|
||||
.map(|(manifest, icon)| PluginDescriptor {
|
||||
id: manifest.id,
|
||||
name: manifest.name,
|
||||
version: manifest.version,
|
||||
author: manifest.author,
|
||||
icon,
|
||||
providers: Vec::new(),
|
||||
@@ -625,6 +626,7 @@ impl PluginRegistry {
|
||||
PluginDescriptor {
|
||||
id: plugin_id.clone(),
|
||||
name: entry.manifest.name.clone(),
|
||||
version: entry.manifest.version.clone(),
|
||||
author: entry.manifest.author.clone(),
|
||||
icon: entry.icon.clone(),
|
||||
providers,
|
||||
|
||||
@@ -17,11 +17,8 @@ use crate::{
|
||||
};
|
||||
|
||||
use super::{
|
||||
<<<<<<< HEAD
|
||||
apply_body_allowlist, apply_openai_prompt_cache_key, merge_extra_params,
|
||||
=======
|
||||
apply_openai_prompt_cache_key, map_sse_error, merge_extra_params, provider_event_error,
|
||||
>>>>>>> main
|
||||
apply_body_allowlist, apply_openai_prompt_cache_key, map_sse_error, merge_extra_params,
|
||||
provider_event_error,
|
||||
recorder::recorded_headers,
|
||||
retry::{send_with_retry, Attempt, RetryPolicy},
|
||||
CallRecorder, FinishReason, ModelEvent, Provider, ProviderStream,
|
||||
|
||||
@@ -14,11 +14,8 @@ use crate::{
|
||||
};
|
||||
|
||||
use super::{
|
||||
<<<<<<< HEAD
|
||||
apply_body_allowlist, apply_openai_prompt_cache_key, merge_extra_params,
|
||||
=======
|
||||
apply_openai_prompt_cache_key, map_sse_error, merge_extra_params, provider_event_error,
|
||||
>>>>>>> main
|
||||
apply_body_allowlist, apply_openai_prompt_cache_key, map_sse_error, merge_extra_params,
|
||||
provider_event_error,
|
||||
recorder::recorded_headers,
|
||||
retry::{send_with_retry, Attempt, RetryPolicy},
|
||||
CallRecorder, FinishReason, ModelEvent, Provider, ProviderStream,
|
||||
|
||||
+77
-135
@@ -14,8 +14,8 @@ use crate::{
|
||||
};
|
||||
|
||||
use super::{
|
||||
normalize::NormalizedProvider, AnthropicProvider, CallRecorder, OpenAiChatProvider,
|
||||
OpenAiResponsesProvider, Provider, ProviderStream,
|
||||
normalize::NormalizedProvider, recorder::CancelOnDrop, AnthropicProvider, CallRecorder,
|
||||
OpenAiChatProvider, OpenAiResponsesProvider, Provider, ProviderStream,
|
||||
};
|
||||
|
||||
const BUILTIN_PROVIDER_RETRIES: u32 = 5;
|
||||
@@ -28,11 +28,12 @@ pub struct ProviderRouter {
|
||||
}
|
||||
|
||||
impl ProviderRouter {
|
||||
<<<<<<< HEAD
|
||||
pub fn new(store: Store, plugins: PluginRegistry, request_timeout: Duration) -> Self {
|
||||
=======
|
||||
pub fn new(store: Store, request_timeout: Duration, stream_idle_timeout: Duration) -> Self {
|
||||
>>>>>>> main
|
||||
pub fn new(
|
||||
store: Store,
|
||||
plugins: PluginRegistry,
|
||||
request_timeout: Duration,
|
||||
stream_idle_timeout: Duration,
|
||||
) -> Self {
|
||||
Self {
|
||||
store,
|
||||
plugins,
|
||||
@@ -54,87 +55,57 @@ impl Provider for ProviderRouter {
|
||||
let stream_idle_timeout = self.stream_idle_timeout;
|
||||
Box::pin(try_stream! {
|
||||
let selected = invocation.request.model.model_id.clone();
|
||||
<<<<<<< HEAD
|
||||
if selected.starts_with(ADAPTER_ID_PREFIX) {
|
||||
// 插件模型与内置模型走完全相同的流程:Recorder、统一事件、
|
||||
// 规范化包装。资源选择与将来的负载均衡都在插件 Provider 内部。
|
||||
let plan = plugins.plan_model(&selected).await?;
|
||||
let recorder = start_recorder(&store, &invocation, &selected, &plan.model.display_name, ProviderType::Plugin, &plan.request_url, &plan.model.model_id).await?;
|
||||
let _cancel_on_drop = recorder.cancel_on_drop();
|
||||
recorder.request(serde_json::json!({}), &crate::plugin::plugin_llm_request(&invocation)?).await?;
|
||||
let mut routed = invocation.clone();
|
||||
routed.request.model.display_name = Some(plan.model.display_name.clone());
|
||||
if let Some(tokens) = plan.model.context_window_tokens {
|
||||
routed.request.model.context_window_tokens.get_or_insert(tokens);
|
||||
}
|
||||
if let Some(tokens) = plan.model.max_output_tokens {
|
||||
routed.request.model.max_output_tokens.get_or_insert(tokens);
|
||||
}
|
||||
let provider: Arc<dyn Provider> = Arc::new(NormalizedProvider::new(Arc::new(PluginModelProvider {
|
||||
registry: plugins.clone(),
|
||||
})));
|
||||
let mut stream = provider.stream(routed, cancellation.clone());
|
||||
while let Some(item) = stream.next().await {
|
||||
match item {
|
||||
Ok(event) => { recorder.event(&event).await?; yield event; }
|
||||
Err(error) => { recorder.failed(&error).await?; Err(error)?; }
|
||||
=======
|
||||
let model = store
|
||||
.model(&selected)
|
||||
.await?
|
||||
.ok_or_else(|| Error::Provider(format!("unknown model: {selected}")))?;
|
||||
let provider_type = model.provider_type();
|
||||
let request_url = model.request_url()?;
|
||||
model.configure(&mut invocation.request.model);
|
||||
invocation.request.model.extra_params = model.extra_params().clone();
|
||||
invocation.request.model.model_id = model.model_id.clone();
|
||||
let recorder = CallRecorder::start(store.clone(), NewLlmCall {
|
||||
call_id: invocation.call_id.clone(),
|
||||
run_id: invocation.run_id.clone(),
|
||||
conversation_id: invocation.conversation_id.clone(),
|
||||
provider_call_index: invocation.provider_call_index.min(i64::MAX as u64) as i64,
|
||||
model_hash: model.model_hash.clone(),
|
||||
provider_type,
|
||||
provider_url: model.base_url.clone(),
|
||||
request_type: provider_type,
|
||||
request_url: request_url.clone(),
|
||||
model_id: model.model_id.clone(),
|
||||
display_name: model.display_name.clone(),
|
||||
reasoning_effort: invocation.request.model.reasoning.effort.clone(),
|
||||
fast: invocation.request.model.latency == ModelLatency::Fast,
|
||||
message_count: invocation.request.history.len(),
|
||||
tool_count: invocation.request.prompt.tools.len(),
|
||||
detailed: false,
|
||||
}).await?;
|
||||
let _cancel_on_drop = recorder.cancel_on_drop();
|
||||
let config = ProviderConfig {
|
||||
kind: match provider_type {
|
||||
ProviderType::OpenAiChat => ProviderKind::OpenAiChat,
|
||||
ProviderType::OpenAiResponses => ProviderKind::OpenAiResponses,
|
||||
ProviderType::Anthropic => ProviderKind::Anthropic,
|
||||
},
|
||||
request_url,
|
||||
api_key: model.api_key.clone(),
|
||||
custom_headers: if model.custom_headers_enabled {
|
||||
custom_headers(&model.custom_headers)?
|
||||
// 两条分支只负责装配 Recorder 与 Provider 流;
|
||||
// 事件消费(空闲超时看门狗、记录、错误规范化)对两者完全一致。
|
||||
let (recorder, _cancel_on_drop, mut stream): (CallRecorder, CancelOnDrop, ProviderStream) =
|
||||
if selected.starts_with(ADAPTER_ID_PREFIX) {
|
||||
// 插件模型与内置模型走完全相同的流程:资源选择与将来的
|
||||
// 负载均衡都在插件 Provider 内部。
|
||||
let plan = plugins.plan_model(&selected).await?;
|
||||
let recorder = start_recorder(&store, &invocation, &selected, &plan.model.display_name, ProviderType::Plugin, &plan.request_url, &plan.model.model_id).await?;
|
||||
let guard = recorder.cancel_on_drop();
|
||||
recorder.request(serde_json::json!({}), &crate::plugin::plugin_llm_request(&invocation)?).await?;
|
||||
let mut routed = invocation.clone();
|
||||
routed.request.model.display_name = Some(plan.model.display_name.clone());
|
||||
if let Some(tokens) = plan.model.context_window_tokens {
|
||||
routed.request.model.context_window_tokens.get_or_insert(tokens);
|
||||
}
|
||||
if let Some(tokens) = plan.model.max_output_tokens {
|
||||
routed.request.model.max_output_tokens.get_or_insert(tokens);
|
||||
}
|
||||
let provider: Arc<dyn Provider> = Arc::new(NormalizedProvider::new(Arc::new(PluginModelProvider {
|
||||
registry: plugins.clone(),
|
||||
})));
|
||||
(recorder, guard, provider.stream(routed, cancellation.clone()))
|
||||
} else {
|
||||
reqwest::header::HeaderMap::new()
|
||||
},
|
||||
max_output_tokens: model.max_output_tokens(),
|
||||
request_timeout,
|
||||
};
|
||||
let client = crate::network::client_builder(&store)
|
||||
.await?
|
||||
.timeout(config.request_timeout)
|
||||
.build()?;
|
||||
let provider = build_observed(&config, recorder.clone(), client)?;
|
||||
let stream_cancellation = cancellation.clone();
|
||||
let mut stream = provider.stream(invocation, cancellation);
|
||||
let mut routed = invocation.clone();
|
||||
let model = store.model(&selected).await?.ok_or_else(|| Error::Provider(format!("unknown model: {selected}")))?;
|
||||
let provider_type = model.provider_type();
|
||||
let request_url = model.request_url()?;
|
||||
model.configure(&mut routed.request.model);
|
||||
routed.request.model.extra_params = model.extra_params().clone();
|
||||
routed.request.model.model_id = model.model_id.clone();
|
||||
let recorder = start_recorder(&store, &invocation, &model.model_hash, &model.display_name, provider_type, &request_url, &model.model_id).await?;
|
||||
let guard = recorder.cancel_on_drop();
|
||||
let config = ProviderConfig {
|
||||
kind: provider_kind(provider_type),
|
||||
request_url,
|
||||
api_key: model.api_key.clone(),
|
||||
custom_headers: if model.custom_headers_enabled { custom_headers(&model.custom_headers)? } else { reqwest::header::HeaderMap::new() },
|
||||
max_output_tokens: model.max_output_tokens(),
|
||||
request_timeout,
|
||||
retry_count: BUILTIN_PROVIDER_RETRIES,
|
||||
allowed_body_fields: None,
|
||||
};
|
||||
let client = crate::network::client_builder(&store).await?.timeout(request_timeout).build()?;
|
||||
let provider = build_observed(&config, recorder.clone(), client)?;
|
||||
(recorder, guard, provider.stream(routed, cancellation.clone()))
|
||||
};
|
||||
|
||||
let stream_started = std::time::Instant::now();
|
||||
tracing::debug!(
|
||||
model = %selected,
|
||||
provider_type = ?provider_type,
|
||||
request_timeout_ms = config.request_timeout.as_millis() as u64,
|
||||
request_timeout_ms = request_timeout.as_millis() as u64,
|
||||
stream_idle_timeout_ms = stream_idle_timeout.as_millis() as u64,
|
||||
"provider stream created"
|
||||
);
|
||||
@@ -163,26 +134,11 @@ impl Provider for ProviderRouter {
|
||||
event_count += 1;
|
||||
match event {
|
||||
Ok(event) => {
|
||||
let event_name = match &event {
|
||||
super::ModelEvent::Start { .. } => "Start",
|
||||
super::ModelEvent::TextStart => "TextStart",
|
||||
super::ModelEvent::TextDelta(_) => "TextDelta",
|
||||
super::ModelEvent::TextEnd => "TextEnd",
|
||||
super::ModelEvent::ThinkingStart => "ThinkingStart",
|
||||
super::ModelEvent::ThinkingDelta(_) => "ThinkingDelta",
|
||||
super::ModelEvent::ThinkingEnd => "ThinkingEnd",
|
||||
super::ModelEvent::ToolCallStart { .. } => "ToolCallStart",
|
||||
super::ModelEvent::ToolCallArgumentsDelta { .. } => "ToolCallArgsDelta",
|
||||
super::ModelEvent::ToolCallEnd { .. } => "ToolCallEnd",
|
||||
super::ModelEvent::ProviderReplayState(_) => "ReplayState",
|
||||
super::ModelEvent::Usage(_) => "Usage",
|
||||
super::ModelEvent::Done(_) => "Done",
|
||||
};
|
||||
if gap_ms > 5000 {
|
||||
tracing::debug!(
|
||||
gap_ms,
|
||||
elapsed_ms,
|
||||
event = event_name,
|
||||
event = event_name(&event),
|
||||
event_count,
|
||||
"slow gap detected between provider events"
|
||||
);
|
||||
@@ -202,46 +158,32 @@ impl Provider for ProviderRouter {
|
||||
);
|
||||
recorder.failed(&error).await?;
|
||||
Err(error)?;
|
||||
>>>>>>> main
|
||||
}
|
||||
}
|
||||
finish_stream(&recorder, &cancellation).await?;
|
||||
} else {
|
||||
let mut routed = invocation.clone();
|
||||
let model = store.model(&selected).await?.ok_or_else(|| Error::Provider(format!("unknown model: {selected}")))?;
|
||||
let provider_type = model.provider_type();
|
||||
let request_url = model.request_url()?;
|
||||
model.configure(&mut routed.request.model);
|
||||
routed.request.model.extra_params = model.extra_params().clone();
|
||||
routed.request.model.model_id = model.model_id.clone();
|
||||
let recorder = start_recorder(&store, &invocation, &model.model_hash, &model.display_name, provider_type, &request_url, &model.model_id).await?;
|
||||
let _cancel_on_drop = recorder.cancel_on_drop();
|
||||
let config = ProviderConfig {
|
||||
kind: provider_kind(provider_type),
|
||||
request_url,
|
||||
api_key: model.api_key.clone(),
|
||||
custom_headers: if model.custom_headers_enabled { custom_headers(&model.custom_headers)? } else { reqwest::header::HeaderMap::new() },
|
||||
max_output_tokens: model.max_output_tokens(),
|
||||
request_timeout,
|
||||
retry_count: BUILTIN_PROVIDER_RETRIES,
|
||||
allowed_body_fields: None,
|
||||
};
|
||||
let client = crate::network::client_builder(&store).await?.timeout(request_timeout).build()?;
|
||||
let provider = build_observed(&config, recorder.clone(), client)?;
|
||||
let mut stream = provider.stream(routed, cancellation.clone());
|
||||
while let Some(item) = stream.next().await {
|
||||
match item {
|
||||
Ok(event) => { recorder.event(&event).await?; yield event; }
|
||||
Err(error) => { recorder.failed(&error).await?; Err(error)?; }
|
||||
}
|
||||
}
|
||||
finish_stream(&recorder, &cancellation).await?;
|
||||
}
|
||||
finish_stream(&recorder, &cancellation).await?;
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
<<<<<<< HEAD
|
||||
fn event_name(event: &super::ModelEvent) -> &'static str {
|
||||
match event {
|
||||
super::ModelEvent::Start { .. } => "Start",
|
||||
super::ModelEvent::TextStart => "TextStart",
|
||||
super::ModelEvent::TextDelta(_) => "TextDelta",
|
||||
super::ModelEvent::TextEnd => "TextEnd",
|
||||
super::ModelEvent::ThinkingStart => "ThinkingStart",
|
||||
super::ModelEvent::ThinkingDelta(_) => "ThinkingDelta",
|
||||
super::ModelEvent::ThinkingEnd => "ThinkingEnd",
|
||||
super::ModelEvent::ToolCallStart { .. } => "ToolCallStart",
|
||||
super::ModelEvent::ToolCallArgumentsDelta { .. } => "ToolCallArgsDelta",
|
||||
super::ModelEvent::ToolCallEnd { .. } => "ToolCallEnd",
|
||||
super::ModelEvent::ProviderReplayState(_) => "ReplayState",
|
||||
super::ModelEvent::Usage(_) => "Usage",
|
||||
super::ModelEvent::Done(_) => "Done",
|
||||
}
|
||||
}
|
||||
|
||||
async fn start_recorder(
|
||||
store: &Store,
|
||||
invocation: &ModelInvocation,
|
||||
@@ -311,7 +253,8 @@ fn provider_kind(provider_type: ProviderType) -> ProviderKind {
|
||||
// 内置模型的 provider_type 只来自 ModelType,不可能是插件。
|
||||
ProviderType::Plugin => unreachable!("plugin models never use built-in provider configs"),
|
||||
}
|
||||
=======
|
||||
}
|
||||
|
||||
async fn next_provider_event(
|
||||
stream: &mut ProviderStream,
|
||||
idle_timeout: Duration,
|
||||
@@ -352,7 +295,6 @@ fn root_error_message(error: &(dyn std::error::Error + 'static)) -> String {
|
||||
current = source;
|
||||
}
|
||||
current.to_string()
|
||||
>>>>>>> main
|
||||
}
|
||||
|
||||
fn custom_headers(value: &serde_json::Value) -> Result<reqwest::header::HeaderMap> {
|
||||
|
||||
Reference in New Issue
Block a user