feat: plugin system

This commit is contained in:
leookun
2026-08-30 20:00:53 +08:00
parent 1609b57433
commit e6130e01a7
78 changed files with 9010 additions and 643 deletions
+4 -1
View File
@@ -133,7 +133,10 @@ async fn bidi_handler(
let conversation_id = decoded.conversation_id().map(str::to_owned);
let trace_metadata = decoded.trace_metadata();
let local = if let Some(model_id) = decoded.model_id() {
if registry.store().model(model_id).await?.is_some() {
// 插件模型 ID 只在本地有意义,永远不转发到 Cursor 官方上游。
if model_id.starts_with(crate::plugin::ADAPTER_ID_PREFIX)
|| registry.store().model(model_id).await?.is_some()
{
tracing::info!(
request_id = decoded.request_id,
model_id,
+8 -2
View File
@@ -13,6 +13,7 @@ use crate::{
transport::TransportRegistry,
},
local_app::CursorHarness,
plugin::{PluginRegistry, PluginRuntime},
provider::ProviderRouter,
search::WebCache,
store::Store,
@@ -37,17 +38,22 @@ 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 provider = std::sync::Arc::new(ProviderRouter::new(
store.clone(),
plugins.clone(),
config.provider_request_timeout,
));
let registry = TransportRegistry::with_web_cache(
let registry = TransportRegistry::with_plugins(
store.clone(),
provider.clone(),
compiler,
WebCache::managed()?,
plugins.clone(),
);
let control = control::ControlService::new(store.clone(), provider)?;
let control =
control::ControlService::new(store.clone(), provider, plugin_runtime, plugins)?;
let harness = control.cursor_harness().clone();
let mut router = api::router(registry.clone())?;
router = match &config.console {
+2
View File
@@ -45,6 +45,8 @@ pub struct ProviderConfig {
pub custom_headers: reqwest::header::HeaderMap,
pub max_output_tokens: Option<u64>,
pub request_timeout: Duration,
pub retry_count: u32,
pub allowed_body_fields: Option<std::collections::HashSet<String>>,
}
#[derive(Clone)]
+33
View File
@@ -137,12 +137,45 @@ pub fn api_router(service: ControlService) -> Router {
)
.route("/__byok-api__/api/llm-calls", get(calls::list))
.route("/__byok-api__/api/llm-calls/{call_id}", get(calls::detail))
.route("/__byok-api__/api/plugins", get(plugins::list))
.route(
"/__byok-api__/api/plugins/runtime",
get(plugins::runtime_status)
.post(plugins::initialize_runtime)
.delete(plugins::cancel_runtime_initialization),
)
.route(
"/__byok-api__/api/plugins/oauth/{session_id}/poll",
post(plugins::oauth_poll),
)
.route(
"/__byok-api__/api/plugins/{plugin_id}",
axum::routing::delete(plugins::remove),
)
.route(
"/__byok-api__/api/plugins/{plugin_id}/resources/{resource_type}/add/{method_id}/begin",
post(plugins::oauth_begin),
)
.route(
"/__byok-api__/api/plugins/{plugin_id}/resources/{resource_type}/import",
post(plugins::import),
)
.route(
"/__byok-api__/api/plugins/{plugin_id}/resources/{resource_type}/export",
get(plugins::export_resources),
)
.route(
"/__byok-api__/api/plugins/{plugin_id}/resources/{resource_type}/{resource_id}",
axum::routing::delete(plugins::delete_resource),
)
.route(
"/__byok-api__/api/plugins/{plugin_id}/resources/{resource_type}/{resource_id}/refresh",
post(plugins::refresh_resource),
)
.route(
"/__byok-api__/api/plugins/{plugin_id}/providers/{provider_id}/models/sync",
post(plugins::sync_models),
)
.route(
"/__byok-api__/api/settings/observability",
get(settings::get).put(settings::update),
+103 -3
View File
@@ -1,10 +1,110 @@
//! Exposes plugin runtime initialization and status endpoints.
use axum::{extract::State, Json};
//! Exposes plugin discovery, resource lifecycle, model sync, and runtime endpoints.
use axum::{
extract::{Path, State},
http::StatusCode,
Json,
};
use crate::{plugin::PluginRuntimeStatus, Result};
use crate::{
plugin::{
ImportResponse, OAuthBeginResponse, OAuthPollResponse, PluginDescriptor,
PluginRuntimeStatus,
},
Result,
};
use super::ControlService;
pub async fn list(State(service): State<ControlService>) -> Result<Json<Vec<PluginDescriptor>>> {
Ok(Json(service.plugins().await))
}
pub async fn remove(
State(service): State<ControlService>,
Path(plugin_id): Path<String>,
) -> Result<StatusCode> {
service.remove_plugin_configuration(&plugin_id).await?;
Ok(StatusCode::NO_CONTENT)
}
pub async fn oauth_begin(
State(service): State<ControlService>,
Path((plugin_id, resource_type, method_id)): Path<(String, String, String)>,
) -> Result<Json<OAuthBeginResponse>> {
Ok(Json(
service
.plugin_oauth_begin(&plugin_id, &resource_type, &method_id)
.await?,
))
}
pub async fn oauth_poll(
State(service): State<ControlService>,
Path(session_id): Path<String>,
) -> Result<Json<OAuthPollResponse>> {
Ok(Json(service.plugin_oauth_poll(&session_id).await?))
}
pub async fn import(
State(service): State<ControlService>,
Path((plugin_id, resource_type)): Path<(String, String)>,
Json(files): Json<serde_json::Value>,
) -> Result<Json<ImportResponse>> {
Ok(Json(
service
.plugin_import(&plugin_id, &resource_type, files)
.await?,
))
}
/// 以附件形式返回账号资源导出文件,便于浏览器直接下载。
pub async fn export_resources(
State(service): State<ControlService>,
Path((plugin_id, resource_type)): Path<(String, String)>,
) -> Result<axum::response::Response> {
let value = service
.plugin_export_resources(&plugin_id, &resource_type)
.await?;
let body = serde_json::to_vec_pretty(&value)?;
let response = axum::response::Response::builder()
.header(axum::http::header::CONTENT_TYPE, "application/json")
.header(
axum::http::header::CONTENT_DISPOSITION,
format!("attachment; filename=\"{plugin_id}-{resource_type}.json\""),
)
.body(axum::body::Body::from(body))
.expect("static export response");
Ok(response)
}
pub async fn refresh_resource(
State(service): State<ControlService>,
Path((plugin_id, resource_type, resource_id)): Path<(String, String, String)>,
) -> Result<StatusCode> {
service
.plugin_refresh_resource(&plugin_id, &resource_type, &resource_id)
.await?;
Ok(StatusCode::NO_CONTENT)
}
pub async fn delete_resource(
State(service): State<ControlService>,
Path((plugin_id, resource_type, resource_id)): Path<(String, String, String)>,
) -> Result<StatusCode> {
service
.plugin_delete_resource(&plugin_id, &resource_type, &resource_id)
.await?;
Ok(StatusCode::NO_CONTENT)
}
pub async fn sync_models(
State(service): State<ControlService>,
Path((plugin_id, provider_id)): Path<(String, String)>,
) -> Result<Json<serde_json::Value>> {
let count = service.plugin_sync_models(&plugin_id, &provider_id).await?;
Ok(Json(serde_json::json!({ "models": count })))
}
pub async fn runtime_status(
State(service): State<ControlService>,
) -> Result<Json<PluginRuntimeStatus>> {
+102 -10
View File
@@ -25,7 +25,7 @@ use crate::{
ModelRequest, ModelSpec, ModelType, Overview, ProjectedContent, ProjectedMessage,
PromptSpec, ProviderType, Role,
},
plugin::{PluginRuntime, PluginRuntimeStatus},
plugin::{PluginDescriptor, PluginRegistry, PluginRuntime, PluginRuntimeStatus},
provider::{is_valid_response_event, ModelEvent, Provider},
store::{
DesktopSettings, PortSettings, ProxySettings, ProxySettingsInput, StatisticsStorage, Store,
@@ -40,6 +40,7 @@ pub struct ControlService {
cursor_harness: CursorHarness,
provider: Arc<dyn Provider>,
plugin_runtime: PluginRuntime,
plugins: PluginRegistry,
model_tests: Arc<Mutex<BTreeMap<String, CancellationToken>>>,
}
@@ -145,12 +146,18 @@ pub struct ObservabilitySettings {
}
impl ControlService {
pub fn new(store: Store, provider: Arc<dyn Provider>) -> Result<Self> {
pub fn new(
store: Store,
provider: Arc<dyn Provider>,
plugin_runtime: PluginRuntime,
plugins: PluginRegistry,
) -> Result<Self> {
Ok(Self {
cursor_harness: CursorHarness::new(store.clone())?,
store,
provider,
plugin_runtime: PluginRuntime::managed()?,
plugin_runtime,
plugins,
model_tests: Arc::new(Mutex::new(BTreeMap::new())),
})
}
@@ -159,6 +166,79 @@ impl ControlService {
&self.cursor_harness
}
pub async fn plugins(&self) -> Vec<PluginDescriptor> {
self.plugins.plugins().await
}
pub async fn plugin_oauth_begin(
&self,
plugin_id: &str,
resource_type: &str,
method_id: &str,
) -> Result<crate::plugin::OAuthBeginResponse> {
self.plugins
.oauth_begin(plugin_id, resource_type, method_id)
.await
}
pub async fn plugin_oauth_poll(
&self,
session_id: &str,
) -> Result<crate::plugin::OAuthPollResponse> {
self.plugins.oauth_poll(session_id).await
}
pub async fn plugin_import(
&self,
plugin_id: &str,
resource_type: &str,
files: serde_json::Value,
) -> Result<crate::plugin::ImportResponse> {
self.plugins
.import_resources(plugin_id, resource_type, files)
.await
}
pub async fn plugin_export_resources(
&self,
plugin_id: &str,
resource_type: &str,
) -> Result<serde_json::Value> {
self.plugins
.export_resources(plugin_id, resource_type)
.await
}
pub async fn plugin_refresh_resource(
&self,
plugin_id: &str,
resource_type: &str,
resource_id: &str,
) -> Result<()> {
self.plugins
.refresh_resource(plugin_id, resource_type, resource_id)
.await
}
pub async fn plugin_delete_resource(
&self,
plugin_id: &str,
resource_type: &str,
resource_id: &str,
) -> Result<()> {
self.plugins
.delete_resource(plugin_id, resource_type, resource_id)
.await
}
pub async fn plugin_sync_models(&self, plugin_id: &str, provider_id: &str) -> Result<usize> {
self.plugins.sync_models(plugin_id, provider_id).await
}
pub async fn remove_plugin_configuration(&self, plugin_id: &str) -> Result<()> {
self.plugins.remove(plugin_id).await
}
pub fn plugin_runtime_status(&self) -> PluginRuntimeStatus {
self.plugin_runtime.status()
}
@@ -308,14 +388,21 @@ impl ControlService {
const TEST_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(45);
const TEST_PROMPT: &str = "Output the numbers 1 through 120 separated by a single space. No commas, no newlines, no explanation.";
let configured = self
.store
.model(model_hash)
.await?
.ok_or_else(|| Error::RunNotFound(format!("model {model_hash}")))?;
let mut model = ModelSpec::new(model_hash);
configured.configure(&mut model);
model.max_output_tokens = Some(configured.max_output_tokens().unwrap_or(65_536));
if model_hash.starts_with(crate::plugin::ADAPTER_ID_PREFIX) {
let descriptor = self.plugins.model_descriptor(model_hash).await?;
model.display_name = Some(descriptor.display_name);
model.context_window_tokens = descriptor.context_window_tokens;
model.max_output_tokens = Some(descriptor.max_output_tokens.unwrap_or(65_536));
} else {
let configured = self
.store
.model(model_hash)
.await?
.ok_or_else(|| Error::RunNotFound(format!("model {model_hash}")))?;
configured.configure(&mut model);
model.max_output_tokens = Some(configured.max_output_tokens().unwrap_or(65_536));
}
let call_id = format!("model-test-{}", uuid::Uuid::new_v4());
let invocation = ModelInvocation {
call_id: call_id.clone(),
@@ -714,6 +801,11 @@ async fn discover_models_from_endpoint(
ProviderType::Anthropic => {
anthropic_models(client, base_url, api_key, custom_headers).await?
}
ProviderType::Plugin => {
return Err(Error::Config(
"plugin providers discover models through their plugin".into(),
))
}
};
models.sort();
models.dedup();
+210 -89
View File
@@ -11,6 +11,7 @@ use crate::{
api::cursor::proxy::{self, CursorProxy},
cursor::{protocol::proto::agent::v1 as agent, transport::TransportRegistry},
model::{format_token_count, parse_token_count, ModelConfig, ModelType},
plugin::PluginModelDescriptor,
Error, Result,
};
@@ -186,9 +187,10 @@ struct UsableModelsAddition {
models: Vec<agent::ModelDetails>,
}
const CONTEXTS: [(&str, &str); 4] = [
const CONTEXTS: [(&str, &str); 5] = [
("200k", "200K"),
("356k", "356K"),
("500k", "500K"),
("800k", "800K"),
("1m", "1M"),
];
@@ -201,12 +203,12 @@ const EFFORTS: [(&str, &str); 5] = [
];
const DEFAULT_CONTEXT: &str = "200k";
fn context_options(model: &ModelConfig) -> Vec<(String, String)> {
fn context_options(context_window_tokens: Option<u64>) -> Vec<(String, String)> {
let mut contexts = CONTEXTS
.into_iter()
.map(|(value, display_name)| (value.to_owned(), display_name.to_owned()))
.collect::<Vec<_>>();
if let Some(tokens) = model.context_window_tokens {
if let Some(tokens) = context_window_tokens {
let value = tokens.to_string();
let duplicate = contexts
.iter()
@@ -224,15 +226,22 @@ pub async fn available_models(
request: Request<Body>,
) -> Result<Response<Body>> {
let models = registry.store().models().await?;
let plugin_models = match registry.plugins() {
Some(plugins) => plugins.configured_models().await,
None => Vec::new(),
};
tracing::info!(
model_count = models.len(),
plugin_model_count = plugin_models.len(),
"appending BYOK models to Cursor AvailableModels"
);
let available_models = models.iter().map(available_model).collect::<Vec<_>>();
let mut available_models = models.iter().map(available_model).collect::<Vec<_>>();
available_models.extend(plugin_models.iter().map(available_plugin_model));
let local = AvailableModelsAddition {
model_names: models
.iter()
.map(|model| model.model_hash.clone())
.chain(plugin_models.iter().map(|model| model.id.clone()))
.collect(),
models: available_models,
}
@@ -252,12 +261,21 @@ pub async fn usable_models(
request: Request<Body>,
) -> Result<Response<Body>> {
let models = registry.store().models().await?;
let plugin_models = match registry.plugins() {
Some(plugins) => plugins.configured_models().await,
None => Vec::new(),
};
tracing::info!(
model_count = models.len(),
plugin_model_count = plugin_models.len(),
"appending BYOK models to Cursor GetUsableModels"
);
let local = UsableModelsAddition {
models: models.iter().map(usable_model).collect(),
models: models
.iter()
.map(usable_model)
.chain(plugin_models.iter().map(usable_plugin_model))
.collect(),
}
.encode_to_vec();
match proxy::forward_buffered(&proxy, request).await {
@@ -319,13 +337,19 @@ fn unary_payload(body: &Bytes) -> Result<(bool, &[u8])> {
}
fn available_model(model: &ModelConfig) -> AvailableModel {
let contexts = context_options(model);
let variants = model_variants(model, &contexts);
let contexts = context_options(model.context_window_tokens);
let tooltip = model_tooltip(model);
let variants = model_variants(
&model.model_hash,
&model.display_name,
&tooltip,
&contexts,
true,
);
let legacy_slugs = variants
.iter()
.filter_map(|variant| variant.legacy_slug.clone())
.collect();
let tooltip = model_tooltip(model);
AvailableModel {
name: model.model_hash.clone(),
default_on: true,
@@ -344,7 +368,7 @@ fn available_model(model: &ModelConfig) -> AvailableModel {
inputbox_short_model_name: Some(model.display_name.clone()),
supports_sandboxing: Some(true),
supports_cmd_k: Some(false),
parameter_definitions: model_parameters(&contexts),
parameter_definitions: model_parameters(&contexts, true),
variants,
legacy_slugs,
named_model_section_index: Some(1),
@@ -364,27 +388,30 @@ fn available_model(model: &ModelConfig) -> AvailableModel {
}
}
fn model_parameters(contexts: &[(String, String)]) -> Vec<ModelParameterDefinition> {
vec![
ModelParameterDefinition {
id: "context".into(),
name: "Context".into(),
markdown_tooltip: Some("Context size used to trigger conversation compaction.".into()),
parameter_type: Some(ModelParameterType {
boolean_parameter: None,
enum_parameter: Some(EnumParameter {
values: contexts
.iter()
.map(|(value, display_name)| EnumParameterValue {
value: value.clone(),
display_name: Some(display_name.clone()),
})
.collect(),
}),
fn model_parameters(
contexts: &[(String, String)],
thinking: bool,
) -> Vec<ModelParameterDefinition> {
let mut parameters = vec![ModelParameterDefinition {
id: "context".into(),
name: "Context".into(),
markdown_tooltip: Some("Context size used to trigger conversation compaction.".into()),
parameter_type: Some(ModelParameterType {
boolean_parameter: None,
enum_parameter: Some(EnumParameter {
values: contexts
.iter()
.map(|(value, display_name)| EnumParameterValue {
value: value.clone(),
display_name: Some(display_name.clone()),
})
.collect(),
}),
is_cycleable_by_hotkey: Some(false),
},
ModelParameterDefinition {
}),
is_cycleable_by_hotkey: Some(false),
}];
if thinking {
parameters.push(ModelParameterDefinition {
id: "reasoning".into(),
name: "Effort".into(),
markdown_tooltip: Some("Effort the model uses to generate its response.".into()),
@@ -401,44 +428,64 @@ fn model_parameters(contexts: &[(String, String)]) -> Vec<ModelParameterDefiniti
}),
}),
is_cycleable_by_hotkey: Some(true),
},
ModelParameterDefinition {
id: "fast".into(),
name: "Fast".into(),
markdown_tooltip: Some("Significantly faster but consumes more usage".into()),
parameter_type: Some(ModelParameterType {
boolean_parameter: Some(BooleanParameter {
values: vec![
BooleanParameterValue {
value: "false".into(),
display_name: None,
increases_model_cost: None,
},
BooleanParameterValue {
value: "true".into(),
display_name: Some("Fast".into()),
increases_model_cost: Some(true),
},
],
}),
enum_parameter: None,
});
}
parameters.push(ModelParameterDefinition {
id: "fast".into(),
name: "Fast".into(),
markdown_tooltip: Some("Significantly faster but consumes more usage".into()),
parameter_type: Some(ModelParameterType {
boolean_parameter: Some(BooleanParameter {
values: vec![
BooleanParameterValue {
value: "false".into(),
display_name: None,
increases_model_cost: None,
},
BooleanParameterValue {
value: "true".into(),
display_name: Some("Fast".into()),
increases_model_cost: Some(true),
},
],
}),
is_cycleable_by_hotkey: Some(false),
},
]
enum_parameter: None,
}),
is_cycleable_by_hotkey: Some(false),
});
parameters
}
fn model_variants(model: &ModelConfig, contexts: &[(String, String)]) -> Vec<ModelVariant> {
let mut variants = Vec::with_capacity(contexts.len() * EFFORTS.len() * 2);
fn model_variants(
name: &str,
display_name: &str,
tooltip: &TooltipData,
contexts: &[(String, String)],
thinking: bool,
) -> Vec<ModelVariant> {
// 非思考模型没有 Effort 轴,变体网格只剩 Context × Fast。
let efforts: &[Option<(&str, &str)>] = if thinking {
&[
Some(EFFORTS[0]),
Some(EFFORTS[1]),
Some(EFFORTS[2]),
Some(EFFORTS[3]),
Some(EFFORTS[4]),
]
} else {
&[None]
};
let mut variants = Vec::with_capacity(contexts.len() * efforts.len() * 2);
for (context, context_name) in contexts {
for (effort, effort_name) in EFFORTS {
for effort in efforts {
for fast in [false, true] {
variants.push(model_variant(
model,
name,
display_name,
tooltip,
context,
context_name,
effort,
effort_name,
*effort,
fast,
));
}
@@ -448,55 +495,67 @@ fn model_variants(model: &ModelConfig, contexts: &[(String, String)]) -> Vec<Mod
}
fn model_variant(
model: &ModelConfig,
name: &str,
display_name: &str,
tooltip: &TooltipData,
context: &str,
context_name: &str,
effort: &str,
effort_name: &str,
effort: Option<(&str, &str)>,
fast: bool,
) -> ModelVariant {
let mut suffix = Vec::with_capacity(3);
if context != DEFAULT_CONTEXT {
suffix.push(context_name);
}
suffix.push(effort_name);
if let Some((_, effort_name)) = effort {
suffix.push(effort_name);
}
if fast {
suffix.push("Fast");
}
let suffix = suffix.join(" ");
let display_name = format!(
"{} <span style=\"color: var(--cursor-text-tertiary);\">{suffix}</span>",
model.display_name
);
let is_default = context == DEFAULT_CONTEXT && effort == "high" && !fast;
let display_name = if suffix.is_empty() {
display_name.to_owned()
} else {
format!(
"{display_name} <span style=\"color: var(--cursor-text-tertiary);\">{suffix}</span>"
)
};
let is_default =
context == DEFAULT_CONTEXT && !fast && effort.is_none_or(|(effort, _)| effort == "high");
let mut parameter_values = vec![ModelParameterValue {
id: "context".into(),
value: context.into(),
}];
if let Some((effort, _)) = effort {
parameter_values.push(ModelParameterValue {
id: "reasoning".into(),
value: effort.into(),
});
}
parameter_values.push(ModelParameterValue {
id: "fast".into(),
value: fast.to_string(),
});
ModelVariant {
parameter_values: vec![
ModelParameterValue {
id: "context".into(),
value: context.into(),
},
ModelParameterValue {
id: "reasoning".into(),
value: effort.into(),
},
ModelParameterValue {
id: "fast".into(),
value: fast.to_string(),
},
],
parameter_values,
display_name: display_name.clone(),
is_max_mode: false,
is_default_max_config: is_default.then_some(true),
is_default_non_max_config: is_default.then_some(true),
tooltip_data: Some(model_tooltip(model)),
tooltip_data: Some(tooltip.clone()),
display_name_outside_picker: Some(display_name),
variant_string_representation: Some(format!(
"{}[context={context},reasoning={effort},fast={fast}]",
model.model_hash
)),
variant_string_representation: Some(match effort {
Some((effort, _)) => {
format!("{name}[context={context},reasoning={effort},fast={fast}]")
}
None => format!("{name}[context={context},fast={fast}]"),
}),
legacy_slug: Some(format!(
"{}-{context}-{effort}{}",
model.model_hash,
"{name}-{context}{}{}",
effort
.map(|(effort, _)| format!("-{effort}"))
.unwrap_or_default(),
if fast { "-fast" } else { "" }
)),
}
@@ -508,6 +567,68 @@ fn model_tooltip(model: &ModelConfig) -> TooltipData {
}
}
fn available_plugin_model(model: &PluginModelDescriptor) -> AvailableModel {
let tooltip = TooltipData {
markdown_content: model.description.clone(),
};
let contexts = context_options(model.context_window_tokens);
let variants = model_variants(
&model.id,
&model.display_name,
&tooltip,
&contexts,
model.thinking,
);
let legacy_slugs = variants
.iter()
.filter_map(|variant| variant.legacy_slug.clone())
.collect();
AvailableModel {
name: model.id.clone(),
default_on: true,
supports_agent: Some(true),
degradation_status: Some(0),
tooltip_data: Some(tooltip.clone()),
supports_thinking: Some(model.thinking),
supports_images: Some(model.images),
supports_max_mode: Some(false),
client_display_name: Some(model.display_name.clone()),
server_model_name: Some(model.id.clone()),
supports_non_max_mode: Some(true),
tooltip_data_for_max_mode: Some(tooltip.clone()),
is_recommended_for_background_composer: Some(false),
supports_plan_mode: Some(true),
inputbox_short_model_name: Some(model.display_name.clone()),
supports_sandboxing: Some(true),
supports_cmd_k: Some(false),
parameter_definitions: model_parameters(&contexts, model.thinking),
variants,
legacy_slugs,
named_model_section_index: Some(1),
vendor_name: Some(model.provider_type.clone()),
vendor: Some(AvailableModelVendor {
id: 6,
display_name: model.provider_type.clone(),
}),
model_picker_badges: vec![ModelPickerBadge {
label: model.provider_type.clone(),
variant: 1,
dismiss_on_selection: false,
}],
}
}
fn usable_plugin_model(model: &PluginModelDescriptor) -> agent::ModelDetails {
agent::ModelDetails {
model_id: model.id.clone(),
display_model_id: model.id.clone(),
display_name: model.display_name.clone(),
display_name_short: model.display_name.clone(),
thinking_details: model.thinking.then(agent::ThinkingDetails::default),
..Default::default()
}
}
fn usable_model(model: &ModelConfig) -> agent::ModelDetails {
agent::ModelDetails {
model_id: model.model_hash.clone(),
+27
View File
@@ -9,6 +9,7 @@ use crate::{
conversation::ConversationRegistry, prompting::PromptCompiler,
services::observability::CursorTraceRecorder,
},
plugin::PluginRegistry,
provider::Provider,
search::WebCache,
store::Store,
@@ -28,6 +29,7 @@ struct RegistryInner {
route_changed: Notify,
store: Store,
web_cache: WebCache,
plugins: Option<PluginRegistry>,
conversations: ConversationRegistry,
}
@@ -47,6 +49,26 @@ impl TransportRegistry {
provider: Arc<dyn Provider>,
compiler: PromptCompiler,
web_cache: WebCache,
) -> Self {
Self::build(store, provider, compiler, web_cache, None)
}
pub fn with_plugins(
store: Store,
provider: Arc<dyn Provider>,
compiler: PromptCompiler,
web_cache: WebCache,
plugins: PluginRegistry,
) -> Self {
Self::build(store, provider, compiler, web_cache, Some(plugins))
}
fn build(
store: Store,
provider: Arc<dyn Provider>,
compiler: PromptCompiler,
web_cache: WebCache,
plugins: Option<PluginRegistry>,
) -> Self {
Self {
inner: Arc::new(RegistryInner {
@@ -61,6 +83,7 @@ impl TransportRegistry {
),
store,
web_cache,
plugins,
}),
}
}
@@ -73,6 +96,10 @@ impl TransportRegistry {
&self.inner.web_cache
}
pub fn plugins(&self) -> Option<&PluginRegistry> {
self.inner.plugins.as_ref()
}
pub fn conversations(&self) -> &ConversationRegistry {
&self.inner.conversations
}
+5
View File
@@ -18,6 +18,9 @@ pub enum ProviderType {
OpenAiResponses,
#[serde(rename = "anthropic")]
Anthropic,
/// 插件执行的调用;协议细节在插件内部,核心只按统一事件流记录。
#[serde(rename = "plugin")]
Plugin,
}
impl ProviderType {
@@ -26,6 +29,7 @@ impl ProviderType {
Self::OpenAiChat => "openai-chat",
Self::OpenAiResponses => "openai-responses",
Self::Anthropic => "anthropic",
Self::Plugin => "plugin",
}
}
}
@@ -44,6 +48,7 @@ impl FromStr for ProviderType {
"openai-chat" => Ok(Self::OpenAiChat),
"openai-responses" => Ok(Self::OpenAiResponses),
"anthropic" => Ok(Self::Anthropic),
"plugin" => Ok(Self::Plugin),
_ => Err(Error::Config(format!("unsupported provider type: {value}"))),
}
}
+3 -1
View File
@@ -23,7 +23,9 @@ mod usage {
pub(crate) fn context_input_tokens(self, provider: ProviderType) -> Option<u64> {
let input = self.input_tokens?;
match provider {
ProviderType::OpenAiChat | ProviderType::OpenAiResponses => Some(input),
ProviderType::OpenAiChat | ProviderType::OpenAiResponses | ProviderType::Plugin => {
Some(input)
}
ProviderType::Anthropic => input
.checked_add(self.cache_read_tokens.unwrap_or_default())?
.checked_add(self.cache_write_tokens.unwrap_or_default()),
+85
View File
@@ -0,0 +1,85 @@
//! Materializes built-in plugins bundled in the binary into the managed dir.
use std::path::PathBuf;
use super::definition::write_if_changed;
use crate::{config, Result};
/// 随二进制打包的内置插件文件;发布构建没有源码目录,靠这里落盘。
const CODEX_AUTH: &[(&str, &str)] = &[
(
"plugin.json",
include_str!(concat!(
env!("CARGO_MANIFEST_DIR"),
"/plugins/build-in/codex-auth/plugin.json"
)),
),
(
"main.ts",
include_str!(concat!(
env!("CARGO_MANIFEST_DIR"),
"/plugins/build-in/codex-auth/main.ts"
)),
),
(
"provider.ts",
include_str!(concat!(
env!("CARGO_MANIFEST_DIR"),
"/plugins/build-in/codex-auth/provider.ts"
)),
),
(
"models.ts",
include_str!(concat!(
env!("CARGO_MANIFEST_DIR"),
"/plugins/build-in/codex-auth/models.ts"
)),
),
(
"oauth.ts",
include_str!(concat!(
env!("CARGO_MANIFEST_DIR"),
"/plugins/build-in/codex-auth/oauth.ts"
)),
),
(
"resources.ts",
include_str!(concat!(
env!("CARGO_MANIFEST_DIR"),
"/plugins/build-in/codex-auth/resources.ts"
)),
),
(
"assets/codex.svg",
include_str!(concat!(
env!("CARGO_MANIFEST_DIR"),
"/plugins/build-in/codex-auth/assets/codex.svg"
)),
),
];
/// 把内置插件写入受管目录并返回该目录,作为插件目录的扫描根之一。
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)
}
fn write_plugin(directory: &std::path::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");
std::fs::create_dir_all(parent)?;
#[cfg(unix)]
{
use std::os::unix::fs::PermissionsExt;
std::fs::set_permissions(parent, std::fs::Permissions::from_mode(0o700))?;
}
write_if_changed(&path, content)?;
#[cfg(unix)]
{
use std::os::unix::fs::PermissionsExt;
std::fs::set_permissions(&path, std::fs::Permissions::from_mode(0o600))?;
}
}
Ok(())
}
+285
View File
@@ -0,0 +1,285 @@
//! Discovers plugin manifests and evaluates serializable TypeScript definitions.
use std::{
collections::BTreeMap,
fs,
path::{Path, PathBuf},
};
use base64::{engine::general_purpose::STANDARD, Engine};
use super::{
definition::PluginDefinitionLoader,
descriptor::PluginModuleDefinition,
manifest::{validate_id, PluginManifest},
};
use crate::{config, Error, Result};
const MANIFEST_FILE_NAME: &str = "plugin.json";
const MAX_ICON_BYTES: u64 = 1024 * 1024;
#[derive(Clone)]
pub struct PluginCatalog {
roots: Vec<PathBuf>,
definition_loader: PluginDefinitionLoader,
}
#[derive(Clone)]
pub(crate) struct PluginEntry {
pub directory: PathBuf,
pub entry: PathBuf,
pub manifest: PluginManifest,
pub definition: PluginModuleDefinition,
pub icon: String,
}
impl PluginCatalog {
pub fn managed() -> Result<Self> {
let installed = config::managed_data_dir()?.join("plugins/installed");
fs::create_dir_all(&installed)?;
#[cfg(unix)]
{
use std::os::unix::fs::PermissionsExt;
fs::set_permissions(&installed, fs::Permissions::from_mode(0o700))?;
}
// 扫描顺序即优先级:用户安装目录 > 源码内置目录(仅 debug,便于热改)
// > 随二进制打包后落盘的内置目录;同 ID 时靠前的覆盖靠后的。
let mut roots = vec![installed];
#[cfg(debug_assertions)]
roots.push(PathBuf::from(env!("CARGO_MANIFEST_DIR")).join("plugins/build-in"));
roots.push(super::builtin::materialize()?);
Ok(Self {
roots,
definition_loader: PluginDefinitionLoader::managed()?,
})
}
pub(crate) fn loader(&self) -> &PluginDefinitionLoader {
&self.definition_loader
}
pub(crate) async fn entries(&self, executable: &Path) -> Vec<PluginEntry> {
let mut plugins = BTreeMap::new();
for root in &self.roots {
let mut directories = match child_directories(root) {
Ok(value) => value,
Err(error) => {
tracing::warn!(path = %root.display(), %error, "failed to scan plugin directory");
continue;
}
};
directories.sort();
for directory in directories {
match load_plugin(&directory, &self.definition_loader, executable).await {
Ok(entry) => {
if plugins.contains_key(&entry.manifest.id) {
tracing::warn!(plugin = %entry.manifest.id, path = %directory.display(), "ignoring duplicate plugin");
} else {
plugins.insert(entry.manifest.id.clone(), entry);
}
}
Err(error) => {
tracing::warn!(path = %directory.display(), %error, "ignoring invalid plugin")
}
}
}
}
plugins.into_values().collect()
}
pub(crate) fn manifests(&self) -> Vec<(PluginManifest, String)> {
let mut plugins = BTreeMap::new();
for root in &self.roots {
let Ok(mut directories) = child_directories(root) else {
continue;
};
directories.sort();
for directory in directories {
let loaded = (|| -> Result<_> {
let manifest: PluginManifest =
serde_json::from_slice(&fs::read(directory.join(MANIFEST_FILE_NAME))?)?;
manifest.validate(&directory)?;
let icon = icon_data_url(&directory, &manifest.icon)?;
Ok((manifest, icon))
})();
if let Ok((manifest, icon)) = loaded {
plugins
.entry(manifest.id.clone())
.or_insert((manifest, icon));
}
}
}
plugins.into_values().collect()
}
}
fn child_directories(root: &Path) -> Result<Vec<PathBuf>> {
if !root.exists() {
return Ok(Vec::new());
}
let mut directories = Vec::new();
for entry in fs::read_dir(root)? {
let entry = entry?;
if entry.file_type()?.is_dir() && !entry.file_name().to_string_lossy().starts_with('.') {
directories.push(entry.path());
}
}
Ok(directories)
}
async fn load_plugin(
directory: &Path,
loader: &PluginDefinitionLoader,
executable: &Path,
) -> Result<PluginEntry> {
let manifest: PluginManifest =
serde_json::from_slice(&fs::read(directory.join(MANIFEST_FILE_NAME))?)?;
manifest.validate(directory)?;
let icon = icon_data_url(directory, &manifest.icon)?;
let entry = directory.join(&manifest.entry).canonicalize()?;
let definition = loader.load(executable, directory, &entry).await?;
validate_definition(&manifest.id, &definition)?;
Ok(PluginEntry {
directory: directory.to_path_buf(),
entry,
manifest,
definition,
icon,
})
}
/// 显示文本必须是非空字符串,或全为非空字符串的 locale 映射。
fn validate_localized_text(value: &serde_json::Value, label: &str) -> Result<()> {
match value {
serde_json::Value::String(text) if !text.trim().is_empty() => Ok(()),
serde_json::Value::Object(map)
if !map.is_empty()
&& map
.values()
.all(|entry| entry.as_str().is_some_and(|text| !text.trim().is_empty())) =>
{
Ok(())
}
_ => Err(Error::Config(format!(
"{label} must be a non-empty string or a locale map of non-empty strings"
))),
}
}
fn validate_definition(plugin_id: &str, definition: &PluginModuleDefinition) -> Result<()> {
if definition.providers.is_empty() {
return Err(Error::Config(format!(
"plugin '{plugin_id}' must define at least one provider"
)));
}
let mut provider_ids = std::collections::HashSet::new();
for provider in &definition.providers {
validate_id(&provider.id, "plugin provider id")?;
validate_localized_text(
&provider.display_name,
&format!(
"plugin '{plugin_id}' provider '{}' displayName",
provider.id
),
)?;
if provider.provider_type.trim().is_empty() {
return Err(Error::Config(format!(
"plugin '{plugin_id}' provider '{}' requires providerType",
provider.id
)));
}
if !provider_ids.insert(provider.id.clone()) {
return Err(Error::Config(format!(
"plugin '{plugin_id}' contains duplicate provider '{}'",
provider.id
)));
}
if let Some(resource_type) = &provider.resource_type {
if !definition
.resources
.iter()
.any(|resource| &resource.resource_type == resource_type)
{
return Err(Error::Config(format!(
"plugin '{plugin_id}' provider '{}' consumes undeclared resource '{resource_type}'",
provider.id
)));
}
}
}
let mut resource_types = std::collections::HashSet::new();
for resource in &definition.resources {
validate_id(&resource.resource_type, "plugin resource type")?;
validate_localized_text(
&resource.display_name,
&format!(
"plugin '{plugin_id}' resource '{}' displayName",
resource.resource_type
),
)?;
if !resource_types.insert(resource.resource_type.clone()) {
return Err(Error::Config(format!(
"plugin '{plugin_id}' contains duplicate resource type '{}'",
resource.resource_type
)));
}
for method in &resource.add {
validate_id(&method.id, "plugin add method id")?;
if method.method_type != super::descriptor::OAUTH2_ADD_METHOD {
return Err(Error::Config(format!(
"plugin '{plugin_id}' add method '{}' uses unsupported type '{}'",
method.id, method.method_type
)));
}
}
}
Ok(())
}
fn icon_data_url(directory: &Path, relative: &str) -> Result<String> {
let root = directory.canonicalize()?;
let path = directory.join(relative).canonicalize()?;
if !path.starts_with(&root) {
return Err(Error::Config(format!(
"plugin icon escapes its directory: {relative}"
)));
}
if fs::metadata(&path)?.len() > MAX_ICON_BYTES {
return Err(Error::Config(format!(
"plugin icon exceeds {MAX_ICON_BYTES} bytes: {relative}"
)));
}
let extension = path
.extension()
.and_then(|value| value.to_str())
.unwrap_or_default()
.to_ascii_lowercase();
let mime = match extension.as_str() {
"svg" => "image/svg+xml",
"png" => "image/png",
"webp" => "image/webp",
_ => {
return Err(Error::Config(format!(
"unsupported plugin icon: {relative}"
)))
}
};
Ok(format!(
"data:{mime};base64,{}",
STANDARD.encode(fs::read(path)?)
))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn repository_examples_have_valid_static_manifests() {
let root = PathBuf::from(env!("CARGO_MANIFEST_DIR")).join("plugins/build-in");
let sdk = tempfile::tempdir().unwrap();
let catalog = PluginCatalog {
roots: vec![root],
definition_loader: PluginDefinitionLoader::for_test(sdk.path()).unwrap(),
};
assert!(!catalog.manifests().is_empty());
}
}
+161
View File
@@ -0,0 +1,161 @@
//! Stores plugin-owned JSON with private permissions and atomic replacement.
use std::{
collections::HashMap,
path::{Path, PathBuf},
sync::Arc,
};
use parking_lot::Mutex;
use tokio::sync::Mutex as AsyncMutex;
use crate::{config, Error, Result};
#[derive(Clone)]
pub struct PluginDataStore {
root: PathBuf,
locks: Arc<Mutex<HashMap<String, Arc<AsyncMutex<()>>>>>,
}
impl PluginDataStore {
pub fn managed() -> Result<Self> {
Self::new(config::managed_data_dir()?.join("plugins/data"))
}
#[cfg(test)]
pub(super) fn for_test(root: PathBuf) -> Result<Self> {
Self::new(root)
}
fn new(root: PathBuf) -> Result<Self> {
std::fs::create_dir_all(&root)?;
set_directory_permissions(&root)?;
Ok(Self {
root,
locks: Arc::new(Mutex::new(HashMap::new())),
})
}
pub async fn read(&self, plugin_id: &str, key: &str) -> Result<serde_json::Value> {
let path = self.path(plugin_id, key)?;
let lock = self.lock(plugin_id);
let _guard = lock.lock().await;
match tokio::fs::read(path).await {
Ok(bytes) => Ok(serde_json::from_slice(&bytes)?),
Err(error) if error.kind() == std::io::ErrorKind::NotFound => {
Ok(serde_json::Value::Null)
}
Err(error) => Err(error.into()),
}
}
pub async fn update(
&self,
plugin_id: &str,
key: &str,
value: &serde_json::Value,
) -> Result<()> {
let path = self.path(plugin_id, key)?;
let lock = self.lock(plugin_id);
let _guard = lock.lock().await;
let directory = path.parent().expect("plugin data path has a parent");
tokio::fs::create_dir_all(directory).await?;
set_directory_permissions(directory)?;
let temporary = directory.join(format!(".{key}.{}.tmp", uuid::Uuid::new_v4()));
let bytes = serde_json::to_vec_pretty(value)?;
tokio::fs::write(&temporary, bytes).await?;
set_file_permissions(&temporary)?;
let file = tokio::fs::OpenOptions::new()
.read(true)
.open(&temporary)
.await?;
file.sync_all().await?;
drop(file);
#[cfg(windows)]
if path.exists() {
tokio::fs::remove_file(&path).await?;
}
tokio::fs::rename(&temporary, &path).await?;
set_file_permissions(&path)?;
Ok(())
}
pub async fn clear(&self, plugin_id: &str) -> Result<()> {
validate_component(plugin_id, "plugin id")?;
let lock = self.lock(plugin_id);
let _guard = lock.lock().await;
let path = self.root.join(plugin_id);
match tokio::fs::remove_dir_all(path).await {
Ok(()) => Ok(()),
Err(error) if error.kind() == std::io::ErrorKind::NotFound => Ok(()),
Err(error) => Err(error.into()),
}
}
fn path(&self, plugin_id: &str, key: &str) -> Result<PathBuf> {
validate_component(plugin_id, "plugin id")?;
validate_component(key, "plugin data key")?;
Ok(self.root.join(plugin_id).join(format!("{key}.json")))
}
fn lock(&self, plugin_id: &str) -> Arc<AsyncMutex<()>> {
self.locks
.lock()
.entry(plugin_id.to_owned())
.or_insert_with(|| Arc::new(AsyncMutex::new(())))
.clone()
}
}
fn validate_component(value: &str, label: &str) -> Result<()> {
if value.is_empty()
|| value.len() > 128
|| !value
.bytes()
.all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'.' | b'_' | b'-'))
{
return Err(Error::Config(format!("invalid {label}: {value}")));
}
Ok(())
}
fn set_directory_permissions(path: &Path) -> Result<()> {
#[cfg(unix)]
{
use std::os::unix::fs::PermissionsExt;
std::fs::set_permissions(path, std::fs::Permissions::from_mode(0o700))?;
}
Ok(())
}
fn set_file_permissions(path: &Path) -> Result<()> {
#[cfg(unix)]
{
use std::os::unix::fs::PermissionsExt;
std::fs::set_permissions(path, std::fs::Permissions::from_mode(0o600))?;
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn writes_reads_and_removes_json() {
let root = tempfile::tempdir().unwrap();
let store = PluginDataStore::new(root.path().join("data")).unwrap();
store
.update(
"com.example",
"state",
&serde_json::json!({"token":"secret"}),
)
.await
.unwrap();
assert_eq!(
store.read("com.example", "state").await.unwrap()["token"],
"secret"
);
store.clear("com.example").await.unwrap();
assert!(store.read("com.example", "state").await.unwrap().is_null());
}
}
+213
View File
@@ -0,0 +1,213 @@
//! Evaluates TypeScript plugin definitions through the host-owned virtual module.
use std::{
path::{Path, PathBuf},
process::Stdio,
time::Duration,
};
use tokio::io::AsyncReadExt;
use super::descriptor::PluginModuleDefinition;
use crate::{config, Error, Result};
const DEFINITION_TIMEOUT: Duration = Duration::from_secs(10);
const MAX_OUTPUT_BYTES: u64 = 2 * 1024 * 1024;
const OUTPUT_PREFIX: &str = "CURSOR_BYOK_PLUGIN_DEFINITION:";
const IMPORT_MAP: &str = r#"{"imports":{"cursor-byok:plugin":"./plugin.ts","cursor-byok:provider":"./provider.ts","cursor-byok:model":"./model.ts","cursor-byok:resource":"./resource.ts","cursor-byok:protocol/openai-responses":"./protocol/openai_responses.ts"}}"#;
#[derive(Clone)]
pub struct PluginDefinitionLoader {
sdk_dir: PathBuf,
import_map: PathBuf,
collector: PathBuf,
worker: PathBuf,
deno_dir: PathBuf,
}
impl PluginDefinitionLoader {
pub fn managed() -> Result<Self> {
Self::in_directory(config::managed_data_dir()?.join("plugins/runtime/sdk/v1"))
}
#[cfg(test)]
pub(super) fn for_test(root: &Path) -> Result<Self> {
Self::in_directory(root.join(".plugin-sdk"))
}
fn in_directory(sdk_dir: PathBuf) -> Result<Self> {
std::fs::create_dir_all(&sdk_dir)?;
std::fs::create_dir_all(sdk_dir.join("protocol"))?;
let import_map = sdk_dir.join("import-map.json");
let collector = sdk_dir.join("collect.ts");
let worker = sdk_dir.join("worker.ts");
let deno_dir = sdk_dir.join("cache");
std::fs::create_dir_all(&deno_dir)?;
let modules = [
(&import_map, IMPORT_MAP),
(&collector, include_str!("sdk/collect.ts")),
(&worker, include_str!("sdk/worker.ts")),
(&sdk_dir.join("plugin.ts"), include_str!("sdk/plugin.ts")),
(
&sdk_dir.join("provider.ts"),
include_str!("sdk/provider.ts"),
),
(&sdk_dir.join("model.ts"), include_str!("sdk/model.ts")),
(
&sdk_dir.join("resource.ts"),
include_str!("sdk/resource.ts"),
),
(
&sdk_dir.join("protocol/openai_responses.ts"),
include_str!("sdk/protocol/openai_responses.ts"),
),
];
for (path, content) in &modules {
write_if_changed(path, content)?;
}
#[cfg(unix)]
{
use std::os::unix::fs::PermissionsExt;
std::fs::set_permissions(&sdk_dir, std::fs::Permissions::from_mode(0o700))?;
std::fs::set_permissions(
sdk_dir.join("protocol"),
std::fs::Permissions::from_mode(0o700),
)?;
std::fs::set_permissions(&deno_dir, std::fs::Permissions::from_mode(0o700))?;
for (path, _) in &modules {
std::fs::set_permissions(path, std::fs::Permissions::from_mode(0o600))?;
}
}
Ok(Self {
sdk_dir,
import_map,
collector,
worker,
deno_dir,
})
}
pub fn worker_path(&self) -> &Path {
&self.worker
}
pub fn import_map(&self) -> &Path {
&self.import_map
}
pub fn sdk_dir(&self) -> &Path {
&self.sdk_dir
}
pub fn deno_dir(&self) -> &Path {
&self.deno_dir
}
pub async fn load(
&self,
executable: &Path,
plugin_directory: &Path,
entry: &Path,
) -> Result<PluginModuleDefinition> {
let entry_url = file_url(entry)?;
let mut command = tokio::process::Command::new(executable);
command
.arg("run")
.arg("--quiet")
.arg("--no-config")
.arg("--no-lock")
.arg("--no-npm")
.arg("--no-remote")
.arg("--no-prompt")
.arg(format!("--allow-read={}", plugin_directory.display()))
.arg(format!("--allow-read={}", self.sdk_dir.display()))
.arg(format!("--import-map={}", self.import_map.display()))
.arg(&self.collector)
.arg(entry_url.as_str())
.env("DENO_DIR", &self.deno_dir)
.env("DENO_NO_UPDATE_CHECK", "1")
.stdin(Stdio::null())
.stdout(Stdio::piped())
.stderr(Stdio::piped())
.kill_on_drop(true);
let mut child = command.spawn()?;
let stdout = child
.stdout
.take()
.ok_or_else(|| Error::Config("cannot capture plugin definition output".into()))?;
let stderr = child
.stderr
.take()
.ok_or_else(|| Error::Config("cannot capture plugin definition error output".into()))?;
let (stdout, stderr, status) = tokio::time::timeout(DEFINITION_TIMEOUT, async move {
let (stdout, stderr, status) =
tokio::join!(read_limited(stdout), read_limited(stderr), child.wait());
Ok::<_, Error>((stdout?, stderr?, status?))
})
.await
.map_err(|_| Error::Config("plugin definition evaluation timed out".into()))??;
if !status.success() {
return Err(Error::Config(format!(
"plugin definition evaluation failed: {}",
String::from_utf8_lossy(&stderr).trim()
)));
}
parse_definition_output(&stdout)
}
}
pub(super) fn file_url(path: &Path) -> Result<url::Url> {
url::Url::from_file_path(path).map_err(|_| {
Error::Config(format!(
"plugin entry path is not a valid file URL: {}",
path.display()
))
})
}
async fn read_limited(reader: impl tokio::io::AsyncRead + Unpin) -> Result<Vec<u8>> {
let mut bytes = Vec::new();
reader
.take(MAX_OUTPUT_BYTES + 1)
.read_to_end(&mut bytes)
.await?;
if bytes.len() as u64 > MAX_OUTPUT_BYTES {
return Err(Error::Config(
"plugin definition output is larger than allowed".into(),
));
}
Ok(bytes)
}
fn parse_definition_output(output: &[u8]) -> Result<PluginModuleDefinition> {
let output = String::from_utf8(output.to_vec()).map_err(|error| {
Error::Config(format!("plugin definition output is not UTF-8: {error}"))
})?;
let json = output
.lines()
.rev()
.find_map(|line| line.strip_prefix(OUTPUT_PREFIX))
.ok_or_else(|| Error::Config("plugin definition did not produce a descriptor".into()))?;
Ok(serde_json::from_str(json)?)
}
pub(super) fn write_if_changed(path: &Path, content: &str) -> Result<()> {
if std::fs::read(path).is_ok_and(|current| current == content.as_bytes()) {
return Ok(());
}
std::fs::write(path, content)?;
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn parses_descriptor_marker() {
let output = br#"CURSOR_BYOK_PLUGIN_DEFINITION:{"providers":[{"id":"codex","displayName":"OpenAI Codex","description":null,"providerType":"openai","resourceType":"chatgpt-account","hasModels":true}],"resources":[{"type":"chatgpt-account","displayName":"ChatGPT accounts","add":[{"type":"oauth2.0","id":"chatgpt-device","displayName":"Sign in","description":null}],"import":{"displayName":"Import","description":null,"accept":[".json"],"multiple":true},"canRefresh":true,"canRemove":false}]}"#;
let descriptor = parse_definition_output(output).unwrap();
assert_eq!(descriptor.providers[0].id, "codex");
assert_eq!(
descriptor.providers[0].resource_type.as_deref(),
Some("chatgpt-account")
);
assert_eq!(descriptor.resources[0].add[0].method_type, "oauth2.0");
assert!(descriptor.resources[0].import.is_some());
}
}
+233
View File
@@ -0,0 +1,233 @@
//! Defines serializable plugin capability definitions and desktop descriptors.
use serde::{Deserialize, Serialize};
use super::state::{ResourceRecord, ResourceState, StoredModel};
/// 由 collect.ts 输出的能力摘要;不含任何可执行内容。
#[derive(Clone, Debug, Deserialize)]
#[serde(rename_all = "camelCase", deny_unknown_fields)]
pub struct PluginModuleDefinition {
pub providers: Vec<ProviderDefinition>,
#[serde(default)]
pub resources: Vec<ResourceDefinition>,
}
/// 插件提供的显示文本:纯字符串或 locale → 文本映射;核心原样透传,由前端解析。
pub type LocalizedText = serde_json::Value;
#[derive(Clone, Debug, Deserialize)]
#[serde(rename_all = "camelCase", deny_unknown_fields)]
pub struct ProviderDefinition {
pub id: String,
pub display_name: LocalizedText,
#[serde(default)]
pub description: LocalizedText,
pub provider_type: String,
#[serde(default)]
pub resource_type: Option<String>,
pub has_models: bool,
}
#[derive(Clone, Debug, Deserialize)]
#[serde(rename_all = "camelCase", deny_unknown_fields)]
pub struct ResourceDefinition {
#[serde(rename = "type")]
pub resource_type: String,
pub display_name: LocalizedText,
#[serde(default)]
pub add: Vec<AddMethodDefinition>,
#[serde(default)]
pub import: Option<ImportDefinition>,
pub can_refresh: bool,
pub can_remove: bool,
}
#[derive(Clone, Debug, Deserialize, Serialize)]
#[serde(rename_all = "camelCase", deny_unknown_fields)]
pub struct AddMethodDefinition {
#[serde(rename = "type")]
pub method_type: String,
pub id: String,
pub display_name: LocalizedText,
#[serde(default)]
pub description: LocalizedText,
}
#[derive(Clone, Debug, Deserialize, Serialize)]
#[serde(rename_all = "camelCase", deny_unknown_fields)]
pub struct ImportDefinition {
pub display_name: LocalizedText,
#[serde(default)]
pub description: LocalizedText,
pub accept: Vec<String>,
pub multiple: bool,
}
pub const OAUTH2_ADD_METHOD: &str = "oauth2.0";
/// 桌面端看到的插件全貌。
#[derive(Clone, Debug, Serialize)]
#[serde(rename_all = "camelCase")]
pub struct PluginDescriptor {
pub id: String,
pub name: String,
pub author: Option<String>,
pub icon: String,
pub providers: Vec<PluginProviderDescriptor>,
pub resources: Vec<PluginResourceDescriptor>,
}
#[derive(Clone, Debug, Serialize)]
#[serde(rename_all = "camelCase")]
pub struct PluginProviderDescriptor {
pub id: String,
pub plugin_id: String,
pub display_name: LocalizedText,
pub description: LocalizedText,
pub provider_type: String,
pub resource_type: Option<String>,
pub has_models: bool,
/// 已满足调用条件:模型目录非空,且需要资源时至少有一条资源。
pub configured: bool,
pub models: Vec<PluginModelDescriptor>,
}
/// 一个可直接被 Cursor 调用的插件模型。
#[derive(Clone, Debug, Serialize)]
#[serde(rename_all = "camelCase")]
pub struct PluginModelDescriptor {
/// 稳定模型 ID:`plugin:<plugin>/<provider>/<model>`。
pub id: String,
pub plugin_id: String,
pub plugin_name: String,
pub provider_id: String,
pub model_id: String,
pub display_name: String,
pub description: Option<String>,
pub icon: String,
pub provider_type: String,
pub context_window_tokens: Option<u64>,
pub max_output_tokens: Option<u64>,
pub thinking: bool,
pub images: bool,
}
#[derive(Clone, Debug, Serialize)]
#[serde(rename_all = "camelCase")]
pub struct PluginResourceDescriptor {
#[serde(rename = "type")]
pub resource_type: String,
pub display_name: LocalizedText,
pub add: Vec<AddMethodDefinition>,
pub import: Option<ImportDefinition>,
pub can_refresh: bool,
pub can_remove: bool,
pub resources: Vec<PluginResourceView>,
}
/// 单条资源的对外投影;凭证保留在核心存储,不进入该结构。
#[derive(Clone, Debug, Serialize)]
#[serde(rename_all = "camelCase")]
pub struct PluginResourceView {
pub id: String,
pub state: ResourceState,
pub display_name: String,
pub description: LocalizedText,
pub metrics: Vec<ResourceMetric>,
pub created_at_ms: i64,
}
#[derive(Clone, Debug, Deserialize, Serialize)]
#[serde(rename_all = "camelCase", deny_unknown_fields)]
pub struct ResourceMetric {
pub id: String,
pub label: LocalizedText,
pub unit: String,
pub value: f64,
#[serde(default)]
pub reset_at_ms: Option<i64>,
}
/// 插件对一条资源的展示投影(resource.present 的返回值)。
#[derive(Clone, Debug, Deserialize)]
#[serde(rename_all = "camelCase", deny_unknown_fields)]
pub struct ResourcePresentation {
pub display_name: String,
#[serde(default)]
pub description: LocalizedText,
#[serde(default)]
pub metrics: Vec<ResourceMetric>,
}
impl PluginResourceView {
pub fn from_record(record: &ResourceRecord, presentation: ResourcePresentation) -> Self {
Self {
id: record.id.clone(),
state: record.state.clone(),
display_name: presentation.display_name,
description: presentation.description,
metrics: presentation.metrics,
created_at_ms: record.created_at_ms,
}
}
}
pub const ADAPTER_ID_PREFIX: &str = "plugin:";
pub fn model_id(plugin_id: &str, provider_id: &str, model_id: &str) -> String {
format!("{ADAPTER_ID_PREFIX}{plugin_id}/{provider_id}/{model_id}")
}
/// 解析稳定模型 ID;上游模型段允许包含 `/`。
pub fn parse_model_id(value: &str) -> Option<(&str, &str, &str)> {
let rest = value.strip_prefix(ADAPTER_ID_PREFIX)?;
let (plugin_id, rest) = rest.split_once('/')?;
let (provider_id, model_id) = rest.split_once('/')?;
(!plugin_id.is_empty() && !provider_id.is_empty() && !model_id.is_empty()).then_some((
plugin_id,
provider_id,
model_id,
))
}
impl PluginModelDescriptor {
pub fn new(
plugin_id: &str,
plugin_name: &str,
icon: &str,
provider: &ProviderDefinition,
model: &StoredModel,
) -> Self {
Self {
id: model_id(plugin_id, &provider.id, &model.id),
plugin_id: plugin_id.to_owned(),
plugin_name: plugin_name.to_owned(),
provider_id: provider.id.clone(),
model_id: model.id.clone(),
display_name: model.display_name.clone(),
description: model.description.clone(),
icon: icon.to_owned(),
provider_type: provider.provider_type.clone(),
context_window_tokens: model.context_window_tokens,
max_output_tokens: model.max_output_tokens,
thinking: model.thinking,
images: model.images,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn parses_stable_model_ids_with_slashes() {
let id = model_id("dev.example", "codex", "org/gpt-5");
assert_eq!(
parse_model_id(&id),
Some(("dev.example", "codex", "org/gpt-5"))
);
assert_eq!(parse_model_id("plugin:only/one"), None);
assert_eq!(parse_model_id("model-hash"), None);
}
}
+4
View File
@@ -50,6 +50,10 @@ pub(super) fn runtime_complete(root: &Path, asset: RuntimeAsset) -> bool {
paths.executable.is_file() && paths.ready_marker.is_file()
}
pub(super) fn runtime_executable(root: &Path, asset: RuntimeAsset) -> PathBuf {
RuntimePaths::new(root, asset).executable
}
async fn download_and_install(
store: &Store,
asset: RuntimeAsset,
+175
View File
@@ -0,0 +1,175 @@
//! Defines and validates the static filesystem plugin manifest.
use std::{collections::HashSet, path::Path};
use regex::Regex;
use serde::Deserialize;
use crate::{Error, Result};
pub const PLUGIN_API_VERSION: u32 = 1;
#[derive(Clone, Debug, Deserialize)]
#[serde(rename_all = "camelCase", deny_unknown_fields)]
pub struct PluginManifest {
pub api_version: u32,
pub id: String,
pub name: String,
#[serde(default)]
pub author: Option<String>,
pub icon: String,
pub entry: String,
#[serde(default)]
pub permissions: PluginPermissions,
}
#[derive(Clone, Debug, Default, Deserialize)]
#[serde(rename_all = "camelCase", deny_unknown_fields)]
pub struct PluginPermissions {
#[serde(default)]
pub network: Vec<String>,
}
impl PluginManifest {
pub fn validate(&self, directory: &Path) -> Result<()> {
if self.api_version != PLUGIN_API_VERSION {
return Err(Error::Config(format!(
"plugin '{}' uses unsupported API version {}",
self.id, self.api_version
)));
}
validate_id(&self.id, "plugin id")?;
required(&self.name, "plugin name")?;
validate_entry_path(directory, &self.entry)?;
validate_asset_path(directory, &self.icon)?;
let mut hosts = HashSet::new();
for host in &self.permissions.network {
validate_network_host(host)?;
if !hosts.insert(host.to_ascii_lowercase()) {
return Err(Error::Config(format!(
"plugin '{}' contains duplicate network host '{host}'",
self.id
)));
}
}
Ok(())
}
}
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());
if expression.is_match(value) {
Ok(())
} else {
Err(Error::Config(format!("invalid {label}: {value}")))
}
}
fn validate_network_host(value: &str) -> Result<()> {
if value.is_empty()
|| value.contains('/')
|| value.contains(':')
|| value.starts_with('.')
|| value.ends_with('.')
{
return Err(Error::Config(format!(
"invalid plugin network host: {value}"
)));
}
let parsed = url::Url::parse(&format!("https://{value}")).map_err(|error| {
Error::Config(format!("invalid plugin network host '{value}': {error}"))
})?;
if parsed.host_str() != Some(value) {
return Err(Error::Config(format!(
"invalid plugin network host: {value}"
)));
}
Ok(())
}
fn required<'a>(value: &'a str, label: &str) -> Result<&'a str> {
let value = value.trim();
if value.is_empty() {
Err(Error::Config(format!("{label} is required")))
} else {
Ok(value)
}
}
fn validate_entry_path(directory: &Path, value: &str) -> Result<()> {
let path = Path::new(value);
if !is_safe_relative_path(path) {
return Err(Error::Config(format!("invalid plugin entry path: {value}")));
}
let extension = path
.extension()
.and_then(|value| value.to_str())
.unwrap_or_default()
.to_ascii_lowercase();
if !matches!(extension.as_str(), "js" | "mjs" | "ts" | "mts") {
return Err(Error::Config(format!(
"unsupported plugin entry format: {value}"
)));
}
let entry = directory.join(path);
if !entry.is_file() {
return Err(Error::Config(format!(
"plugin entry does not exist: {value}"
)));
}
let root = directory.canonicalize()?;
let entry = entry.canonicalize()?;
if !entry.starts_with(root) {
return Err(Error::Config(format!(
"plugin entry escapes its directory: {value}"
)));
}
Ok(())
}
fn is_safe_relative_path(path: &Path) -> bool {
!path.is_absolute()
&& !path.components().any(|component| {
matches!(
component,
std::path::Component::ParentDir
| std::path::Component::RootDir
| std::path::Component::Prefix(_)
)
})
}
fn validate_asset_path(directory: &Path, value: &str) -> Result<()> {
let path = Path::new(value);
if !is_safe_relative_path(path) {
return Err(Error::Config(format!("invalid plugin asset path: {value}")));
}
let extension = path
.extension()
.and_then(|value| value.to_str())
.unwrap_or_default()
.to_ascii_lowercase();
if !matches!(extension.as_str(), "svg" | "png" | "webp") {
return Err(Error::Config(format!(
"unsupported plugin icon format: {value}"
)));
}
if !directory.join(path).is_file() {
return Err(Error::Config(format!(
"plugin icon does not exist: {value}"
)));
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn rejects_urls_in_network_host_allowlist() {
assert!(validate_network_host("https://example.com").is_err());
assert!(validate_network_host("example.com:443").is_err());
assert!(validate_network_host("example.com").is_ok());
}
}
+18 -1
View File
@@ -1,6 +1,23 @@
//! Owns plugin runtime installation and lifecycle infrastructure.
//! Owns filesystem plugin discovery, sandboxed workers, and plugin providers.
mod asset;
mod builtin;
mod catalog;
mod data;
mod definition;
mod descriptor;
mod installation;
mod manifest;
mod protocol;
mod registry;
mod runtime;
mod state;
mod wire;
mod worker;
pub use descriptor::{
parse_model_id, PluginDescriptor, PluginModelDescriptor, PluginProviderDescriptor,
PluginResourceDescriptor, PluginResourceView, ADAPTER_ID_PREFIX,
};
pub use registry::{ImportResponse, OAuthBeginResponse, OAuthPollResponse, PluginRegistry};
pub use runtime::{PluginRuntime, PluginRuntimePhase, PluginRuntimeState, PluginRuntimeStatus};
pub(crate) use wire::llm_request as plugin_llm_request;
+92
View File
@@ -0,0 +1,92 @@
//! Defines newline-delimited messages exchanged with a plugin worker.
use serde::{Deserialize, Serialize};
#[derive(Debug, Serialize)]
#[serde(tag = "type", rename_all = "snake_case")]
pub enum HostMessage<'a> {
Request {
id: &'a str,
method: &'a str,
params: &'a serde_json::Value,
},
Cancel {
id: &'a str,
},
HostResult {
id: &'a str,
result: &'a serde_json::Value,
},
HostError {
id: &'a str,
error: &'a str,
},
}
#[derive(Debug, Deserialize)]
#[serde(tag = "type", rename_all = "snake_case")]
pub enum WorkerMessage {
Result {
id: String,
#[serde(default)]
result: serde_json::Value,
#[serde(default)]
error: Option<String>,
},
/// 流式请求(provider.invoke)在最终 Result 之前发出的模型事件。
Event {
id: String,
event: serde_json::Value,
},
HostCall {
id: String,
#[serde(rename = "requestId")]
request_id: String,
method: String,
#[serde(default)]
params: serde_json::Value,
},
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn parses_multiplexed_host_call_and_events() {
let message: WorkerMessage = serde_json::from_value(serde_json::json!({
"type": "host_call",
"id": "host-2",
"requestId": "request-1",
"method": "network.fetch",
"params": { "url": "https://example.com" }
}))
.unwrap();
match message {
WorkerMessage::HostCall {
id,
request_id,
method,
..
} => {
assert_eq!(id, "host-2");
assert_eq!(request_id, "request-1");
assert_eq!(method, "network.fetch");
}
_ => panic!("expected host call"),
}
let message: WorkerMessage = serde_json::from_value(serde_json::json!({
"type": "event",
"id": "request-1",
"event": { "type": "text-delta", "text": "hi" }
}))
.unwrap();
match message {
WorkerMessage::Event { id, event } => {
assert_eq!(id, "request-1");
assert_eq!(event["type"], "text-delta");
}
_ => panic!("expected event"),
}
}
}
+956
View File
@@ -0,0 +1,956 @@
//! Orchestrates plugin capabilities: resources, model catalogs, and invocation.
use std::{collections::HashMap, path::Path, sync::Arc};
use async_stream::try_stream;
use serde::Serialize;
use tokio::sync::{Mutex, RwLock};
use tokio_util::sync::CancellationToken;
use super::{
catalog::{PluginCatalog, PluginEntry},
data::PluginDataStore,
descriptor::{
parse_model_id, PluginDescriptor, PluginModelDescriptor, PluginProviderDescriptor,
PluginResourceDescriptor, PluginResourceView, ProviderDefinition, ResourceDefinition,
ResourcePresentation, OAUTH2_ADD_METHOD,
},
runtime::PluginRuntime,
state::{now_ms, PluginStateStore, ResourceDraft, ResourcePatch, ResourceRecord, StoredModel},
wire,
worker::{PluginWorker, WorkerStreamItem},
};
use crate::{
model::ModelInvocation, provider::ModelEvent, provider::ProviderStream, store::Store, Error,
Result,
};
const OAUTH_SLOW_DOWN_STEP_MS: i64 = 5_000;
const MAX_IMPORT_DRAFTS: usize = 256;
#[derive(Clone)]
pub struct PluginRegistry {
inner: Arc<RegistryInner>,
}
struct RegistryInner {
store: Store,
runtime: PluginRuntime,
catalog: PluginCatalog,
state: PluginStateStore,
entries: RwLock<Option<Vec<PluginEntry>>>,
workers: Mutex<HashMap<String, Arc<PluginWorker>>>,
oauth_sessions: Mutex<HashMap<String, OAuthSession>>,
}
struct OAuthSession {
plugin_id: String,
resource_type: String,
method_id: String,
session: serde_json::Value,
expires_at_ms: i64,
poll_interval_ms: i64,
next_poll_at_ms: i64,
}
#[derive(Clone, Debug, Serialize)]
#[serde(rename_all = "camelCase")]
pub struct OAuthBeginResponse {
pub session_id: String,
pub user_code: String,
pub verification_url: String,
pub verification_url_complete: Option<String>,
pub expires_at_ms: i64,
pub poll_interval_ms: i64,
}
#[derive(Clone, Debug, Serialize)]
#[serde(rename_all = "camelCase", tag = "status")]
pub enum OAuthPollResponse {
#[serde(rename_all = "camelCase")]
Pending { poll_interval_ms: i64 },
#[serde(rename_all = "camelCase")]
Completed {
added: usize,
updated: usize,
model_sync_error: Option<String>,
},
#[serde(rename_all = "camelCase")]
Denied { message: Option<String> },
#[serde(rename_all = "camelCase")]
Failed { message: String },
}
#[derive(Clone, Debug, Serialize)]
#[serde(rename_all = "camelCase")]
pub struct ImportResponse {
pub added: usize,
pub updated: usize,
pub warnings: Vec<String>,
pub model_sync_error: Option<String>,
}
/// 路由分支在建立 Recorder 时需要的插件模型元数据。
#[derive(Clone, Debug)]
pub struct PluginInvocationPlan {
pub model: PluginModelDescriptor,
pub request_url: String,
}
impl PluginRegistry {
pub fn managed(store: Store, runtime: PluginRuntime) -> Result<Self> {
let data = PluginDataStore::managed()?;
Ok(Self {
inner: Arc::new(RegistryInner {
store,
runtime,
catalog: PluginCatalog::managed()?,
state: PluginStateStore::new(data),
entries: RwLock::new(None),
workers: Mutex::new(HashMap::new()),
oauth_sessions: Mutex::new(HashMap::new()),
}),
})
}
pub async fn plugins(&self) -> Vec<PluginDescriptor> {
let Some(executable) = self.inner.runtime.executable() else {
return self
.inner
.catalog
.manifests()
.into_iter()
.map(|(manifest, icon)| PluginDescriptor {
id: manifest.id,
name: manifest.name,
author: manifest.author,
icon,
providers: Vec::new(),
resources: Vec::new(),
})
.collect();
};
let mut plugins = Vec::new();
for entry in self.entries(&executable).await {
plugins.push(self.descriptor(&entry, &executable).await);
}
plugins
}
/// 已满足调用条件的全部插件模型;每个模型独立进入 Cursor 目录。
pub async fn configured_models(&self) -> Vec<PluginModelDescriptor> {
let Some(executable) = self.inner.runtime.executable() else {
return Vec::new();
};
let mut models = Vec::new();
for entry in self.entries(&executable).await {
for provider in &entry.definition.providers {
if !self.provider_configured(&entry, provider).await {
continue;
}
let stored = self
.inner
.state
.models(&entry.manifest.id, &provider.id)
.await
.unwrap_or_default();
models.extend(stored.iter().map(|model| {
PluginModelDescriptor::new(
&entry.manifest.id,
&entry.manifest.name,
&entry.icon,
provider,
model,
)
}));
}
}
models
}
pub async fn model_descriptor(&self, model_id: &str) -> Result<PluginModelDescriptor> {
let (plugin_id, provider_id, upstream_id) = parse_model_id(model_id)
.ok_or_else(|| Error::Provider(format!("invalid plugin model ID: {model_id}")))?;
let executable = self.executable()?;
let entry = self.find_entry(&executable, plugin_id).await?;
let provider = find_provider(&entry, provider_id)?;
let stored = self
.inner
.state
.models(plugin_id, provider_id)
.await?
.into_iter()
.find(|model| model.id == upstream_id)
.ok_or_else(|| Error::RunNotFound(format!("plugin model {model_id}")))?;
Ok(PluginModelDescriptor::new(
plugin_id,
&entry.manifest.name,
&entry.icon,
provider,
&stored,
))
}
pub async fn plan_model(&self, model_id: &str) -> Result<PluginInvocationPlan> {
let model = self.model_descriptor(model_id).await?;
let request_url = format!("plugin://{}/{}", model.plugin_id, model.provider_id);
Ok(PluginInvocationPlan { model, request_url })
}
/// 插件模型的统一 Provider 流:选首个可用资源,经 Worker 执行,
/// 事件与内置 Provider 走同一管道。未来的负载均衡在这里换资源重试。
pub fn stream_model(
&self,
invocation: ModelInvocation,
cancellation: CancellationToken,
) -> ProviderStream {
let registry = self.clone();
Box::pin(try_stream! {
let model_id = invocation.request.model.model_id.clone();
let (plugin_id, provider_id, upstream_id) = parse_model_id(&model_id)
.map(|(plugin, provider, model)| (plugin.to_owned(), provider.to_owned(), model.to_owned()))
.ok_or_else(|| Error::Provider(format!("invalid plugin model ID: {model_id}")))?;
let executable = registry.executable()?;
let entry = registry.find_entry(&executable, &plugin_id).await?;
let provider = find_provider(&entry, &provider_id)?.clone();
let stored = registry.inner.state.models(&plugin_id, &provider_id).await?
.into_iter()
.find(|model| model.id == upstream_id)
.ok_or_else(|| Error::RunNotFound(format!("plugin model {model_id}")))?;
let resource = match &provider.resource_type {
Some(resource_type) => Some((
resource_type.clone(),
registry.select_resource(&plugin_id, resource_type).await?,
)),
None => None,
};
let request = wire::llm_request(&invocation)?;
let params = serde_json::json!({
"providerId": provider_id,
"model": stored.snapshot(),
"resource": resource.as_ref().map(|(resource_type, record)| record.snapshot(resource_type)),
"request": request,
});
let worker = registry.worker(&entry, &executable).await;
let mut items = worker.invoke_streaming("provider.invoke", params, cancellation.clone()).await?;
yield ModelEvent::Start { model_call_id: invocation.call_id.clone() };
while let Some(item) = items.recv().await {
match item {
WorkerStreamItem::Event(event) => {
yield wire::model_event(&event)?;
}
WorkerStreamItem::Result(result) => {
let value = result?;
let status = value.get("status").and_then(serde_json::Value::as_str).unwrap_or_default();
let patch = value.get("patch")
.filter(|patch| !patch.is_null())
.map(|patch| serde_json::from_value::<ResourcePatch>(patch.clone()))
.transpose()?;
if let (Some(patch), Some((resource_type, record))) = (patch, resource.as_ref()) {
if let Err(error) = registry.inner.state
.apply_patch(&plugin_id, resource_type, &record.id, patch).await
{
tracing::warn!(plugin = %plugin_id, %error, "failed to apply plugin resource patch");
}
}
match status {
"completed" => return,
"resource-error" | "request-error" => {
let message = value.get("message")
.and_then(serde_json::Value::as_str)
.unwrap_or("plugin provider call failed");
Err(Error::Provider(message.to_owned()))?;
}
status => {
Err(Error::Protocol(format!("unknown plugin provider result: {status}")))?;
}
}
}
}
}
Err(Error::Provider(format!("plugin '{plugin_id}' worker stopped mid-stream")))?;
})
}
pub async fn oauth_begin(
&self,
plugin_id: &str,
resource_type: &str,
method_id: &str,
) -> Result<OAuthBeginResponse> {
let executable = self.executable()?;
let entry = self.find_entry(&executable, plugin_id).await?;
let resource = find_resource(&entry, resource_type)?;
let method = resource
.add
.iter()
.find(|method| method.id == method_id && method.method_type == OAUTH2_ADD_METHOD)
.ok_or_else(|| {
Error::Config(format!(
"plugin '{plugin_id}' does not define OAuth method '{method_id}'"
))
})?;
let value = self
.worker(&entry, &executable)
.await
.invoke(
"oauth.begin",
serde_json::json!({ "resourceType": resource_type, "methodId": method.id }),
CancellationToken::new(),
)
.await?;
let begin: OAuth2Begin = serde_json::from_value(value)?;
let session_id = uuid::Uuid::new_v4().to_string();
self.inner.oauth_sessions.lock().await.insert(
session_id.clone(),
OAuthSession {
plugin_id: plugin_id.to_owned(),
resource_type: resource_type.to_owned(),
method_id: method_id.to_owned(),
session: begin.session,
expires_at_ms: begin.expires_at_ms,
poll_interval_ms: begin.poll_interval_ms.max(1_000),
next_poll_at_ms: now_ms() + begin.poll_interval_ms.max(1_000),
},
);
Ok(OAuthBeginResponse {
session_id,
user_code: begin.user_code,
verification_url: begin.verification_url,
verification_url_complete: begin.verification_url_complete,
expires_at_ms: begin.expires_at_ms,
poll_interval_ms: begin.poll_interval_ms.max(1_000),
})
}
pub async fn oauth_poll(&self, session_id: &str) -> Result<OAuthPollResponse> {
let now = now_ms();
let (plugin_id, resource_type, method_id, session, poll_interval_ms) = {
let mut sessions = self.inner.oauth_sessions.lock().await;
let Some(state) = sessions.get_mut(session_id) else {
return Ok(OAuthPollResponse::Failed {
message: "authorization session no longer exists".into(),
});
};
if now >= state.expires_at_ms {
sessions.remove(session_id);
return Ok(OAuthPollResponse::Failed {
message: "device authorization expired".into(),
});
}
if now < state.next_poll_at_ms {
return Ok(OAuthPollResponse::Pending {
poll_interval_ms: state.poll_interval_ms,
});
}
state.next_poll_at_ms = now + state.poll_interval_ms;
(
state.plugin_id.clone(),
state.resource_type.clone(),
state.method_id.clone(),
state.session.clone(),
state.poll_interval_ms,
)
};
let executable = self.executable()?;
let entry = self.find_entry(&executable, &plugin_id).await?;
let value = self
.worker(&entry, &executable)
.await
.invoke(
"oauth.poll",
serde_json::json!({
"resourceType": resource_type,
"methodId": method_id,
"session": session,
}),
CancellationToken::new(),
)
.await?;
let poll: OAuth2Poll = serde_json::from_value(value)?;
match poll {
OAuth2Poll::Pending { session } => {
self.update_session(session_id, session, None).await;
Ok(OAuthPollResponse::Pending { poll_interval_ms })
}
OAuth2Poll::SlowDown { session } => {
let interval = poll_interval_ms + OAUTH_SLOW_DOWN_STEP_MS;
self.update_session(session_id, session, Some(interval))
.await;
Ok(OAuthPollResponse::Pending {
poll_interval_ms: interval,
})
}
OAuth2Poll::Completed { resources } => {
self.inner.oauth_sessions.lock().await.remove(session_id);
let outcome = self
.inner
.state
.upsert_resources(&plugin_id, &resource_type, resources)
.await?;
let model_sync_error = self
.sync_provider_models_for_resource(&entry, &executable, &resource_type)
.await;
Ok(OAuthPollResponse::Completed {
added: outcome.added,
updated: outcome.updated,
model_sync_error,
})
}
OAuth2Poll::Denied { message } => {
self.inner.oauth_sessions.lock().await.remove(session_id);
Ok(OAuthPollResponse::Denied { message })
}
OAuth2Poll::Failed { message } => {
self.inner.oauth_sessions.lock().await.remove(session_id);
Ok(OAuthPollResponse::Failed { message })
}
}
}
pub async fn import_resources(
&self,
plugin_id: &str,
resource_type: &str,
files: serde_json::Value,
) -> Result<ImportResponse> {
let executable = self.executable()?;
let entry = self.find_entry(&executable, plugin_id).await?;
let resource = find_resource(&entry, resource_type)?;
if resource.import.is_none() {
return Err(Error::Config(format!(
"plugin '{plugin_id}' resource '{resource_type}' does not support import"
)));
}
let value = self
.worker(&entry, &executable)
.await
.invoke(
"import.parse",
serde_json::json!({ "resourceType": resource_type, "files": files }),
CancellationToken::new(),
)
.await?;
let parsed: ImportParseResult = serde_json::from_value(value)?;
if parsed.resources.is_empty() {
return Err(Error::Config(
parsed
.warnings
.first()
.cloned()
.unwrap_or_else(|| "import produced no resources".into()),
));
}
if parsed.resources.len() > MAX_IMPORT_DRAFTS {
return Err(Error::Config(format!(
"import produced more than {MAX_IMPORT_DRAFTS} resources"
)));
}
let outcome = self
.inner
.state
.upsert_resources(plugin_id, resource_type, parsed.resources)
.await?;
let model_sync_error = self
.sync_provider_models_for_resource(&entry, &executable, resource_type)
.await;
Ok(ImportResponse {
added: outcome.added,
updated: outcome.updated,
warnings: parsed.warnings,
model_sync_error,
})
}
/// 导出某资源类型的全部私有数据,供备份或迁移;格式与批量导入兼容。
pub async fn export_resources(
&self,
plugin_id: &str,
resource_type: &str,
) -> Result<serde_json::Value> {
let executable = self.executable()?;
let entry = self.find_entry(&executable, plugin_id).await?;
find_resource(&entry, resource_type)?;
let records = self.inner.state.resources(plugin_id, resource_type).await?;
Ok(serde_json::json!({
"accounts": records
.iter()
.map(|record| record.private_data.clone())
.collect::<Vec<_>>(),
}))
}
pub async fn refresh_resource(
&self,
plugin_id: &str,
resource_type: &str,
resource_id: &str,
) -> Result<()> {
let executable = self.executable()?;
let entry = self.find_entry(&executable, plugin_id).await?;
let resource = find_resource(&entry, resource_type)?;
if !resource.can_refresh {
return Err(Error::Config(format!(
"plugin '{plugin_id}' resource '{resource_type}' does not support refresh"
)));
}
let record = self
.find_record(plugin_id, resource_type, resource_id)
.await?;
let value = self
.worker(&entry, &executable)
.await
.invoke(
"resource.refresh",
serde_json::json!({
"resourceType": resource_type,
"resource": record.snapshot(resource_type),
}),
CancellationToken::new(),
)
.await?;
let patch: ResourcePatch = serde_json::from_value(value)?;
self.inner
.state
.apply_patch(plugin_id, resource_type, resource_id, patch)
.await
}
pub async fn delete_resource(
&self,
plugin_id: &str,
resource_type: &str,
resource_id: &str,
) -> Result<()> {
let executable = self.executable()?;
let entry = self.find_entry(&executable, plugin_id).await?;
let resource = find_resource(&entry, resource_type)?;
let record = self
.find_record(plugin_id, resource_type, resource_id)
.await?;
if resource.can_remove {
// 上游撤销失败不阻塞本地删除:用户必须能移除已失效的资源。
if let Err(error) = self
.worker(&entry, &executable)
.await
.invoke(
"resource.remove",
serde_json::json!({
"resourceType": resource_type,
"resource": record.snapshot(resource_type),
}),
CancellationToken::new(),
)
.await
{
tracing::warn!(plugin = %plugin_id, %error, "plugin resource remove hook failed");
}
}
self.inner
.state
.remove_resource(plugin_id, resource_type, resource_id)
.await?;
Ok(())
}
pub async fn sync_models(&self, plugin_id: &str, provider_id: &str) -> Result<usize> {
let executable = self.executable()?;
let entry = self.find_entry(&executable, plugin_id).await?;
let provider = find_provider(&entry, provider_id)?.clone();
self.sync_provider_models(&entry, &executable, &provider)
.await
}
pub async fn remove(&self, plugin_id: &str) -> Result<()> {
if let Some(worker) = self.inner.workers.lock().await.remove(plugin_id) {
worker.stop().await;
}
self.inner.state.clear(plugin_id).await
}
async fn descriptor(&self, entry: &PluginEntry, executable: &Path) -> PluginDescriptor {
let plugin_id = &entry.manifest.id;
let mut providers = Vec::new();
for provider in &entry.definition.providers {
let stored = self
.inner
.state
.models(plugin_id, &provider.id)
.await
.unwrap_or_default();
let configured = self.provider_configured(entry, provider).await;
providers.push(PluginProviderDescriptor {
id: provider.id.clone(),
plugin_id: plugin_id.clone(),
display_name: provider.display_name.clone(),
description: provider.description.clone(),
provider_type: provider.provider_type.clone(),
resource_type: provider.resource_type.clone(),
has_models: provider.has_models,
configured,
models: stored
.iter()
.map(|model| {
PluginModelDescriptor::new(
plugin_id,
&entry.manifest.name,
&entry.icon,
provider,
model,
)
})
.collect(),
});
}
let mut resources = Vec::new();
for definition in &entry.definition.resources {
let records = self
.inner
.state
.resources(plugin_id, &definition.resource_type)
.await
.unwrap_or_default();
let views = self
.present_resources(entry, executable, definition, &records)
.await;
resources.push(PluginResourceDescriptor {
resource_type: definition.resource_type.clone(),
display_name: definition.display_name.clone(),
add: definition.add.clone(),
import: definition.import.clone(),
can_refresh: definition.can_refresh,
can_remove: definition.can_remove,
resources: views,
});
}
PluginDescriptor {
id: plugin_id.clone(),
name: entry.manifest.name.clone(),
author: entry.manifest.author.clone(),
icon: entry.icon.clone(),
providers,
resources,
}
}
async fn present_resources(
&self,
entry: &PluginEntry,
executable: &Path,
definition: &ResourceDefinition,
records: &[ResourceRecord],
) -> Vec<PluginResourceView> {
if records.is_empty() {
return Vec::new();
}
let snapshots = records
.iter()
.map(|record| record.snapshot(&definition.resource_type))
.collect::<Vec<_>>();
let presented = self
.worker(entry, executable)
.await
.invoke(
"resource.present",
serde_json::json!({
"resourceType": definition.resource_type,
"resources": snapshots,
}),
CancellationToken::new(),
)
.await
.and_then(|value| {
serde_json::from_value::<Vec<ResourcePresentation>>(value).map_err(Error::from)
});
match presented {
Ok(views) if views.len() == records.len() => records
.iter()
.zip(views)
.map(|(record, view)| PluginResourceView::from_record(record, view))
.collect(),
Ok(_) | Err(_) => records
.iter()
.map(|record| {
PluginResourceView::from_record(
record,
ResourcePresentation {
display_name: record.key.clone(),
description: serde_json::Value::Null,
metrics: Vec::new(),
},
)
})
.collect(),
}
}
async fn provider_configured(
&self,
entry: &PluginEntry,
provider: &ProviderDefinition,
) -> bool {
let plugin_id = &entry.manifest.id;
if provider.has_models {
let models = self
.inner
.state
.models(plugin_id, &provider.id)
.await
.unwrap_or_default();
if models.is_empty() {
return false;
}
}
match &provider.resource_type {
Some(resource_type) => !self
.inner
.state
.resources(plugin_id, resource_type)
.await
.unwrap_or_default()
.is_empty(),
None => true,
}
}
/// 资源到位后刷新使用该资源类型的 Provider 模型目录;失败只报告不中断。
async fn sync_provider_models_for_resource(
&self,
entry: &PluginEntry,
executable: &Path,
resource_type: &str,
) -> Option<String> {
let mut errors = Vec::new();
for provider in entry.definition.providers.clone() {
if provider.resource_type.as_deref() != Some(resource_type) || !provider.has_models {
continue;
}
if let Err(error) = self
.sync_provider_models(entry, executable, &provider)
.await
{
errors.push(format!("{}: {error}", provider.id));
}
}
(!errors.is_empty()).then(|| errors.join("; "))
}
async fn sync_provider_models(
&self,
entry: &PluginEntry,
executable: &Path,
provider: &ProviderDefinition,
) -> Result<usize> {
if !provider.has_models {
return Err(Error::Config(format!(
"plugin provider '{}' does not enumerate models",
provider.id
)));
}
let plugin_id = &entry.manifest.id;
let resource = match &provider.resource_type {
Some(resource_type) => {
let record = self.select_resource(plugin_id, resource_type).await?;
Some(record.snapshot(resource_type))
}
None => None,
};
let value = self
.worker(entry, executable)
.await
.invoke(
"models.list",
serde_json::json!({ "providerId": provider.id, "resource": resource }),
CancellationToken::new(),
)
.await?;
let definitions = value
.as_array()
.ok_or_else(|| Error::Protocol("plugin models.list must return an array".into()))?;
let mut models = Vec::with_capacity(definitions.len());
let mut seen = std::collections::HashSet::new();
for definition in definitions {
let model = StoredModel::from_definition(definition)?;
if seen.insert(model.id.clone()) {
models.push(model);
}
}
if models.is_empty() {
return Err(Error::Provider(format!(
"plugin provider '{}' returned no models",
provider.id
)));
}
self.inner
.state
.replace_models(plugin_id, &provider.id, &models)
.await?;
Ok(models.len())
}
/// 第一版选择策略:按创建顺序取首个可用资源;冷却到期视为可用。
async fn select_resource(
&self,
plugin_id: &str,
resource_type: &str,
) -> Result<ResourceRecord> {
let records = self.inner.state.resources(plugin_id, resource_type).await?;
if records.is_empty() {
return Err(Error::Provider(format!(
"plugin '{plugin_id}' has no '{resource_type}' resource; add one first"
)));
}
let now = now_ms();
records
.iter()
.find(|record| record.state.is_ready(now))
.or_else(|| records.first())
.cloned()
.ok_or_else(|| Error::Provider("no plugin resource is available".into()))
}
async fn find_record(
&self,
plugin_id: &str,
resource_type: &str,
resource_id: &str,
) -> Result<ResourceRecord> {
self.inner
.state
.resources(plugin_id, resource_type)
.await?
.into_iter()
.find(|record| record.id == resource_id)
.ok_or_else(|| Error::RunNotFound(format!("plugin resource {resource_id}")))
}
async fn update_session(
&self,
session_id: &str,
session: Option<serde_json::Value>,
poll_interval_ms: Option<i64>,
) {
let mut sessions = self.inner.oauth_sessions.lock().await;
if let Some(state) = sessions.get_mut(session_id) {
if let Some(session) = session {
state.session = session;
}
if let Some(interval) = poll_interval_ms {
state.poll_interval_ms = interval;
}
}
}
fn executable(&self) -> Result<std::path::PathBuf> {
self.inner
.runtime
.executable()
.ok_or_else(|| Error::Config("plugin runtime is not ready".into()))
}
async fn entries(&self, executable: &Path) -> Vec<PluginEntry> {
if let Some(entries) = self.inner.entries.read().await.as_ref() {
return entries.clone();
}
let loaded = self.inner.catalog.entries(executable).await;
*self.inner.entries.write().await = Some(loaded.clone());
loaded
}
async fn find_entry(&self, executable: &Path, plugin_id: &str) -> Result<PluginEntry> {
self.entries(executable)
.await
.into_iter()
.find(|entry| entry.manifest.id == plugin_id)
.ok_or_else(|| Error::RunNotFound(format!("plugin {plugin_id}")))
}
async fn worker(&self, entry: &PluginEntry, executable: &Path) -> Arc<PluginWorker> {
let mut workers = self.inner.workers.lock().await;
workers
.entry(entry.manifest.id.clone())
.or_insert_with(|| {
Arc::new(PluginWorker::new(
entry,
executable.to_path_buf(),
self.inner.catalog.loader().clone(),
self.inner.store.clone(),
))
})
.clone()
}
}
#[derive(Debug, serde::Deserialize)]
#[serde(rename_all = "camelCase", deny_unknown_fields)]
struct OAuth2Begin {
session: serde_json::Value,
user_code: String,
verification_url: String,
#[serde(default)]
verification_url_complete: Option<String>,
expires_at_ms: i64,
poll_interval_ms: i64,
}
#[derive(Debug, serde::Deserialize)]
#[serde(rename_all = "kebab-case", tag = "status")]
enum OAuth2Poll {
#[serde(rename_all = "camelCase")]
Pending {
#[serde(default)]
session: Option<serde_json::Value>,
},
#[serde(rename_all = "camelCase")]
SlowDown {
#[serde(default)]
session: Option<serde_json::Value>,
},
#[serde(rename_all = "camelCase")]
Completed { resources: Vec<ResourceDraft> },
#[serde(rename_all = "camelCase")]
Denied {
#[serde(default)]
message: Option<String>,
},
#[serde(rename_all = "camelCase")]
Failed { message: String },
}
#[derive(Debug, serde::Deserialize)]
#[serde(rename_all = "camelCase", deny_unknown_fields)]
struct ImportParseResult {
resources: Vec<ResourceDraft>,
#[serde(default)]
warnings: Vec<String>,
}
fn find_provider<'a>(entry: &'a PluginEntry, provider_id: &str) -> Result<&'a ProviderDefinition> {
entry
.definition
.providers
.iter()
.find(|provider| provider.id == provider_id)
.ok_or_else(|| {
Error::RunNotFound(format!(
"plugin '{}' provider {provider_id}",
entry.manifest.id
))
})
}
fn find_resource<'a>(
entry: &'a PluginEntry,
resource_type: &str,
) -> Result<&'a ResourceDefinition> {
entry
.definition
.resources
.iter()
.find(|resource| resource.resource_type == resource_type)
.ok_or_else(|| {
Error::RunNotFound(format!(
"plugin '{}' resource type {resource_type}",
entry.manifest.id
))
})
}
+8
View File
@@ -142,6 +142,14 @@ impl PluginRuntime {
status.clone()
}
pub fn executable(&self) -> Option<PathBuf> {
let asset = self.inner.asset?;
if self.status().state != PluginRuntimeState::Ready {
return None;
}
Some(installation::runtime_executable(&self.inner.root, asset))
}
pub fn initialize(&self, store: Store) -> PluginRuntimeStatus {
let Some(asset) = self.inner.asset else {
return self.status();
+5
View File
@@ -0,0 +1,5 @@
import { __descriptor, __getRegisteredPlugin } from "cursor-byok:plugin";
if (Deno.args.length !== 1) throw new Error("plugin entry URL is required");
await import(Deno.args[0]);
console.log("CURSOR_BYOK_PLUGIN_DEFINITION:" + JSON.stringify(__descriptor(__getRegisteredPlugin())));
+5
View File
@@ -0,0 +1,5 @@
{
"fmt": {
"lineWidth": 200
}
}
+9
View File
@@ -0,0 +1,9 @@
{
"imports": {
"cursor-byok:plugin": "./plugin.ts",
"cursor-byok:provider": "./provider.ts",
"cursor-byok:model": "./model.ts",
"cursor-byok:resource": "./resource.ts",
"cursor-byok:protocol/openai-responses": "./protocol/openai_responses.ts"
}
}
+31
View File
@@ -0,0 +1,31 @@
import type { JsonValue, PluginContext } from "./plugin.ts";
import type { ResourceSnapshot } from "./resource.ts";
export type ModelCapabilities = {
thinking?: boolean;
images?: boolean;
};
export type ModelDefinition = {
id: string;
displayName: string;
description?: string;
contextWindowTokens?: number;
maxOutputTokens?: number;
capabilities?: ModelCapabilities;
/** 之后的调用原样传回;永远不会展示给用户。 */
privateData?: JsonValue;
};
/** 宿主目录中持久化的一条模型。 */
export type ModelSnapshot = ModelDefinition;
export type ModelListInput = {
/** 模型发现需要认证时为首个可用资源,否则为 null。 */
resource: ResourceSnapshot | null;
};
export type ModelSupport = {
/** 列举成功后,宿主用返回值整体替换该 Provider 的模型目录。 */
list(input: ModelListInput, context: PluginContext): Promise<ModelDefinition[]>;
};
+100
View File
@@ -0,0 +1,100 @@
import type { ProviderSupport } from "./provider.ts";
import type { ResourceSupport } from "./resource.ts";
export type JsonPrimitive = string | number | boolean | null;
export type JsonValue = JsonPrimitive | JsonValue[] | { [key: string]: JsonValue };
/**
* 可本地化文本:纯字符串,或 locale → 文本 的映射
* (如 { "zh-CN": "账号", "en-US": "Accounts" })。
* 宿主原样透传,由界面按当前语言解析;模型名等来自上游的数据保持纯字符串。
*/
export type LocalizedText = string | { [locale: string]: string };
export type NetworkRequestInit = {
method?: string;
headers?: Record<string, string>;
body?: string;
};
export type NetworkResponse = {
status: number;
headers: Record<string, string>;
body: string;
};
/** 流式响应体,按行随到随交付(用于 SSE)。 */
export type NetworkEventStream = {
status: number;
headers: Record<string, string>;
lines: AsyncIterable<string>;
};
/**
* 每次能力调用收到的宿主服务。网络请求仅限 plugin.json 声明的 HTTPS 主机;
* 宿主取消本次调用时通过 `signal` 中止。
*/
export type PluginContext = {
network: {
fetch(url: string, init?: NetworkRequestInit): Promise<NetworkResponse>;
stream(url: string, init?: NetworkRequestInit): Promise<NetworkEventStream>;
};
signal: AbortSignal;
};
/**
* Provider 插件定义:一组能力实现的集合。插件不持有任何持久状态——
* 资源与模型目录由宿主存储,每次调用所需的数据都通过参数传入。
*/
export type ProviderPluginDefinition = {
providers: ProviderSupport[];
resources?: ResourceSupport[];
};
let registered: ProviderPluginDefinition | undefined;
/** 注册 Provider 插件;每个插件入口只能调用一次。 */
export function defineProviderPlugin(definition: ProviderPluginDefinition): ProviderPluginDefinition {
if (registered) throw new Error("defineProviderPlugin can only be called once");
registered = definition;
return definition;
}
export function __getRegisteredPlugin(): ProviderPluginDefinition {
if (!registered) throw new Error("plugin entry must call defineProviderPlugin");
return registered;
}
/** 可序列化的能力摘要,宿主收集它时不调用任何能力方法。 */
export function __descriptor(definition: ProviderPluginDefinition) {
return {
providers: definition.providers.map((provider) => ({
id: provider.id,
displayName: provider.displayName,
description: provider.description ?? null,
providerType: provider.providerType,
resourceType: provider.resourceType ?? null,
hasModels: provider.models !== undefined,
})),
resources: (definition.resources ?? []).map((resource) => ({
type: resource.type,
displayName: resource.displayName,
add: (resource.add ?? []).map((method) => ({
type: method.type,
id: method.id,
displayName: method.displayName,
description: method.description ?? null,
})),
import: resource.import
? {
displayName: resource.import.displayName,
description: resource.import.description ?? null,
accept: resource.import.accept,
multiple: resource.import.multiple ?? false,
}
: null,
canRefresh: resource.refresh !== undefined,
canRemove: resource.remove !== undefined,
})),
};
}
@@ -0,0 +1,425 @@
import type { JsonValue, PluginContext } from "../plugin.ts";
import type { LlmContentPart, LlmRequest, ModelEvent, ProviderOutput } from "../provider.ts";
/** 本协议产生的回放状态种类;与宿主内置 Responses Provider 一致,可互相回放。 */
export const REPLAY_KIND = "openai_responses";
/** 上游返回非 2xx 时抛出,携带完整响应体供调用方分类。 */
export class HttpError extends Error {
constructor(readonly status: number, readonly body: string) {
super(`HTTP ${status}: ${body}`);
}
}
export type OpenAiResponsesCall = {
url: string;
model: string;
request: LlmRequest;
headers?: Record<string, string>;
/** 最后合并进请求体,如 { store: false }。 */
extraBody?: Record<string, JsonValue>;
};
function record(value: unknown): Record<string, unknown> | null {
return value !== null && typeof value === "object" && !Array.isArray(value) ? value as Record<string, unknown> : null;
}
function text(value: unknown): string | null {
return typeof value === "string" ? value : null;
}
function count(value: unknown): number | null {
return typeof value === "number" && Number.isFinite(value) ? value : null;
}
function contentParts(parts: LlmContentPart[], textType: "input_text" | "output_text"): JsonValue[] {
const content: JsonValue[] = [];
for (const part of parts) {
if (part.type === "text") {
if (part.text) content.push({ type: textType, text: part.text });
} else {
content.push({
type: "input_image",
detail: "auto",
image_url: `data:${part.mediaType};base64,${part.dataBase64}`,
});
}
}
return content;
}
function replayItems(value: JsonValue): JsonValue[] {
const items = record(value)?.items;
if (!Array.isArray(items)) {
throw new Error("OpenAI Responses replay state is missing items");
}
return items;
}
export function buildResponsesBody(call: OpenAiResponsesCall): Record<string, JsonValue> {
const input: JsonValue[] = [];
for (const message of call.request.messages) {
if (message.role === "assistant") {
if (message.replayState?.providerKind === REPLAY_KIND) {
input.push(...replayItems(message.replayState.value));
}
if (message.text) {
input.push({
type: "message",
role: "assistant",
content: [{ type: "output_text", text: message.text }],
});
}
for (const toolCall of message.toolCalls) {
input.push({
type: "function_call",
call_id: toolCall.callId,
name: toolCall.name,
arguments: JSON.stringify(toolCall.arguments),
});
}
} else if (message.role === "tool") {
input.push({
type: "function_call_output",
call_id: message.callId,
output: message.parts.length === 0 ? message.content : contentParts(message.parts, "input_text"),
});
} else {
const content = contentParts(message.content, "input_text");
if (content.length > 0) input.push({ type: "message", role: message.role, content });
}
}
const body: Record<string, JsonValue> = {
model: call.model,
input,
stream: true,
instructions: call.request.instructions,
include: ["reasoning.encrypted_content"],
};
if (call.request.tools.length > 0) {
body.tools = call.request.tools.map((tool) => ({
type: "function",
name: tool.name,
description: tool.description,
parameters: tool.parameters,
strict: false,
}));
}
if (call.request.maxOutputTokens !== null) body.max_output_tokens = call.request.maxOutputTokens;
const reasoning = call.request.reasoning;
if (reasoning.enabled || reasoning.effort !== null) {
body.reasoning = {
summary: "auto",
...(reasoning.effort !== null ? { effort: reasoning.effort } : {}),
};
}
// OpenAI 的规范 tier 值是 priority;"fast" 只是客户端别名,上游不接受。
if (call.request.latency === "fast") body.service_tier = "priority";
// 会话级缓存键把请求钉到同一缓存分片,前缀缓存才能稳定命中。
if (call.request.cacheKey !== null) body.prompt_cache_key = call.request.cacheKey;
return { ...body, ...call.extraBody };
}
type ToolState = {
callId: string | null;
name: string | null;
arguments: string;
emitted: number;
started: boolean;
ended: boolean;
};
type ToolArguments =
| { kind: "none" }
| { kind: "delta"; delta: string }
| { kind: "snapshot"; snapshot: string };
function updateTool(
index: number,
item: Record<string, unknown> | null,
args: ToolArguments,
done: boolean,
tools: Map<number, ToolState>,
): ModelEvent[] {
let tool = tools.get(index);
if (!tool) {
tool = { callId: null, name: null, arguments: "", emitted: 0, started: false, ended: false };
tools.set(index, tool);
}
tool.callId ??= text(item?.call_id);
tool.name ??= text(item?.name);
if (args.kind === "delta") {
tool.arguments += args.delta;
} else if (args.kind === "snapshot" && args.snapshot !== tool.arguments) {
if (!args.snapshot.startsWith(tool.arguments)) {
throw new Error("OpenAI Responses final tool arguments do not match streamed arguments");
}
tool.arguments += args.snapshot.slice(tool.arguments.length);
}
const events: ModelEvent[] = [];
if (!tool.started && tool.callId !== null && tool.name !== null) {
tool.started = true;
events.push({ type: "tool-call-start", index, callId: tool.callId, name: tool.name });
}
if (tool.started && tool.emitted < tool.arguments.length) {
events.push({ type: "tool-call-arguments-delta", index, delta: tool.arguments.slice(tool.emitted) });
tool.emitted = tool.arguments.length;
}
if (done && !tool.ended) {
if (!tool.started) {
throw new Error("OpenAI Responses function call is missing call_id or name");
}
tool.ended = true;
events.push({ type: "tool-call-end", index });
}
return events;
}
function itemText(item: Record<string, unknown>): string | null {
const content = item.content;
if (!Array.isArray(content)) return null;
return content
.map((part) => record(part))
.filter((part) => part?.type === "output_text")
.map((part) => text(part?.text) ?? "")
.join("");
}
function requiredIndex(value: Record<string, unknown>): number {
const index = count(value.output_index);
if (index === null) throw new Error("OpenAI Responses event is missing output_index");
return index;
}
function usageEvent(value: unknown): ModelEvent {
const usage = record(value) ?? {};
return {
type: "usage",
usage: {
inputTokens: count(usage.input_tokens),
outputTokens: count(usage.output_tokens),
totalTokens: count(usage.total_tokens),
cacheReadTokens: count(record(usage.input_tokens_details)?.cached_tokens),
cacheWriteTokens: null,
reasoningTokens: count(record(usage.output_tokens_details)?.reasoning_tokens),
},
};
}
async function readBody(lines: AsyncIterable<string>): Promise<string> {
const collected: string[] = [];
for await (const line of lines) collected.push(line);
return collected.join("\n");
}
/**
* 执行一次 Responses API 流式调用,发出与宿主统一事件集一致的标准化事件,
* 包括文本/思考边界、工具参数增量与加密推理回放状态。非 2xx 响应抛出
* `HttpError`,流内失败抛出 `Error`,由调用方分类额度与授权问题。
*/
export async function streamOpenAiResponses(
call: OpenAiResponsesCall,
output: ProviderOutput,
context: PluginContext,
): Promise<void> {
const response = await context.network.stream(call.url, {
method: "POST",
headers: {
accept: "text/event-stream",
"content-type": "application/json",
...call.headers,
},
body: JSON.stringify(buildResponsesBody(call)),
});
if (response.status < 200 || response.status >= 300) {
throw new HttpError(response.status, await readBody(response.lines));
}
let textOpen = false;
let streamedText = "";
let thinkingOpen = false;
const tools = new Map<number, ToolState>();
const reasoningItems: JsonValue[] = [];
let sawTool = false;
let sawCompletedItem = false;
let terminal = false;
const closeThinking = () => {
if (thinkingOpen) {
thinkingOpen = false;
output.emit({ type: "thinking-end" });
}
};
const closeText = () => {
if (textOpen) {
textOpen = false;
output.emit({ type: "text-end" });
}
};
// 流式增量可能落后于最终文本;补发缺失的后缀。
const reconcileText = (finalText: string) => {
if (finalText.startsWith(streamedText) && finalText.length > streamedText.length) {
if (!textOpen) {
textOpen = true;
output.emit({ type: "text-start" });
}
output.emit({ type: "text-delta", text: finalText.slice(streamedText.length) });
streamedText = finalText;
}
};
const endStartedTools = () => {
for (const [index, tool] of tools) {
if (tool.started && !tool.ended) {
tool.ended = true;
output.emit({ type: "tool-call-end", index });
}
}
};
const emitReplayState = () => {
if (reasoningItems.length > 0) {
output.emit({ type: "replay-state", providerKind: REPLAY_KIND, value: { items: reasoningItems.slice() } });
reasoningItems.length = 0;
}
};
for await (const line of response.lines) {
if (!line.startsWith("data:")) continue;
const payload = line.slice(5).trim();
if (!payload) continue;
if (payload === "[DONE]") break;
let value: Record<string, unknown>;
try {
value = record(JSON.parse(payload)) ?? {};
} catch {
throw new Error("OpenAI Responses SSE returned invalid JSON");
}
switch (value.type) {
case "response.output_text.delta": {
closeThinking();
if (!textOpen) {
textOpen = true;
output.emit({ type: "text-start" });
}
const delta = text(value.delta);
if (delta !== null) {
streamedText += delta;
output.emit({ type: "text-delta", text: delta });
}
break;
}
case "response.output_text.done": {
const finalText = text(value.text);
if (finalText !== null) reconcileText(finalText);
closeText();
break;
}
case "response.reasoning_summary_text.delta":
case "response.reasoning_text.delta": {
if (!thinkingOpen) {
thinkingOpen = true;
output.emit({ type: "thinking-start" });
}
const delta = text(value.delta);
if (delta !== null) output.emit({ type: "thinking-delta", text: delta });
break;
}
case "response.reasoning_summary_text.done":
case "response.reasoning_text.done":
closeThinking();
break;
case "response.output_item.added": {
const item = record(value.item);
if (item?.type !== "function_call") break;
sawTool = true;
for (const event of updateTool(requiredIndex(value), item, { kind: "none" }, false, tools)) {
output.emit(event);
}
break;
}
case "response.output_item.done": {
const item = record(value.item);
if (item?.type === "reasoning") {
closeThinking();
reasoningItems.push(item as JsonValue);
} else if (item?.type === "message") {
sawCompletedItem = true;
const finalText = itemText(item);
if (finalText !== null) reconcileText(finalText);
closeText();
} else if (item?.type === "function_call") {
sawCompletedItem = true;
sawTool = true;
const snapshot = text(item.arguments);
const args: ToolArguments = snapshot === null ? { kind: "none" } : { kind: "snapshot", snapshot };
for (const event of updateTool(requiredIndex(value), item, args, true, tools)) {
output.emit(event);
}
}
break;
}
case "response.function_call_arguments.delta": {
const delta = text(value.delta);
if (delta === null) break;
sawTool = true;
for (const event of updateTool(requiredIndex(value), null, { kind: "delta", delta }, false, tools)) {
output.emit(event);
}
break;
}
case "response.function_call_arguments.done": {
const snapshot = text(value.arguments);
// 空快照不代表结束;等 output_item.done 收尾。
const args: ToolArguments = snapshot === null || snapshot === "" ? { kind: "none" } : { kind: "snapshot", snapshot };
const done = snapshot !== null && snapshot !== "";
for (const event of updateTool(requiredIndex(value), null, args, done, tools)) {
output.emit(event);
}
break;
}
case "response.completed": {
const usage = record(value.response)?.usage;
if (usage !== undefined) output.emit(usageEvent(usage));
closeThinking();
closeText();
endStartedTools();
for (const tool of tools.values()) {
if (!tool.started) {
throw new Error("OpenAI Responses completed with incomplete tool metadata");
}
}
terminal = true;
emitReplayState();
output.emit({ type: "done", reason: sawTool ? "tool-use" : "stop" });
break;
}
case "response.incomplete": {
closeThinking();
closeText();
endStartedTools();
terminal = true;
output.emit({ type: "done", reason: "length" });
break;
}
case "response.failed":
throw new Error(`OpenAI Responses failed: ${payload}`);
}
if (terminal) break;
}
if (!terminal && sawCompletedItem) {
closeThinking();
closeText();
for (const tool of tools.values()) {
if (!tool.ended) {
throw new Error("OpenAI Responses stream ended with an incomplete tool call");
}
}
terminal = true;
emitReplayState();
output.emit({ type: "done", reason: sawTool ? "tool-use" : "stop" });
}
if (!terminal) {
throw new Error("OpenAI Responses stream ended without response.completed or response.incomplete");
}
}
+129
View File
@@ -0,0 +1,129 @@
import type { JsonValue, LocalizedText, PluginContext } from "./plugin.ts";
import type { ModelSnapshot, ModelSupport } from "./model.ts";
import type { ResourcePatch, ResourceSnapshot } from "./resource.ts";
/**
* LLM 请求契约。宿主把它的规范会话(ProjectedMessage)投影成这个形状;
* 插件负责把它适配成上游 Provider 的协议。
*/
export type LlmContentPart =
| { type: "text"; text: string }
| { type: "image"; mediaType: string; dataBase64: string };
/** 不透明的 Provider 回放状态(如加密推理项);回放时按 providerKind 过滤。 */
export type LlmReplayState = {
providerKind: string;
value: JsonValue;
};
export type LlmToolCall = {
/** 同一轮内的稳定序号。 */
index: number;
callId: string;
name: string;
/** 已解析的 JSON 参数。 */
arguments: JsonValue;
};
export type LlmMessage =
| { role: "system" | "user"; content: LlmContentPart[] }
| {
role: "assistant";
text: string;
thinking: string;
replayState: LlmReplayState | null;
toolCalls: LlmToolCall[];
}
| {
role: "tool";
callId: string;
name: string;
content: string;
isError: boolean;
/** 非空时优先于 content,承载图片等富工具结果。 */
parts: LlmContentPart[];
};
export type LlmTool = {
name: string;
description: string;
/** 工具参数的 JSON Schema。 */
parameters: JsonValue;
};
export type LlmRequest = {
/** 系统指令;空字符串表示没有。 */
instructions: string;
messages: LlmMessage[];
tools: LlmTool[];
reasoning: { enabled: boolean; effort: string | null };
latency: "fast" | "standard";
maxOutputTokens: number | null;
/** 会话级稳定缓存键,用于上游前缀缓存的路由亲和(如 prompt_cache_key)。 */
cacheKey: string | null;
};
export type ModelUsage = {
inputTokens: number | null;
outputTokens: number | null;
totalTokens: number | null;
cacheReadTokens: number | null;
cacheWriteTokens: number | null;
reasoningTokens: number | null;
};
/**
* 标准化输出契约,与宿主统一流事件一一对应。插件边接收上游数据边发出事件;
* 文本、思考和每个工具调用都有显式的开始/结束边界,工具参数以增量交付。
* 回放状态在流结束前发出一次,宿主存入 assistant 消息供下一轮回放。
*/
export type ModelEvent =
| { type: "text-start" }
| { type: "text-delta"; text: string }
| { type: "text-end" }
| { type: "thinking-start" }
| { type: "thinking-delta"; text: string }
| { type: "thinking-end" }
| { type: "tool-call-start"; index: number; callId: string; name: string }
| { type: "tool-call-arguments-delta"; index: number; delta: string }
| { type: "tool-call-end"; index: number }
| { type: "replay-state"; providerKind: string; value: JsonValue }
| { type: "usage"; usage: ModelUsage }
| { type: "done"; reason: "stop" | "length" | "tool-use" };
export type ProviderOutput = {
emit(event: ModelEvent): void;
};
export type ProviderInvokeInput = {
model: ModelSnapshot;
/** 宿主为本次调用选中的资源;无资源 Provider 为 null。 */
resource: ResourceSnapshot | null;
request: LlmRequest;
};
/**
* `resource-error` 把失败归因到选中的资源,宿主据此更新资源状态,
* 并可在尚未发出任何事件时(未来)换一个资源重试。`patch` 同时用于
* 持久化成功调用的副作用,例如刷新后的 access token。
*/
export type ProviderResult =
| { status: "completed"; patch?: ResourcePatch }
| { status: "resource-error"; message: string; patch: ResourcePatch }
| { status: "request-error"; message: string; patch?: ResourcePatch };
export type ProviderSupport = {
id: string;
displayName: LocalizedText;
description?: LocalizedText;
/** 产品身份,用于归类与图标,如 "openai"。 */
providerType: string;
/** 每次调用消费的资源类型;无资源 Provider 可省略。 */
resourceType?: string;
models?: ModelSupport;
invoke(
input: ProviderInvokeInput,
output: ProviderOutput,
context: PluginContext,
): Promise<ProviderResult>;
};
+118
View File
@@ -0,0 +1,118 @@
import type { JsonValue, LocalizedText, PluginContext } from "./plugin.ts";
/**
* 资源是插件定义的私有记录(通常是上游账号),由 Provider 消费。
* 宿主负责持久化、列表和每次调用的资源选择;插件只负责创建、投影和解释资源。
*/
export type ResourceState =
| { status: "ready" }
| { status: "cooling"; retryAtMs?: number; message?: string }
| { status: "invalid"; message?: string };
/** 由添加流程或导入产生的新资源。 */
export type ResourceDraft = {
/** 去重键:宿主按 (资源类型, key) 执行 upsert。 */
key: string;
/** 凭证与插件私有字段;永远不会展示给用户。 */
privateData: JsonValue;
/** 缺省为 ready。 */
state?: ResourceState;
};
/** 宿主已持久化的一条资源。 */
export type ResourceSnapshot = {
/** 宿主分配的标识,区别于插件的去重键。 */
id: string;
type: string;
key: string;
privateData: JsonValue;
state: ResourceState;
};
/** 宿主原子应用到单条资源上的部分更新。 */
export type ResourcePatch = {
privateData?: JsonValue;
state?: ResourceState;
};
export type ResourceMetric = {
id: string;
label: LocalizedText;
unit: "percent" | "count";
/** percent 指标表示剩余占比,0..100。 */
value: number;
resetAtMs?: number;
};
/** 单条资源的用户可见投影;不得泄露凭证。displayName 是数据(如邮箱),保持纯字符串。 */
export type ResourceView = {
displayName: string;
description?: LocalizedText;
metrics?: ResourceMetric[];
};
/**
* OAuth 2.0 设备码式添加流程。宿主负责绘制 UI、驱动轮询循环
* (间隔、slow-down 退避、超时判定),并在流程存续期内在内存中持有
* `session`;插件只实现两次 HTTP 状态转移。
*/
export type OAuth2AddMethod = {
type: "oauth2.0";
id: string;
displayName: LocalizedText;
description?: LocalizedText;
begin(context: PluginContext): Promise<OAuth2Begin>;
poll(session: JsonValue, context: PluginContext): Promise<OAuth2Poll>;
};
export type OAuth2Begin = {
/** 不透明流程状态(设备码、PKCE verifier 等);永远不会持久化。 */
session: JsonValue;
userCode: string;
verificationUrl: string;
verificationUrlComplete?: string;
expiresAtMs: number;
pollIntervalMs: number;
};
export type OAuth2Poll =
| { status: "pending"; session?: JsonValue }
| { status: "slow-down"; session?: JsonValue }
| { status: "completed"; resources: ResourceDraft[] }
| { status: "denied"; message?: string }
| { status: "failed"; message: string };
export type ResourceAddMethod = OAuth2AddMethod;
export type ResourceImportFile = {
name: string;
/** 文件原文;解析和校验由插件负责。 */
content: string;
};
export type ResourceImportSupport = {
displayName: LocalizedText;
description?: LocalizedText;
/** 宿主文件选择器接受的扩展名,如 [".json"]。 */
accept: string[];
multiple?: boolean;
parse(files: ResourceImportFile[], context: PluginContext): Promise<ResourceImportResult>;
};
export type ResourceImportResult = {
resources: ResourceDraft[];
/** 单个文件的问题,值得提示但不必使整次导入失败。 */
warnings?: string[];
};
export type ResourceSupport = {
type: string;
displayName: LocalizedText;
add?: ResourceAddMethod[];
import?: ResourceImportSupport;
present(resource: ResourceSnapshot): ResourceView;
/** 用户主动触发时重新读取上游状态(额度、凭证有效性)。 */
refresh?(resource: ResourceSnapshot, context: PluginContext): Promise<ResourcePatch>;
/** 可选的上游撤销;宿主随后删除本地记录。 */
remove?(resource: ResourceSnapshot, context: PluginContext): Promise<void>;
};
+185
View File
@@ -0,0 +1,185 @@
import { __getRegisteredPlugin, type JsonValue, type NetworkEventStream, type PluginContext } from "cursor-byok:plugin";
import type { ModelEvent, ProviderSupport } from "cursor-byok:provider";
import type { ResourceAddMethod, ResourceSupport } from "cursor-byok:resource";
if (Deno.args.length !== 1) throw new Error("plugin entry URL is required");
await import(Deno.args[0]);
const plugin = __getRegisteredPlugin();
const encoder = new TextEncoder();
const writer = Deno.stdout.writable.getWriter();
const pendingHost = new Map<string, { resolve(value: unknown): void; reject(error: Error): void }>();
const controllers = new Map<string, AbortController>();
let hostSequence = 0;
// 事件与最终结果共用一条串行写队列,保证顺序。
let writeQueue = Promise.resolve();
function send(value: unknown): Promise<void> {
const operation = writeQueue.then(() => writer.write(encoder.encode(JSON.stringify(value) + "\n")));
writeQueue = operation.catch(() => undefined);
return operation;
}
function hostCall(requestId: string, method: string, params: unknown): Promise<unknown> {
const id = `${requestId}:host:${++hostSequence}`;
return new Promise((resolve, reject) => {
pendingHost.set(id, { resolve, reject });
void send({ type: "host_call", id, requestId, method, params });
});
}
async function* streamLines(requestId: string, streamId: string): AsyncGenerator<string> {
try {
for (;;) {
const chunk = await hostCall(requestId, "network.stream.read", { streamId }) as {
lines: string[];
done: boolean;
};
for (const line of chunk.lines) yield line;
if (chunk.done) return;
}
} finally {
void hostCall(requestId, "network.stream.close", { streamId }).catch(() => undefined);
}
}
function contextFor(requestId: string, signal: AbortSignal): PluginContext {
return {
network: {
fetch: (url, init = {}) => hostCall(requestId, "network.fetch", { url, ...init }) as ReturnType<PluginContext["network"]["fetch"]>,
stream: async (url, init = {}): Promise<NetworkEventStream> => {
const opened = await hostCall(requestId, "network.stream.open", { url, ...init }) as {
streamId: string;
status: number;
headers: Record<string, string>;
};
return {
status: opened.status,
headers: opened.headers,
lines: streamLines(requestId, opened.streamId),
};
},
},
signal,
};
}
function provider(id: unknown): ProviderSupport {
const found = plugin.providers.find((provider) => provider.id === id);
if (!found) throw new Error(`unknown plugin provider: ${id}`);
return found;
}
function resourceSupport(type: unknown): ResourceSupport {
const found = (plugin.resources ?? []).find((resource) => resource.type === type);
if (!found) throw new Error(`unknown plugin resource type: ${type}`);
return found;
}
function addMethod(support: ResourceSupport, methodId: unknown): ResourceAddMethod {
const found = (support.add ?? []).find((method) => method.id === methodId);
if (!found) throw new Error(`unknown plugin add method: ${methodId}`);
return found;
}
async function dispatch(message: { id: string; method: string; params?: JsonValue }) {
const controller = new AbortController();
controllers.set(message.id, controller);
const context = contextFor(message.id, controller.signal);
const params = (message.params ?? {}) as Record<string, JsonValue>;
try {
let result: unknown;
switch (message.method) {
case "provider.invoke": {
const output = {
emit: (event: ModelEvent) => void send({ type: "event", id: message.id, event }),
};
result = await provider(params.providerId).invoke(
{
model: params.model as never,
resource: (params.resource ?? null) as never,
request: params.request as never,
},
output,
context,
);
break;
}
case "models.list": {
const models = provider(params.providerId).models;
if (!models) throw new Error(`plugin provider ${params.providerId} has no models`);
result = await models.list({ resource: (params.resource ?? null) as never }, context);
break;
}
case "resource.present": {
const support = resourceSupport(params.resourceType);
const resources = Array.isArray(params.resources) ? params.resources : [];
result = resources.map((resource) => support.present(resource as never));
break;
}
case "resource.refresh": {
const support = resourceSupport(params.resourceType);
if (!support.refresh) throw new Error(`resource ${params.resourceType} has no refresh`);
result = await support.refresh(params.resource as never, context);
break;
}
case "resource.remove": {
const support = resourceSupport(params.resourceType);
await support.remove?.(params.resource as never, context);
result = null;
break;
}
case "oauth.begin": {
const support = resourceSupport(params.resourceType);
result = await addMethod(support, params.methodId).begin(context);
break;
}
case "oauth.poll": {
const support = resourceSupport(params.resourceType);
result = await addMethod(support, params.methodId).poll(params.session ?? null, context);
break;
}
case "import.parse": {
const support = resourceSupport(params.resourceType);
if (!support.import) throw new Error(`resource ${params.resourceType} has no import`);
const files = Array.isArray(params.files) ? params.files : [];
result = await support.import.parse(files as never, context);
break;
}
default:
throw new Error(`unknown plugin method: ${message.method}`);
}
await send({ type: "result", id: message.id, result: result ?? null });
} catch (error) {
await send({ type: "result", id: message.id, error: error instanceof Error ? error.message : String(error) });
} finally {
controllers.delete(message.id);
}
}
let buffered = "";
for await (const chunk of Deno.stdin.readable.pipeThrough(new TextDecoderStream())) {
buffered += chunk;
for (;;) {
const newline = buffered.indexOf("\n");
if (newline < 0) break;
const line = buffered.slice(0, newline);
buffered = buffered.slice(newline + 1);
if (!line.trim()) continue;
const message = JSON.parse(line);
if (message.type === "request") {
void dispatch(message);
} else if (message.type === "cancel") {
controllers.get(message.id)?.abort();
} else if (message.type === "host_result") {
const pending = pendingHost.get(message.id);
if (!pending) continue;
pendingHost.delete(message.id);
pending.resolve(message.result);
} else if (message.type === "host_error") {
const pending = pendingHost.get(message.id);
if (!pending) continue;
pendingHost.delete(message.id);
pending.reject(new Error(message.error));
}
}
}
+466
View File
@@ -0,0 +1,466 @@
//! Owns core-side persistence of plugin resources and model catalogs.
use serde::{Deserialize, Serialize};
use super::data::PluginDataStore;
use crate::{Error, Result};
/// 核心理解的资源运行状态;插件只能通过 draft/patch/report 改变它。
#[derive(Clone, Debug, Deserialize, Serialize, PartialEq)]
#[serde(tag = "status", rename_all = "snake_case")]
pub enum ResourceState {
Ready,
Cooling {
#[serde(default, skip_serializing_if = "Option::is_none")]
retry_at_ms: Option<i64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
message: Option<String>,
},
Invalid {
#[serde(default, skip_serializing_if = "Option::is_none")]
message: Option<String>,
},
}
impl ResourceState {
/// 冷却到期后自动恢复可用。
pub fn is_ready(&self, now_ms: i64) -> bool {
match self {
Self::Ready => true,
Self::Cooling { retry_at_ms, .. } => retry_at_ms.is_some_and(|at| at <= now_ms),
Self::Invalid { .. } => false,
}
}
}
/// 核心持久化的一条插件资源。`private_data` 只回传给插件。
#[derive(Clone, Debug, Deserialize, Serialize)]
pub struct ResourceRecord {
pub id: String,
pub key: String,
pub private_data: serde_json::Value,
pub state: ResourceState,
pub created_at_ms: i64,
pub updated_at_ms: i64,
}
impl ResourceRecord {
/// 传给插件的快照形状(SDK 的 ResourceSnapshot)。
pub fn snapshot(&self, resource_type: &str) -> serde_json::Value {
serde_json::json!({
"id": self.id,
"type": resource_type,
"key": self.key,
"privateData": self.private_data,
"state": state_json(&self.state),
})
}
}
fn state_json(state: &ResourceState) -> serde_json::Value {
match state {
ResourceState::Ready => serde_json::json!({ "status": "ready" }),
ResourceState::Cooling {
retry_at_ms,
message,
} => serde_json::json!({
"status": "cooling",
"retryAtMs": retry_at_ms,
"message": message,
}),
ResourceState::Invalid { message } => serde_json::json!({
"status": "invalid",
"message": message,
}),
}
}
/// 插件返回的新资源(SDK 的 ResourceDraft)。
#[derive(Clone, Debug, Deserialize)]
#[serde(rename_all = "camelCase", deny_unknown_fields)]
pub struct ResourceDraft {
pub key: String,
pub private_data: serde_json::Value,
#[serde(default)]
pub state: Option<ResourceStateInput>,
}
/// 插件对单条资源的部分更新(SDK 的 ResourcePatch)。
#[derive(Clone, Debug, Default, Deserialize)]
#[serde(rename_all = "camelCase", deny_unknown_fields)]
pub struct ResourcePatch {
#[serde(default)]
pub private_data: Option<serde_json::Value>,
#[serde(default)]
pub state: Option<ResourceStateInput>,
}
/// SDK 侧 camelCase 状态输入,转换成核心存储形状。
#[derive(Clone, Debug, Deserialize)]
#[serde(tag = "status", rename_all = "kebab-case", deny_unknown_fields)]
pub enum ResourceStateInput {
Ready,
Cooling {
#[serde(default, rename = "retryAtMs")]
retry_at_ms: Option<i64>,
#[serde(default)]
message: Option<String>,
},
Invalid {
#[serde(default)]
message: Option<String>,
},
}
impl From<ResourceStateInput> for ResourceState {
fn from(input: ResourceStateInput) -> Self {
match input {
ResourceStateInput::Ready => Self::Ready,
ResourceStateInput::Cooling {
retry_at_ms,
message,
} => Self::Cooling {
retry_at_ms,
message,
},
ResourceStateInput::Invalid { message } => Self::Invalid { message },
}
}
}
/// 插件发现的一个模型(SDK 的 ModelDefinition),由核心整体替换目录。
#[derive(Clone, Debug, Deserialize, Serialize)]
pub struct StoredModel {
pub id: String,
pub display_name: String,
#[serde(default)]
pub description: Option<String>,
#[serde(default)]
pub context_window_tokens: Option<u64>,
#[serde(default)]
pub max_output_tokens: Option<u64>,
#[serde(default)]
pub thinking: bool,
#[serde(default)]
pub images: bool,
#[serde(default)]
pub private_data: serde_json::Value,
}
impl StoredModel {
pub fn from_definition(value: &serde_json::Value) -> Result<Self> {
let object = value
.as_object()
.ok_or_else(|| Error::Protocol("plugin model definition must be an object".into()))?;
let id = object
.get("id")
.and_then(serde_json::Value::as_str)
.filter(|id| !id.trim().is_empty())
.ok_or_else(|| Error::Protocol("plugin model definition requires id".into()))?;
let display_name = object
.get("displayName")
.and_then(serde_json::Value::as_str)
.filter(|name| !name.trim().is_empty())
.ok_or_else(|| {
Error::Protocol("plugin model definition requires displayName".into())
})?;
let capabilities = object
.get("capabilities")
.and_then(|value| value.as_object());
let capability = |name: &str| {
capabilities
.and_then(|value| value.get(name))
.and_then(serde_json::Value::as_bool)
.unwrap_or(false)
};
Ok(Self {
id: id.to_owned(),
display_name: display_name.to_owned(),
description: object
.get("description")
.and_then(serde_json::Value::as_str)
.map(str::to_owned),
context_window_tokens: object
.get("contextWindowTokens")
.and_then(serde_json::Value::as_u64),
max_output_tokens: object
.get("maxOutputTokens")
.and_then(serde_json::Value::as_u64),
thinking: capability("thinking"),
images: capability("images"),
private_data: object
.get("privateData")
.cloned()
.unwrap_or(serde_json::Value::Null),
})
}
/// 传给插件的模型快照(SDK 的 ModelSnapshot)。
pub fn snapshot(&self) -> serde_json::Value {
serde_json::json!({
"id": self.id,
"displayName": self.display_name,
"description": self.description,
"contextWindowTokens": self.context_window_tokens,
"maxOutputTokens": self.max_output_tokens,
"capabilities": { "thinking": self.thinking, "images": self.images },
"privateData": self.private_data,
})
}
}
/// 资源与模型目录的核心存储,构建在插件私有 JSON 文件之上。
#[derive(Clone)]
pub struct PluginStateStore {
data: PluginDataStore,
}
pub struct UpsertOutcome {
pub added: usize,
pub updated: usize,
}
impl PluginStateStore {
pub fn new(data: PluginDataStore) -> Self {
Self { data }
}
pub async fn resources(
&self,
plugin_id: &str,
resource_type: &str,
) -> Result<Vec<ResourceRecord>> {
let value = self
.data
.read(plugin_id, &resource_key(resource_type))
.await?;
if value.is_null() {
return Ok(Vec::new());
}
Ok(serde_json::from_value(value)?)
}
pub async fn upsert_resources(
&self,
plugin_id: &str,
resource_type: &str,
drafts: Vec<ResourceDraft>,
) -> Result<UpsertOutcome> {
let mut records = self.resources(plugin_id, resource_type).await?;
let now = now_ms();
let mut outcome = UpsertOutcome {
added: 0,
updated: 0,
};
for draft in drafts {
if draft.key.trim().is_empty() {
return Err(Error::Protocol("plugin resource draft requires key".into()));
}
let state = draft
.state
.map_or(ResourceState::Ready, ResourceState::from);
match records.iter_mut().find(|record| record.key == draft.key) {
Some(existing) => {
existing.private_data = draft.private_data;
existing.state = state;
existing.updated_at_ms = now;
outcome.updated += 1;
}
None => {
records.push(ResourceRecord {
id: uuid::Uuid::new_v4().to_string(),
key: draft.key,
private_data: draft.private_data,
state,
created_at_ms: now,
updated_at_ms: now,
});
outcome.added += 1;
}
}
}
self.save_resources(plugin_id, resource_type, &records)
.await?;
Ok(outcome)
}
pub async fn apply_patch(
&self,
plugin_id: &str,
resource_type: &str,
resource_id: &str,
patch: ResourcePatch,
) -> Result<()> {
let mut records = self.resources(plugin_id, resource_type).await?;
let record = records
.iter_mut()
.find(|record| record.id == resource_id)
.ok_or_else(|| Error::RunNotFound(format!("plugin resource {resource_id}")))?;
if let Some(private_data) = patch.private_data {
record.private_data = private_data;
}
if let Some(state) = patch.state {
record.state = state.into();
}
record.updated_at_ms = now_ms();
self.save_resources(plugin_id, resource_type, &records)
.await
}
pub async fn remove_resource(
&self,
plugin_id: &str,
resource_type: &str,
resource_id: &str,
) -> Result<ResourceRecord> {
let mut records = self.resources(plugin_id, resource_type).await?;
let index = records
.iter()
.position(|record| record.id == resource_id)
.ok_or_else(|| Error::RunNotFound(format!("plugin resource {resource_id}")))?;
let removed = records.remove(index);
self.save_resources(plugin_id, resource_type, &records)
.await?;
Ok(removed)
}
pub async fn models(&self, plugin_id: &str, provider_id: &str) -> Result<Vec<StoredModel>> {
let value = self.data.read(plugin_id, &model_key(provider_id)).await?;
if value.is_null() {
return Ok(Vec::new());
}
Ok(serde_json::from_value(value)?)
}
pub async fn replace_models(
&self,
plugin_id: &str,
provider_id: &str,
models: &[StoredModel],
) -> Result<()> {
self.data
.update(
plugin_id,
&model_key(provider_id),
&serde_json::to_value(models)?,
)
.await
}
pub async fn clear(&self, plugin_id: &str) -> Result<()> {
self.data.clear(plugin_id).await
}
async fn save_resources(
&self,
plugin_id: &str,
resource_type: &str,
records: &[ResourceRecord],
) -> Result<()> {
self.data
.update(
plugin_id,
&resource_key(resource_type),
&serde_json::to_value(records)?,
)
.await
}
}
pub fn now_ms() -> i64 {
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|duration| duration.as_millis().min(i64::MAX as u128) as i64)
.unwrap_or_default()
}
fn resource_key(resource_type: &str) -> String {
format!("resources-{resource_type}")
}
fn model_key(provider_id: &str) -> String {
format!("models-{provider_id}")
}
#[cfg(test)]
mod tests {
use super::*;
fn store() -> (tempfile::TempDir, PluginStateStore) {
let root = tempfile::tempdir().unwrap();
let data = PluginDataStore::for_test(root.path().join("data")).unwrap();
(root, PluginStateStore::new(data))
}
#[tokio::test]
async fn upserts_resources_by_key_and_applies_patches() {
let (_root, store) = store();
let outcome = store
.upsert_resources(
"dev.example",
"account",
vec![ResourceDraft {
key: "acct-1".into(),
private_data: serde_json::json!({"token":"one"}),
state: None,
}],
)
.await
.unwrap();
assert_eq!(outcome.added, 1);
let outcome = store
.upsert_resources(
"dev.example",
"account",
vec![ResourceDraft {
key: "acct-1".into(),
private_data: serde_json::json!({"token":"two"}),
state: None,
}],
)
.await
.unwrap();
assert_eq!(outcome.updated, 1);
let records = store.resources("dev.example", "account").await.unwrap();
assert_eq!(records.len(), 1);
assert_eq!(records[0].private_data["token"], "two");
store
.apply_patch(
"dev.example",
"account",
&records[0].id,
ResourcePatch {
private_data: None,
state: Some(ResourceStateInput::Cooling {
retry_at_ms: Some(200),
message: None,
}),
},
)
.await
.unwrap();
let records = store.resources("dev.example", "account").await.unwrap();
assert!(!records[0].state.is_ready(100));
assert!(records[0].state.is_ready(300), "cooling expires over time");
}
#[tokio::test]
async fn replaces_model_catalogs() {
let (_root, store) = store();
let model = StoredModel::from_definition(&serde_json::json!({
"id": "gpt-test",
"displayName": "GPT Test",
"capabilities": {"thinking": true},
"privateData": {"reasoningEfforts": ["low"]},
}))
.unwrap();
store
.replace_models("dev.example", "codex", &[model])
.await
.unwrap();
let models = store.models("dev.example", "codex").await.unwrap();
assert_eq!(models.len(), 1);
assert!(models[0].thinking);
assert_eq!(models[0].private_data["reasoningEfforts"][0], "low");
}
}
+297
View File
@@ -0,0 +1,297 @@
//! Translates between core model types and the plugin SDK wire contract.
use base64::{engine::general_purpose::STANDARD, Engine};
use crate::{
model::{
ContentPart, ModelInvocation, ModelLatency, ProjectedContent, ProjectedMessage,
ProviderReplayState, Role, Usage,
},
provider::{FinishReason, ModelEvent},
Error, Result,
};
/// 把一次核心模型调用投影成 SDK 的 LlmRequest。
pub fn llm_request(invocation: &ModelInvocation) -> Result<serde_json::Value> {
let request = &invocation.request;
let messages = request
.history
.iter()
.map(wire_message)
.collect::<Result<Vec<_>>>()?;
Ok(serde_json::json!({
"instructions": request.prompt.instructions,
"messages": messages,
"tools": request.prompt.tools.iter().map(|tool| serde_json::json!({
"name": tool.name,
"description": tool.description,
"parameters": tool.parameters,
})).collect::<Vec<_>>(),
"reasoning": {
"enabled": request.model.reasoning.enabled,
"effort": request.model.reasoning.effort,
},
"latency": match request.model.latency {
ModelLatency::Fast => "fast",
_ => "standard",
},
"maxOutputTokens": request.model.max_output_tokens,
"cacheKey": invocation.conversation_id,
}))
}
fn wire_message(message: &ProjectedMessage) -> Result<serde_json::Value> {
match &message.content {
ProjectedContent::Parts(parts) => match message.role {
Role::System | Role::User => Ok(serde_json::json!({
"role": if message.role == Role::System { "system" } else { "user" },
"content": wire_parts(parts),
})),
// 纯文本 assistant 历史消息投影成无工具调用的 assistant。
Role::Assistant => Ok(serde_json::json!({
"role": "assistant",
"text": joined_text(parts),
"thinking": "",
"replayState": serde_json::Value::Null,
"toolCalls": [],
})),
Role::Tool => Err(Error::Protocol(
"tool messages must carry a tool result".into(),
)),
},
ProjectedContent::Assistant {
text,
thinking,
replay_state,
calls,
} => Ok(serde_json::json!({
"role": "assistant",
"text": text,
"thinking": thinking,
"replayState": replay_state.as_ref().map(|state| serde_json::json!({
"providerKind": state.provider_kind,
"value": state.value,
})),
"toolCalls": calls.iter().map(|call| serde_json::json!({
"index": call.index,
"callId": call.call_id,
"name": call.name,
"arguments": call.arguments,
})).collect::<Vec<_>>(),
})),
ProjectedContent::ToolResult(result) => Ok(serde_json::json!({
"role": "tool",
"callId": result.call_id,
"name": result.name,
"content": result.content,
"isError": result.is_error,
"parts": wire_parts(&result.provider_parts),
})),
}
}
fn wire_parts(parts: &[ContentPart]) -> Vec<serde_json::Value> {
parts
.iter()
.map(|part| match part {
ContentPart::Text { text } => serde_json::json!({ "type": "text", "text": text }),
ContentPart::Image { mime_type, data } => serde_json::json!({
"type": "image",
"mediaType": mime_type,
"dataBase64": STANDARD.encode(data),
}),
})
.collect()
}
fn joined_text(parts: &[ContentPart]) -> String {
parts
.iter()
.filter_map(|part| match part {
ContentPart::Text { text } => Some(text.as_str()),
ContentPart::Image { .. } => None,
})
.collect()
}
/// 把插件发出的标准化事件解析为核心 ModelEvent。
pub fn model_event(value: &serde_json::Value) -> Result<ModelEvent> {
let kind = value
.get("type")
.and_then(serde_json::Value::as_str)
.ok_or_else(|| Error::Protocol("plugin model event requires type".into()))?;
let event = match kind {
"text-start" => ModelEvent::TextStart,
"text-delta" => ModelEvent::TextDelta(required_str(value, "text")?.to_owned()),
"text-end" => ModelEvent::TextEnd,
"thinking-start" => ModelEvent::ThinkingStart,
"thinking-delta" => ModelEvent::ThinkingDelta(required_str(value, "text")?.to_owned()),
"thinking-end" => ModelEvent::ThinkingEnd,
"tool-call-start" => ModelEvent::ToolCallStart {
index: required_index(value)?,
call_id: required_str(value, "callId")?.to_owned(),
name: required_str(value, "name")?.to_owned(),
},
"tool-call-arguments-delta" => ModelEvent::ToolCallArgumentsDelta {
index: required_index(value)?,
delta: required_str(value, "delta")?.to_owned(),
},
"tool-call-end" => ModelEvent::ToolCallEnd {
index: required_index(value)?,
},
"replay-state" => ModelEvent::ProviderReplayState(ProviderReplayState {
provider_kind: required_str(value, "providerKind")?.to_owned(),
value: value.get("value").cloned().unwrap_or_default(),
}),
"usage" => {
let usage = value
.get("usage")
.ok_or_else(|| Error::Protocol("plugin usage event requires usage".into()))?;
let tokens = |name: &str| usage.get(name).and_then(serde_json::Value::as_u64);
ModelEvent::Usage(Usage {
input_tokens: tokens("inputTokens"),
output_tokens: tokens("outputTokens"),
total_tokens: tokens("totalTokens"),
cache_read_tokens: tokens("cacheReadTokens"),
cache_write_tokens: tokens("cacheWriteTokens"),
reasoning_tokens: tokens("reasoningTokens"),
})
}
"done" => ModelEvent::Done(match required_str(value, "reason")? {
"stop" => FinishReason::Stop,
"length" => FinishReason::Length,
"tool-use" => FinishReason::ToolUse,
reason => {
return Err(Error::Protocol(format!(
"unknown plugin finish reason: {reason}"
)))
}
}),
kind => {
return Err(Error::Protocol(format!(
"unknown plugin model event: {kind}"
)))
}
};
Ok(event)
}
fn required_str<'a>(value: &'a serde_json::Value, key: &str) -> Result<&'a str> {
value
.get(key)
.and_then(serde_json::Value::as_str)
.ok_or_else(|| Error::Protocol(format!("plugin model event requires string '{key}'")))
}
fn required_index(value: &serde_json::Value) -> Result<usize> {
value
.get("index")
.and_then(serde_json::Value::as_u64)
.map(|index| index as usize)
.ok_or_else(|| Error::Protocol("plugin model event requires index".into()))
}
#[cfg(test)]
mod tests {
use super::*;
use crate::model::{
ModelRequest, ModelSpec, ProjectedContent, PromptSpec, ToolCallContent, ToolResultContent,
};
#[test]
fn projects_history_into_wire_messages() {
let invocation = ModelInvocation {
call_id: "call".into(),
run_id: "run".into(),
conversation_id: "conversation".into(),
provider_call_index: 0,
request: ModelRequest {
prompt: PromptSpec {
instructions: "be brief".into(),
tools: Vec::new(),
},
model: ModelSpec::new("plugin:p/c/m"),
history: vec![
ProjectedMessage {
message_id: "m1".into(),
role: Role::User,
content: ProjectedContent::Parts(vec![ContentPart::Text {
text: "hi".into(),
}]),
},
ProjectedMessage {
message_id: "m2".into(),
role: Role::Assistant,
content: ProjectedContent::Assistant {
text: "".into(),
thinking: "t".into(),
replay_state: Some(ProviderReplayState {
provider_kind: "openai_responses".into(),
value: serde_json::json!({"items": []}),
}),
calls: vec![ToolCallContent {
index: 0,
call_id: "c1".into(),
name: "read".into(),
arguments: serde_json::json!({"path":"a"}),
}],
},
},
ProjectedMessage {
message_id: "m3".into(),
role: Role::Tool,
content: ProjectedContent::ToolResult(ToolResultContent {
call_id: "c1".into(),
name: "read".into(),
content: "data".into(),
is_error: false,
image: None,
provider_parts: Vec::new(),
}),
},
],
},
};
let request = llm_request(&invocation).unwrap();
assert_eq!(request["instructions"], "be brief");
assert_eq!(request["latency"], "standard");
assert_eq!(request["cacheKey"], "conversation");
let messages = request["messages"].as_array().unwrap();
assert_eq!(messages[0]["role"], "user");
assert_eq!(
messages[1]["replayState"]["providerKind"],
"openai_responses"
);
assert_eq!(messages[1]["toolCalls"][0]["callId"], "c1");
assert_eq!(messages[2]["role"], "tool");
assert_eq!(messages[2]["isError"], false);
}
#[test]
fn parses_plugin_events_into_model_events() {
assert_eq!(
model_event(&serde_json::json!({"type":"text-delta","text":"hi"})).unwrap(),
ModelEvent::TextDelta("hi".into())
);
assert_eq!(
model_event(&serde_json::json!({"type":"done","reason":"tool-use"})).unwrap(),
ModelEvent::Done(FinishReason::ToolUse)
);
let usage = model_event(&serde_json::json!({
"type":"usage",
"usage":{"inputTokens":10,"outputTokens":2,"cacheReadTokens":4}
}))
.unwrap();
assert_eq!(
usage,
ModelEvent::Usage(Usage {
input_tokens: Some(10),
output_tokens: Some(2),
total_tokens: None,
cache_read_tokens: Some(4),
cache_write_tokens: None,
reasoning_tokens: None,
})
);
assert!(model_event(&serde_json::json!({"type":"mystery"})).is_err());
}
}
+629
View File
@@ -0,0 +1,629 @@
//! Runs one long-lived, sandboxed Deno process per active plugin.
use std::{
collections::{HashMap, HashSet},
path::PathBuf,
process::Stdio,
sync::Arc,
time::Duration,
};
use tokio::{
io::{AsyncBufReadExt, AsyncWriteExt, BufReader},
process::{Child, ChildStdin},
sync::{mpsc, Mutex},
};
use tokio_util::sync::CancellationToken;
use super::{
catalog::PluginEntry,
definition::{file_url, PluginDefinitionLoader},
protocol::{HostMessage, WorkerMessage},
};
use crate::{store::Store, Error, Result};
const INVOCATION_TIMEOUT: Duration = Duration::from_secs(10 * 60);
const MAX_NETWORK_RESPONSE_BYTES: u64 = 16 * 1024 * 1024;
const MAX_STREAM_BYTES: u64 = 256 * 1024 * 1024;
/// 一次流式调用的输出:零或多个事件,然后恰好一个最终结果。
#[derive(Debug)]
pub enum WorkerStreamItem {
Event(serde_json::Value),
Result(Result<serde_json::Value>),
}
type Pending = Arc<Mutex<HashMap<String, mpsc::UnboundedSender<WorkerStreamItem>>>>;
type StreamLines = Arc<Mutex<mpsc::Receiver<Result<String>>>>;
#[derive(Clone)]
pub struct PluginWorker {
inner: Arc<PluginWorkerInner>,
}
struct PluginWorkerInner {
plugin_id: String,
executable: PathBuf,
directory: PathBuf,
entry: PathBuf,
loader: PluginDefinitionLoader,
host: HostContext,
process: Mutex<Option<WorkerProcess>>,
pending: Pending,
}
struct WorkerProcess {
child: Child,
stdin: Arc<Mutex<ChildStdin>>,
}
#[derive(Clone)]
struct HostContext {
plugin_id: String,
network_hosts: Arc<HashSet<String>>,
store: Store,
cancellations: Arc<Mutex<HashMap<String, CancellationToken>>>,
streams: Arc<Mutex<HashMap<String, StreamLines>>>,
}
impl PluginWorker {
pub fn new(
plugin: &PluginEntry,
executable: PathBuf,
loader: PluginDefinitionLoader,
store: Store,
) -> Self {
let plugin_id = plugin.manifest.id.clone();
Self {
inner: Arc::new(PluginWorkerInner {
host: HostContext {
plugin_id: plugin_id.clone(),
network_hosts: Arc::new(
plugin
.manifest
.permissions
.network
.iter()
.map(|host| host.to_ascii_lowercase())
.collect(),
),
store,
cancellations: Arc::new(Mutex::new(HashMap::new())),
streams: Arc::new(Mutex::new(HashMap::new())),
},
plugin_id,
executable,
directory: plugin.directory.clone(),
entry: plugin.entry.clone(),
loader,
process: Mutex::new(None),
pending: Arc::new(Mutex::new(HashMap::new())),
}),
}
}
/// 一元调用:忽略事件,等待最终结果,受统一超时约束。
pub async fn invoke(
&self,
method: &str,
params: serde_json::Value,
cancellation: CancellationToken,
) -> Result<serde_json::Value> {
let mut items = self.invoke_streaming(method, params, cancellation).await?;
let result = tokio::time::timeout(INVOCATION_TIMEOUT, async {
while let Some(item) = items.recv().await {
if let WorkerStreamItem::Result(result) = item {
return result;
}
}
Err(Error::Provider(format!(
"plugin '{}' worker stopped",
self.inner.plugin_id
)))
})
.await;
match result {
Ok(result) => result,
Err(_) => Err(Error::Provider(format!(
"plugin '{}' invocation timed out",
self.inner.plugin_id
))),
}
}
/// 流式调用:事件按序转发,最终以恰好一个 Result 收尾。
/// 取消通过传入的令牌传播到 Worker 与其挂起的宿主网络请求。
pub async fn invoke_streaming(
&self,
method: &str,
params: serde_json::Value,
cancellation: CancellationToken,
) -> Result<mpsc::UnboundedReceiver<WorkerStreamItem>> {
let id = uuid::Uuid::new_v4().to_string();
let request_cancellation = CancellationToken::new();
self.inner
.host
.cancellations
.lock()
.await
.insert(id.clone(), request_cancellation.clone());
let (sender, receiver) = mpsc::unbounded_channel();
self.inner
.pending
.lock()
.await
.insert(id.clone(), sender.clone());
let send_result = async {
let stdin = self.stdin().await?;
write_message(
&stdin,
&HostMessage::Request {
id: &id,
method,
params: &params,
},
)
.await
}
.await;
if let Err(error) = send_result {
self.cleanup(&id).await;
return Err(error);
}
// 取消监视:通知 Worker,同时中止该请求挂起的宿主网络调用。
let inner = self.inner.clone();
let request_id = id.clone();
tokio::spawn(async move {
tokio::select! {
_ = cancellation.cancelled() => {
request_cancellation.cancel();
if let Some(process) = inner.process.lock().await.as_ref() {
let _ = write_message(&process.stdin, &HostMessage::Cancel { id: &request_id }).await;
}
let _ = sender.send(WorkerStreamItem::Result(Err(Error::Cancelled)));
inner.pending.lock().await.remove(&request_id);
inner.host.cancellations.lock().await.remove(&request_id);
}
_ = sender.closed() => {
inner.host.cancellations.lock().await.remove(&request_id);
}
}
});
Ok(receiver)
}
pub async fn stop(&self) {
if let Some(mut process) = self.inner.process.lock().await.take() {
let _ = process.child.kill().await;
}
fail_pending(&self.inner.pending, "plugin worker stopped").await;
}
async fn cleanup(&self, id: &str) {
self.inner.pending.lock().await.remove(id);
self.inner.host.cancellations.lock().await.remove(id);
}
async fn stdin(&self) -> Result<Arc<Mutex<ChildStdin>>> {
let mut process = self.inner.process.lock().await;
let dead = match process.as_mut() {
Some(current) => current.child.try_wait()?.is_some(),
None => true,
};
if dead {
*process = Some(self.spawn().await?);
}
Ok(process
.as_ref()
.expect("plugin worker was started")
.stdin
.clone())
}
async fn spawn(&self) -> Result<WorkerProcess> {
let entry_url = file_url(&self.inner.entry)?;
let mut command = tokio::process::Command::new(&self.inner.executable);
command
.arg("run")
.arg("--quiet")
.arg("--no-config")
.arg("--no-lock")
.arg("--no-npm")
.arg("--no-remote")
.arg("--no-prompt")
.arg(format!("--allow-read={}", self.inner.directory.display()))
.arg(format!(
"--allow-read={}",
self.inner.loader.sdk_dir().display()
))
.arg(format!(
"--import-map={}",
self.inner.loader.import_map().display()
))
.arg(self.inner.loader.worker_path())
.arg(entry_url.as_str())
.env("DENO_DIR", self.inner.loader.deno_dir())
.env("DENO_NO_UPDATE_CHECK", "1")
.current_dir(&self.inner.directory)
.stdin(Stdio::piped())
.stdout(Stdio::piped())
.stderr(Stdio::piped())
.kill_on_drop(true);
let mut child = command.spawn()?;
let stdin =
Arc::new(Mutex::new(child.stdin.take().ok_or_else(|| {
Error::Config("cannot open plugin worker stdin".into())
})?));
let stdout = child
.stdout
.take()
.ok_or_else(|| Error::Config("cannot open plugin worker stdout".into()))?;
let stderr = child
.stderr
.take()
.ok_or_else(|| Error::Config("cannot open plugin worker stderr".into()))?;
spawn_stdout_reader(
self.inner.plugin_id.clone(),
stdout,
stdin.clone(),
self.inner.pending.clone(),
self.inner.host.clone(),
);
spawn_stderr_reader(self.inner.plugin_id.clone(), stderr);
Ok(WorkerProcess { child, stdin })
}
}
fn spawn_stdout_reader(
plugin_id: String,
stdout: tokio::process::ChildStdout,
stdin: Arc<Mutex<ChildStdin>>,
pending: Pending,
host: HostContext,
) {
tokio::spawn(async move {
let mut lines = BufReader::new(stdout).lines();
while let Ok(Some(line)) = lines.next_line().await {
let message = match serde_json::from_str::<WorkerMessage>(&line) {
Ok(message) => message,
Err(error) => {
tracing::warn!(plugin = %plugin_id, %error, "plugin worker wrote an invalid message");
continue;
}
};
match message {
WorkerMessage::Result { id, result, error } => {
if let Some(sender) = pending.lock().await.remove(&id) {
let value = match error {
Some(error) => {
Err(Error::Provider(format!("plugin '{plugin_id}': {error}")))
}
None => Ok(result),
};
let _ = sender.send(WorkerStreamItem::Result(value));
}
}
WorkerMessage::Event { id, event } => {
if let Some(sender) = pending.lock().await.get(&id) {
let _ = sender.send(WorkerStreamItem::Event(event));
}
}
WorkerMessage::HostCall {
id,
request_id,
method,
params,
} => {
let host = host.clone();
let stdin = stdin.clone();
tokio::spawn(async move {
let result = host.call(&request_id, &method, params).await;
match result {
Ok(result) => {
let _ = write_message(
&stdin,
&HostMessage::HostResult {
id: &id,
result: &result,
},
)
.await;
}
Err(error) => {
let text = error.to_string();
let _ = write_message(
&stdin,
&HostMessage::HostError {
id: &id,
error: &text,
},
)
.await;
}
}
});
}
}
}
fail_pending(&pending, &format!("plugin '{plugin_id}' worker exited")).await;
});
}
fn spawn_stderr_reader(plugin_id: String, stderr: tokio::process::ChildStderr) {
tokio::spawn(async move {
let mut lines = BufReader::new(stderr).lines();
while let Ok(Some(line)) = lines.next_line().await {
tracing::warn!(plugin = %plugin_id, message = %line, "plugin worker stderr");
}
});
}
async fn write_message(stdin: &Arc<Mutex<ChildStdin>>, message: &HostMessage<'_>) -> Result<()> {
let mut bytes = serde_json::to_vec(message)?;
bytes.push(b'\n');
let mut stdin = stdin.lock().await;
stdin.write_all(&bytes).await?;
stdin.flush().await?;
Ok(())
}
async fn fail_pending(pending: &Pending, message: &str) {
for (_, sender) in std::mem::take(&mut *pending.lock().await) {
let _ = sender.send(WorkerStreamItem::Result(Err(Error::Provider(
message.into(),
))));
}
}
impl HostContext {
async fn call(
&self,
request_id: &str,
method: &str,
params: serde_json::Value,
) -> Result<serde_json::Value> {
match method {
"network.fetch" => self.fetch(request_id, params).await,
"network.stream.open" => self.stream_open(request_id, params).await,
"network.stream.read" => self.stream_read(params).await,
"network.stream.close" => {
self.streams
.lock()
.await
.remove(required_string(&params, "streamId")?);
Ok(serde_json::Value::Null)
}
_ => Err(Error::Protocol(format!(
"unsupported plugin host method: {method}"
))),
}
}
async fn request(
&self,
request_id: &str,
params: &serde_json::Value,
) -> Result<(reqwest::RequestBuilder, CancellationToken)> {
let raw_url = required_string(params, "url")?;
let url = url::Url::parse(raw_url)
.map_err(|error| Error::Config(format!("invalid plugin network URL: {error}")))?;
if url.scheme() != "https" || !url.username().is_empty() || url.password().is_some() {
return Err(Error::Config(
"plugin network URL must be HTTPS without credentials".into(),
));
}
let host = url
.host_str()
.ok_or_else(|| Error::Config("plugin network URL has no host".into()))?
.to_ascii_lowercase();
if !self.network_hosts.contains(&host) {
return Err(Error::Config(format!(
"plugin '{}' cannot access host '{host}'",
self.plugin_id
)));
}
let method = params
.get("method")
.and_then(serde_json::Value::as_str)
.unwrap_or("GET")
.parse::<reqwest::Method>()
.map_err(|error| Error::Config(format!("invalid plugin HTTP method: {error}")))?;
let client = crate::network::client_builder(&self.store)
.await?
.redirect(reqwest::redirect::Policy::none())
.connect_timeout(Duration::from_secs(30))
.build()?;
let mut request = client.request(method, url);
if let Some(headers) = params.get("headers").and_then(serde_json::Value::as_object) {
for (name, value) in headers {
let value = value.as_str().ok_or_else(|| {
Error::Config(format!("plugin HTTP header '{name}' must be a string"))
})?;
request = request.header(name, value);
}
}
if let Some(body) = params.get("body").and_then(serde_json::Value::as_str) {
request = request.body(body.to_owned());
}
let cancellation = self
.cancellations
.lock()
.await
.get(request_id)
.cloned()
.unwrap_or_default();
Ok((request, cancellation))
}
async fn fetch(
&self,
request_id: &str,
params: serde_json::Value,
) -> Result<serde_json::Value> {
let (request, cancellation) = self.request(request_id, &params).await?;
let request = request.timeout(Duration::from_secs(60));
let response = tokio::select! {
_ = cancellation.cancelled() => return Err(Error::Cancelled),
response = request.send() => response?,
};
let status = response.status().as_u16();
if response
.content_length()
.is_some_and(|size| size > MAX_NETWORK_RESPONSE_BYTES)
{
return Err(Error::Provider(
"plugin network response is larger than allowed".into(),
));
}
let headers = header_map(&response);
let body = tokio::select! {
_ = cancellation.cancelled() => return Err(Error::Cancelled),
body = response.bytes() => body?,
};
if body.len() as u64 > MAX_NETWORK_RESPONSE_BYTES {
return Err(Error::Provider(
"plugin network response is larger than allowed".into(),
));
}
Ok(
serde_json::json!({ "status": status, "headers": headers, "body": String::from_utf8_lossy(&body) }),
)
}
/// 打开流式响应:立即返回状态与响应头,响应体按行经 stream.read 拉取。
async fn stream_open(
&self,
request_id: &str,
params: serde_json::Value,
) -> Result<serde_json::Value> {
let (request, cancellation) = self.request(request_id, &params).await?;
let response = tokio::select! {
_ = cancellation.cancelled() => return Err(Error::Cancelled),
response = request.send() => response?,
};
let status = response.status().as_u16();
let headers = header_map(&response);
let (sender, receiver) = mpsc::channel::<Result<String>>(256);
tokio::spawn(async move {
use futures_util::StreamExt;
let mut body = response.bytes_stream();
let mut buffered = Vec::<u8>::new();
let mut total = 0_u64;
loop {
let chunk = tokio::select! {
_ = cancellation.cancelled() => {
let _ = sender.send(Err(Error::Cancelled)).await;
return;
}
chunk = body.next() => chunk,
};
let Some(chunk) = chunk else { break };
let chunk = match chunk {
Ok(chunk) => chunk,
Err(error) => {
let _ = sender.send(Err(Error::from(error))).await;
return;
}
};
total += chunk.len() as u64;
if total > MAX_STREAM_BYTES {
let _ = sender
.send(Err(Error::Provider(
"plugin network stream is larger than allowed".into(),
)))
.await;
return;
}
buffered.extend_from_slice(&chunk);
while let Some(position) = buffered.iter().position(|byte| *byte == b'\n') {
let mut line = buffered.drain(..=position).collect::<Vec<u8>>();
line.pop();
if line.last() == Some(&b'\r') {
line.pop();
}
if sender
.send(Ok(String::from_utf8_lossy(&line).into_owned()))
.await
.is_err()
{
return;
}
}
}
if !buffered.is_empty() {
let _ = sender
.send(Ok(String::from_utf8_lossy(&buffered).into_owned()))
.await;
}
});
let stream_id = uuid::Uuid::new_v4().to_string();
self.streams
.lock()
.await
.insert(stream_id.clone(), Arc::new(Mutex::new(receiver)));
Ok(serde_json::json!({
"streamId": stream_id,
"status": status,
"headers": headers,
}))
}
async fn stream_read(&self, params: serde_json::Value) -> Result<serde_json::Value> {
let stream_id = required_string(&params, "streamId")?;
let lines_handle = self
.streams
.lock()
.await
.get(stream_id)
.cloned()
.ok_or_else(|| Error::Protocol(format!("unknown plugin stream: {stream_id}")))?;
let mut receiver = lines_handle.lock().await;
let mut lines = Vec::new();
match receiver.recv().await {
Some(Ok(line)) => lines.push(line),
Some(Err(error)) => {
drop(receiver);
self.streams.lock().await.remove(stream_id);
return Err(error);
}
None => {
drop(receiver);
self.streams.lock().await.remove(stream_id);
return Ok(serde_json::json!({ "lines": [], "done": true }));
}
}
// 把已就绪的行一并带走,减少往返。
while lines.len() < 256 {
match receiver.try_recv() {
Ok(Ok(line)) => lines.push(line),
Ok(Err(error)) => {
drop(receiver);
self.streams.lock().await.remove(stream_id);
return Err(error);
}
Err(_) => break,
}
}
Ok(serde_json::json!({ "lines": lines, "done": false }))
}
}
fn header_map(response: &reqwest::Response) -> std::collections::BTreeMap<String, String> {
response
.headers()
.iter()
.filter_map(|(name, value)| {
value
.to_str()
.ok()
.map(|value| (name.to_string(), value.to_string()))
})
.collect()
}
fn required_string<'a>(params: &'a serde_json::Value, key: &str) -> Result<&'a str> {
params
.get(key)
.and_then(serde_json::Value::as_str)
.ok_or_else(|| Error::Protocol(format!("plugin host call requires string '{key}'")))
}
+1 -1
View File
@@ -96,7 +96,7 @@ impl Provider for AnthropicProvider {
.header("x-api-key", &config.api_key).header("anthropic-version", "2023-06-01")
.headers(config.custom_headers.clone())
.json(&body),
RetryPolicy::default(),
RetryPolicy { retries: config.retry_count, ..RetryPolicy::default() },
&cancellation,
recorder.as_ref(),
request_headers,
+13
View File
@@ -60,6 +60,19 @@ fn merge_extra_params(body: &mut serde_json::Value, extra: &serde_json::Value) -
Ok(())
}
fn apply_body_allowlist(
body: &mut serde_json::Value,
allowed: Option<&std::collections::HashSet<String>>,
) -> Result<()> {
let Some(allowed) = allowed else {
return Ok(());
};
body.as_object_mut()
.ok_or_else(|| crate::Error::Provider("provider request body must be an object".into()))?
.retain(|name, _| allowed.contains(name));
Ok(())
}
fn apply_openai_prompt_cache_key(body: &mut serde_json::Value, model_id: &str) -> Result<()> {
if !model_id.to_ascii_lowercase().contains("gpt") {
return Ok(());
+3 -2
View File
@@ -17,7 +17,7 @@ use crate::{
};
use super::{
apply_openai_prompt_cache_key, merge_extra_params,
apply_body_allowlist, apply_openai_prompt_cache_key, merge_extra_params,
recorder::recorded_headers,
retry::{send_with_retry, Attempt, RetryPolicy},
CallRecorder, FinishReason, ModelEvent, Provider, ProviderStream,
@@ -86,6 +86,7 @@ impl Provider for OpenAiChatProvider {
apply_model(&mut body, &request.model, config.max_output_tokens)?;
merge_extra_params(&mut body, &request.model.extra_params)?;
apply_openai_prompt_cache_key(&mut body, &request.model.model_id)?;
apply_body_allowlist(&mut body, config.allowed_body_fields.as_ref())?;
let request_headers = recorded_headers(&config, &[("content-type", "application/json")]);
if let Some(recorder) = &recorder {
recorder.request(request_headers.clone(), &body).await?;
@@ -94,7 +95,7 @@ impl Provider for OpenAiChatProvider {
"OpenAI Chat",
|| client.post(&config.request_url)
.bearer_auth(&config.api_key).headers(config.custom_headers.clone()).json(&body),
RetryPolicy::default(),
RetryPolicy { retries: config.retry_count, ..RetryPolicy::default() },
&cancellation,
recorder.as_ref(),
request_headers,
+3 -2
View File
@@ -14,7 +14,7 @@ use crate::{
};
use super::{
apply_openai_prompt_cache_key, merge_extra_params,
apply_body_allowlist, apply_openai_prompt_cache_key, merge_extra_params,
recorder::recorded_headers,
retry::{send_with_retry, Attempt, RetryPolicy},
CallRecorder, FinishReason, ModelEvent, Provider, ProviderStream,
@@ -83,6 +83,7 @@ impl Provider for OpenAiResponsesProvider {
apply_model(&mut body, &request.model, config.max_output_tokens)?;
merge_extra_params(&mut body, &request.model.extra_params)?;
apply_openai_prompt_cache_key(&mut body, &request.model.model_id)?;
apply_body_allowlist(&mut body, config.allowed_body_fields.as_ref())?;
let request_headers = recorded_headers(&config, &[("content-type", "application/json")]);
if let Some(recorder) = &recorder {
recorder.request(request_headers.clone(), &body).await?;
@@ -91,7 +92,7 @@ impl Provider for OpenAiResponsesProvider {
"OpenAI Responses",
|| client.post(&config.request_url)
.bearer_auth(&config.api_key).headers(config.custom_headers.clone()).json(&body),
RetryPolicy::default(),
RetryPolicy { retries: config.retry_count, ..RetryPolicy::default() },
&cancellation,
recorder.as_ref(),
request_headers,
-3
View File
@@ -60,9 +60,6 @@ where
String::from_utf8_lossy(&bytes)
));
if attempt == policy.retries {
if let Some(recorder) = recorder {
recorder.failed(&error).await?;
}
return Err(error);
}
tracing::warn!(
+133 -123
View File
@@ -1,4 +1,4 @@
//! Routes model requests to the configured provider.
//! Routes model requests to built-in configurations or stable plugin model IDs.
use std::{sync::Arc, time::Duration};
use async_stream::try_stream;
@@ -8,6 +8,7 @@ use tokio_util::sync::CancellationToken;
use crate::{
config::{ProviderConfig, ProviderKind},
model::{ModelInvocation, ModelLatency, NewLlmCall, ProviderType},
plugin::{PluginRegistry, ADAPTER_ID_PREFIX},
store::Store,
Error, Result,
};
@@ -17,15 +18,19 @@ use super::{
OpenAiResponsesProvider, Provider, ProviderStream,
};
const BUILTIN_PROVIDER_RETRIES: u32 = 5;
pub struct ProviderRouter {
store: Store,
plugins: PluginRegistry,
request_timeout: Duration,
}
impl ProviderRouter {
pub fn new(store: Store, request_timeout: Duration) -> Self {
pub fn new(store: Store, plugins: PluginRegistry, request_timeout: Duration) -> Self {
Self {
store,
plugins,
request_timeout,
}
}
@@ -34,141 +39,146 @@ impl ProviderRouter {
impl Provider for ProviderRouter {
fn stream(
&self,
mut invocation: ModelInvocation,
invocation: ModelInvocation,
cancellation: CancellationToken,
) -> ProviderStream {
let store = self.store.clone();
let plugins = self.plugins.clone();
let request_timeout = self.request_timeout;
Box::pin(try_stream! {
let selected = invocation.request.model.model_id.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 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)?
} 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 stream_started = std::time::Instant::now();
tracing::debug!(
model = %selected,
provider_type = ?provider_type,
timeout_ms = config.request_timeout.as_millis() as u64,
"provider stream created"
);
let mut last_event_time = std::time::Instant::now();
let mut event_count: u64 = 0;
while let Some(event) = stream.next().await {
let now = std::time::Instant::now();
let gap_ms = now.duration_since(last_event_time).as_millis() as u64;
let elapsed_ms = now.duration_since(stream_started).as_millis() as u64;
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_count,
"slow gap detected between provider events"
);
}
recorder.event(&event).await?;
last_event_time = now;
yield event;
}
Err(error) => {
tracing::debug!(
error = %error,
elapsed_ms,
gap_ms,
event_count,
"provider stream error"
);
recorder.failed(&error).await?;
Err(error)?;
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)?; }
}
}
}
if !recorder.is_finished() {
let elapsed_ms = stream_started.elapsed().as_millis() as u64;
if stream_cancellation.is_cancelled() {
tracing::debug!(elapsed_ms, event_count, "provider stream ended after cancellation");
recorder.cancelled().await?;
} else {
let error = Error::Provider("provider stream ended without Done".into());
tracing::warn!(
elapsed_ms,
event_count,
"provider stream ended without Done"
);
recorder.failed(&error).await?;
Err(error)?;
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?;
}
})
}
}
async fn start_recorder(
store: &Store,
invocation: &ModelInvocation,
model_hash: &str,
display_name: &str,
provider_type: ProviderType,
request_url: &str,
model_id: &str,
) -> Result<CallRecorder> {
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_hash.into(),
provider_type,
provider_url: request_url.into(),
request_type: provider_type,
request_url: request_url.into(),
model_id: model_id.into(),
display_name: display_name.into(),
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
}
async fn finish_stream(recorder: &CallRecorder, cancellation: &CancellationToken) -> Result<()> {
if recorder.is_finished() {
return Ok(());
}
if cancellation.is_cancelled() {
recorder.cancelled().await
} else {
let error = Error::Provider("provider stream ended without Done".into());
recorder.failed(&error).await?;
Err(error)
}
}
/// 插件模型的 Provider 实现;对路由与规范化层完全等同于内置 Provider。
struct PluginModelProvider {
registry: PluginRegistry,
}
impl Provider for PluginModelProvider {
fn stream(
&self,
invocation: ModelInvocation,
cancellation: CancellationToken,
) -> ProviderStream {
self.registry.stream_model(invocation, cancellation)
}
}
fn provider_kind(provider_type: ProviderType) -> ProviderKind {
match provider_type {
ProviderType::OpenAiChat => ProviderKind::OpenAiChat,
ProviderType::OpenAiResponses => ProviderKind::OpenAiResponses,
ProviderType::Anthropic => ProviderKind::Anthropic,
// 内置模型的 provider_type 只来自 ModelType,不可能是插件。
ProviderType::Plugin => unreachable!("plugin models never use built-in provider configs"),
}
}
fn custom_headers(value: &serde_json::Value) -> Result<reqwest::header::HeaderMap> {
let object = value
.as_object()
+49
View File
@@ -412,3 +412,52 @@ fn summary_from_row(row: sqlx::sqlite::SqliteRow) -> Result<LlmCallSummary> {
detailed: row.try_get("detailed")?,
})
}
#[cfg(test)]
mod tests {
use super::*;
/// 插件模型不在 model_configs 中,调用记录必须照常落库并可按其稳定 ID 筛选。
#[tokio::test]
async fn plugin_calls_record_without_a_model_config_row() {
let directory = tempfile::tempdir().unwrap();
let store = Store::connect(&format!(
"sqlite://{}",
directory.path().join("test.db").display()
))
.await
.unwrap();
let plugin_model = "plugin:dev.example/codex/gpt-test";
store
.start_llm_call(&NewLlmCall {
call_id: "plugin-call".into(),
run_id: "run".into(),
conversation_id: "conversation".into(),
provider_call_index: 0,
model_hash: plugin_model.into(),
provider_type: ProviderType::Plugin,
provider_url: "plugin://dev.example/codex".into(),
request_type: ProviderType::Plugin,
request_url: "plugin://dev.example/codex".into(),
model_id: "gpt-test".into(),
display_name: "GPT Test".into(),
reasoning_effort: None,
fast: false,
message_count: 1,
tool_count: 0,
detailed: false,
})
.await
.unwrap();
store
.finish_llm_call("plugin-call", "completed", None, 10, None, None)
.await
.unwrap();
let overview = store
.overview(None, None, Some(&format!("[\"{plugin_model}\"]")))
.await
.unwrap();
assert_eq!(overview.metrics.llm_calls, 1);
assert_eq!(overview.metrics.successful_calls, 1);
}
}
+1 -1
View File
@@ -446,7 +446,7 @@ mod tests {
.unwrap();
assert_eq!(checksum_after, checksum_before);
assert_eq!(versions, vec![1, 2, 3, 4, 5, 6]);
assert_eq!(versions, vec![1, 2, 3, 4, 5, 6, 7]);
assert_eq!(checkpoint_table_exists, 1);
}
}