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:
leookun
2026-08-30 20:45:27 +08:00
parent 05181f9e8a
commit a5bbe67845
27 changed files with 651 additions and 270 deletions
+5 -1
View File
@@ -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(),
+4
View File
@@ -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
View File
@@ -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());
}
}
+39 -7
View File
@@ -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());
}
+1
View File
@@ -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>,
+38
View File
@@ -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"));
}
}
+4 -2
View File
@@ -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,
+2 -5
View File
@@ -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,
+2 -5
View File
@@ -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
View File
@@ -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> {