mirror of
https://wget.la/https://github.com/leookun/cursor-byok
synced 2026-10-06 05:12:03 +08:00
feat: plugin system
This commit is contained in:
@@ -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
@@ -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 {
|
||||
|
||||
@@ -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)]
|
||||
|
||||
@@ -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),
|
||||
|
||||
@@ -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
@@ -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();
|
||||
|
||||
@@ -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(),
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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}"))),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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()),
|
||||
|
||||
@@ -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(())
|
||||
}
|
||||
@@ -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());
|
||||
}
|
||||
}
|
||||
@@ -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());
|
||||
}
|
||||
}
|
||||
@@ -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());
|
||||
}
|
||||
}
|
||||
@@ -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);
|
||||
}
|
||||
}
|
||||
@@ -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,
|
||||
|
||||
@@ -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());
|
||||
}
|
||||
}
|
||||
@@ -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;
|
||||
|
||||
@@ -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"),
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
))
|
||||
})
|
||||
}
|
||||
@@ -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();
|
||||
|
||||
@@ -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())));
|
||||
@@ -0,0 +1,5 @@
|
||||
{
|
||||
"fmt": {
|
||||
"lineWidth": 200
|
||||
}
|
||||
}
|
||||
@@ -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"
|
||||
}
|
||||
}
|
||||
@@ -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[]>;
|
||||
};
|
||||
@@ -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");
|
||||
}
|
||||
}
|
||||
@@ -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>;
|
||||
};
|
||||
@@ -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>;
|
||||
};
|
||||
@@ -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));
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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");
|
||||
}
|
||||
}
|
||||
@@ -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());
|
||||
}
|
||||
}
|
||||
@@ -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: ¶ms,
|
||||
},
|
||||
)
|
||||
.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(¶ms, "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, ¶ms).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, ¶ms).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(¶ms, "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}'")))
|
||||
}
|
||||
@@ -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,
|
||||
|
||||
@@ -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(());
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
@@ -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()
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user