From 924b5e59267c1e3264e5e3908a3266b3c40bb835 Mon Sep 17 00:00:00 2001 From: leokun Date: Thu, 3 Sep 2026 16:01:25 +0800 Subject: [PATCH] =?UTF-8?q?feat(cursor):=20cli=20=E6=8E=A5=E5=85=A5?= =?UTF-8?q?=E6=9C=AC=E5=9C=B0=E6=A8=A1=E5=9E=8B=E8=B7=AF=E7=94=B1=E4=B8=8E?= =?UTF-8?q?=E9=85=8D=E7=BD=AE?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - feat(服务配置): 本地响应配置并禁用 HTTP/2 传输 - feat(模型目录): 提供默认模型接口及本地路由凭据 - fix(待办状态): 基于检查点合并增量待办并保留未变更项 - fix(模型参数): 忽略未知参数以兼容新版 Cursor 请求 - test(本地路由): 覆盖模型元数据路由与凭据行为 --- server/src/api/cursor/handlers.rs | 28 ++- server/src/cursor/checkpoint/derived.rs | 58 ++++- server/src/cursor/compile/model.rs | 29 ++- server/src/cursor/prompting/derived_state.rs | 87 ++++++- server/src/cursor/services/mod.rs | 1 + server/src/cursor/services/model_catalog.rs | 216 ++++++++++++++++++ server/src/cursor/services/server_config.rs | 51 +++++ .../tools/tool_call_result/exec/output.rs | 1 + server/src/local_app/proxy.rs | 25 ++ server/tests/connect_wire.rs | 4 +- 10 files changed, 489 insertions(+), 11 deletions(-) create mode 100644 server/src/cursor/services/server_config.rs diff --git a/server/src/api/cursor/handlers.rs b/server/src/api/cursor/handlers.rs index bfc9ad0..d2def44 100644 --- a/server/src/api/cursor/handlers.rs +++ b/server/src/api/cursor/handlers.rs @@ -19,7 +19,9 @@ use crate::{ connect, proto::{agent::v1 as agent, aiserver::v1 as ai}, }, - services::{account, analytics, commit_message, knowledge, model_catalog, tab}, + services::{ + account, analytics, commit_message, knowledge, model_catalog, server_config, tab, + }, transport::{TransportParent, TransportRegistry}, }, Result, @@ -44,6 +46,14 @@ fn router_with_proxy( .route("/__byok-api__/healthz", get(health)) .route("/agent.v1.AgentService/RunSSE", post(run_sse_handler)) .route("/aiserver.v1.BidiService/BidiAppend", post(bidi_handler)) + .route( + "/aiserver.v1.AiService/GetServerConfig", + post(server_config::get), + ) + .route( + "/aiserver.v1.ServerConfigService/GetServerConfig", + post(server_config::get), + ) .route( "/aiserver.v1.AiService/AvailableModels", post(model_catalog::available_models), @@ -64,6 +74,22 @@ fn router_with_proxy( "/aiserver.v1.NetworkService/IsConnected", post(is_connected), ) + .route( + "/agent.v1.AgentService/GetDefaultModelForCli", + post(model_catalog::default_model_for_cli), + ) + .route( + "/aiserver.v1.AiService/GetDefaultModelForCli", + post(model_catalog::default_model_for_cli), + ) + .route( + "/aiserver.v1.AiService/GetDefaultModel", + post(model_catalog::default_model), + ) + .route( + "/aiserver.v1.AiService/GetDefaultModelNudgeData", + post(model_catalog::default_model_nudge), + ) .route( "/aiserver.v1.AuthService/GetEmail", post(account::get_email), diff --git a/server/src/cursor/checkpoint/derived.rs b/server/src/cursor/checkpoint/derived.rs index 285061a..dfd0151 100644 --- a/server/src/cursor/checkpoint/derived.rs +++ b/server/src/cursor/checkpoint/derived.rs @@ -4,7 +4,10 @@ use std::collections::HashMap; use prost::Message; use crate::{ - cursor::{prompting::fold_derived_state, protocol::proto::agent::v1 as pb}, + cursor::{ + prompting::{fold_derived_state, fold_derived_state_from, DerivedState}, + protocol::proto::agent::v1 as pb, + }, model::{CanonicalMessage, MessageContent}, store::BlobId, Error, Result, @@ -17,7 +20,18 @@ impl CheckpointBuilder { &self, messages: &[CanonicalMessage], ) -> Result<(Vec, Option)> { - let state = fold_derived_state(messages); + let changes = fold_derived_state(messages); + let state = if changes.todos.is_some() && !self.base.todos.is_empty() { + fold_derived_state_from( + messages, + DerivedState { + todos: Some(self.base_todo_state().await?), + plan: None, + }, + ) + } else { + changes + }; let todo_values = state .todos .as_ref() @@ -28,7 +42,15 @@ impl CheckpointBuilder { .ok_or_else(|| Error::Protocol("TodoWrite state is missing todos[]".into())) }) .transpose()?; - let mut todo_ids = Vec::new(); + let mut todo_ids = if todo_values.is_none() { + self.base + .todos + .iter() + .map(|raw| BlobId::from_bytes(raw)) + .collect::>>()? + } else { + Vec::new() + }; for (index, todo) in todo_values.into_iter().flatten().enumerate() { let status = match todo .get("status") @@ -97,6 +119,36 @@ impl CheckpointBuilder { }; Ok((todo_ids, plan_id)) } + + async fn base_todo_state(&self) -> Result { + let mut todos = Vec::with_capacity(self.base.todos.len()); + for raw_id in &self.base.todos { + let id = BlobId::from_bytes(raw_id)?; + let data = self.sync.get(&id).await?.ok_or_else(|| { + Error::Protocol(format!("Cursor Todo Blob is missing: {}", id.to_base64())) + })?; + let todo = pb::TodoItem::decode(data.as_slice())?; + let status = match pb::TodoStatus::try_from(todo.status) { + Ok(pb::TodoStatus::InProgress) => "in_progress", + Ok(pb::TodoStatus::Completed) => "completed", + Ok(pb::TodoStatus::Cancelled) => "cancelled", + Ok(pb::TodoStatus::Pending) => "pending", + Ok(pb::TodoStatus::Unspecified) | Err(_) => { + return Err(Error::Protocol(format!( + "unknown Cursor Todo status: {}", + todo.status + ))) + } + }; + todos.push(serde_json::json!({ + "id": todo.id, + "content": todo.content, + "status": status, + "dependencies": todo.dependencies, + })); + } + Ok(serde_json::json!({"merge": false, "todos": todos})) + } } pub(super) fn update_current_step_state( diff --git a/server/src/cursor/compile/model.rs b/server/src/cursor/compile/model.rs index 50133bf..ac5f97b 100644 --- a/server/src/cursor/compile/model.rs +++ b/server/src/cursor/compile/model.rs @@ -119,11 +119,7 @@ fn from_requested( )) })?); } - other => { - return Err(Error::Protocol(format!( - "unsupported Cursor model parameter: {other}" - ))) - } + _ => {} } } Ok(spec) @@ -139,3 +135,26 @@ fn parse_bool(parameter: &pb::requested_model::ModelParameterValue) -> Result DerivedState { - let mut state = DerivedState::default(); + fold_derived_state_from(messages, DerivedState::default()) +} + +pub fn fold_derived_state_from( + messages: &[CanonicalMessage], + mut state: DerivedState, +) -> DerivedState { let mut calls = std::collections::HashMap::::new(); for message in messages { match &message.content { @@ -81,3 +87,82 @@ fn normalize(value: &str) -> String { .flat_map(char::to_lowercase) .collect() } + +#[cfg(test)] +mod tests { + use serde_json::{json, Value}; + + use super::{fold_derived_state_from, DerivedState}; + use crate::model::{ + CanonicalMessage, MessageContent, Origin, Role, ToolCallContent, ToolResultContent, + }; + + #[test] + fn merge_patch_inherits_content_from_checkpoint_todo_state() { + let messages = todo_write_messages(json!({ + "merge": true, + "todos": [{"id": "tests", "status": "completed"}], + })); + let initial = DerivedState { + todos: Some(json!({ + "merge": false, + "todos": [{ + "id": "tests", + "content": "Run focused tests", + "status": "in_progress", + }], + })), + plan: None, + }; + + let state = fold_derived_state_from(&messages, initial); + assert_eq!( + state.todos, + Some(json!({ + "merge": false, + "todos": [{ + "id": "tests", + "content": "Run focused tests", + "status": "completed", + }], + })) + ); + } + + fn todo_write_messages(arguments: Value) -> Vec { + vec![ + CanonicalMessage { + message_id: "assistant".into(), + role: Role::Assistant, + origin: Origin::Assistant, + content: MessageContent::Assistant { + text: String::new(), + thinking: String::new(), + tool_round_id: None, + replay_state: None, + tool_calls: vec![ToolCallContent { + index: 0, + call_id: "todo-call".into(), + name: "TodoWrite".into(), + arguments, + }], + }, + runtime_event_id: None, + }, + CanonicalMessage { + message_id: "result".into(), + role: Role::Tool, + origin: Origin::Tool, + content: MessageContent::ToolResult(ToolResultContent { + call_id: "todo-call".into(), + name: "TodoWrite".into(), + content: "{}".into(), + is_error: false, + image: None, + provider_parts: Vec::new(), + }), + runtime_event_id: None, + }, + ] + } +} diff --git a/server/src/cursor/services/mod.rs b/server/src/cursor/services/mod.rs index 33c87e8..d2952be 100644 --- a/server/src/cursor/services/mod.rs +++ b/server/src/cursor/services/mod.rs @@ -8,5 +8,6 @@ pub mod context_sync; pub mod knowledge; pub mod model_catalog; pub mod observability; +pub mod server_config; pub mod tab; pub mod usage; diff --git a/server/src/cursor/services/model_catalog.rs b/server/src/cursor/services/model_catalog.rs index eb82f87..b63e4f3 100644 --- a/server/src/cursor/services/model_catalog.rs +++ b/server/src/cursor/services/model_catalog.rs @@ -187,6 +187,32 @@ struct UsableModelsAddition { models: Vec, } +#[derive(Clone, PartialEq, Message)] +struct DefaultModelResponse { + #[prost(string, tag = "1")] + model: String, + #[prost(string, tag = "2")] + thinking_model: String, + #[prost(bool, tag = "3")] + max_mode: bool, + #[prost(string, tag = "4")] + next_default_set_date: String, +} + +#[derive(Clone, PartialEq, Message)] +struct DefaultModelNudgeDataResponse { + #[prost(string, tag = "1")] + nudge_date: String, + #[prost(bool, tag = "2")] + should_default_switch_on_new_chat: bool, + #[prost(string, repeated, tag = "3")] + models_with_no_default_switch: Vec, + #[prost(string, tag = "4")] + conversion_model_override: String, +} + +const CLI_LOCAL_MODEL_API_KEY: &str = "cursor-byok-local"; + const CONTEXTS: [(&str, &str); 5] = [ ("200k", "200K"), ("356k", "356K"), @@ -287,6 +313,94 @@ pub async fn usable_models( } } +pub async fn default_model_for_cli( + State(registry): State, +) -> Result> { + let models = registry.store().models().await?; + let plugin_models = configured_plugin_models(®istry).await; + Ok(local_response( + agent::GetDefaultModelForCliResponse { + model: default_model_details(&models, &plugin_models), + } + .encode_to_vec(), + )) +} + +pub async fn default_model(State(registry): State) -> Result> { + let models = registry.store().models().await?; + let plugin_models = configured_plugin_models(®istry).await; + Ok(local_response( + default_model_response(&models, &plugin_models).encode_to_vec(), + )) +} + +pub async fn default_model_nudge( + State(registry): State, +) -> Result> { + let models = registry.store().models().await?; + let plugin_models = configured_plugin_models(®istry).await; + Ok(local_response( + default_model_nudge_response(&models, &plugin_models).encode_to_vec(), + )) +} + +async fn configured_plugin_models(registry: &TransportRegistry) -> Vec { + match registry.plugins() { + Some(plugins) => plugins.configured_models().await, + None => Vec::new(), + } +} + +fn default_model_details( + models: &[ModelConfig], + plugin_models: &[PluginModelDescriptor], +) -> Option { + models + .first() + .map(usable_model) + .or_else(|| plugin_models.first().map(usable_plugin_model)) +} + +fn default_model_id<'a>( + models: &'a [ModelConfig], + plugin_models: &'a [PluginModelDescriptor], +) -> &'a str { + models + .first() + .map(|model| model.model_hash.as_str()) + .or_else(|| plugin_models.first().map(|model| model.id.as_str())) + .unwrap_or_default() +} + +fn default_model_response( + models: &[ModelConfig], + plugin_models: &[PluginModelDescriptor], +) -> DefaultModelResponse { + let model = default_model_id(models, plugin_models).to_owned(); + DefaultModelResponse { + thinking_model: model.clone(), + model, + max_mode: false, + next_default_set_date: String::new(), + } +} + +fn default_model_nudge_response( + models: &[ModelConfig], + plugin_models: &[PluginModelDescriptor], +) -> DefaultModelNudgeDataResponse { + DefaultModelNudgeDataResponse { + nudge_date: "0".into(), + should_default_switch_on_new_chat: false, + models_with_no_default_switch: models + .iter() + .map(|model| model.model_hash.clone()) + .chain(plugin_models.iter().map(|model| model.id.clone())) + .collect(), + conversion_model_override: String::new(), + } +} + fn merge_response(upstream: proxy::BufferedResponse, extra: Vec) -> Result> { if !upstream.status.is_success() { tracing::warn!(status = %upstream.status, "Cursor model catalog upstream rejected request; using local catalog"); @@ -622,6 +736,13 @@ fn available_plugin_model(model: &PluginModelDescriptor) -> AvailableModel { } } +fn cli_local_model_credentials() -> agent::model_details::Credentials { + agent::model_details::Credentials::ApiKeyCredentials(agent::ApiKeyCredentials { + api_key: CLI_LOCAL_MODEL_API_KEY.into(), + base_url: None, + }) +} + fn usable_plugin_model(model: &PluginModelDescriptor) -> agent::ModelDetails { agent::ModelDetails { model_id: model.id.clone(), @@ -629,6 +750,7 @@ fn usable_plugin_model(model: &PluginModelDescriptor) -> agent::ModelDetails { display_name: model.display_name.clone(), display_name_short: model.display_name.clone(), thinking_details: Some(agent::ThinkingDetails::default()), + credentials: Some(cli_local_model_credentials()), ..Default::default() } } @@ -640,6 +762,100 @@ fn usable_model(model: &ModelConfig) -> agent::ModelDetails { display_name: model.display_name.clone(), display_name_short: model.display_name.clone(), thinking_details: Some(agent::ThinkingDetails::default()), + credentials: Some(cli_local_model_credentials()), ..Default::default() } } + +#[cfg(test)] +mod tests { + use super::*; + use crate::model::{ModelType, OPENAI_CHAT_ENDPOINT}; + + fn model() -> ModelConfig { + ModelConfig { + model_hash: "local-model-hash".into(), + sort_order: 0, + display_name: "Local Model".into(), + group_name: None, + model_type: ModelType::OpenAi, + base_url: "https://provider.example/v1/chat/completions".into(), + use_full_url: true, + api_key: "provider-secret".into(), + tooltip_data: "Local Model".into(), + model_id: "upstream-model".into(), + reasoning_effort: None, + openai_endpoint: OPENAI_CHAT_ENDPOINT.into(), + openai_extra_params_enabled: false, + openai_extra_params: serde_json::json!({}), + custom_headers_enabled: false, + custom_headers: serde_json::json!({}), + anthropic_extra_params_enabled: false, + anthropic_extra_params: serde_json::json!({}), + context_window_tokens: None, + max_completion_tokens: None, + anthropic_max_tokens: None, + anthropic_thinking_effort: None, + thinking_budget_tokens: None, + created_at_ms: 0, + updated_at_ms: 0, + } + } + + #[test] + fn cli_model_details_use_local_routing_credentials() { + let details = usable_model(&model()); + assert_eq!(details.model_id, "local-model-hash"); + assert_eq!(details.display_name, "Local Model"); + let agent::model_details::Credentials::ApiKeyCredentials(credentials) = + details.credentials.expect("API credentials") + else { + panic!("expected API key credentials"); + }; + assert_eq!(credentials.api_key, CLI_LOCAL_MODEL_API_KEY); + assert_eq!(credentials.base_url, None); + assert_ne!(credentials.api_key, "provider-secret"); + } + + #[test] + fn cli_plugin_model_details_use_local_routing_credentials() { + let details = usable_plugin_model(&PluginModelDescriptor { + id: "plugin:test/provider/model".into(), + plugin_id: "plugin:test".into(), + plugin_name: "Test Plugin".into(), + provider_id: "provider".into(), + model_id: "model".into(), + display_name: "Plugin Model".into(), + description: None, + icon: String::new(), + provider_type: "test".into(), + max_output_tokens: None, + images: false, + }); + assert_eq!(details.model_id, "plugin:test/provider/model"); + let agent::model_details::Credentials::ApiKeyCredentials(credentials) = + details.credentials.expect("API credentials") + else { + panic!("expected API key credentials"); + }; + assert_eq!(credentials.api_key, CLI_LOCAL_MODEL_API_KEY); + assert_eq!(credentials.base_url, None); + } + + #[test] + fn cli_default_responses_use_the_local_model_hash() { + let models = vec![model()]; + let details = default_model_details(&models, &[]).expect("default model"); + assert_eq!(details.model_id, "local-model-hash"); + + let response = default_model_response(&models, &[]); + assert_eq!(response.model, "local-model-hash"); + assert_eq!(response.thinking_model, "local-model-hash"); + + let nudge = default_model_nudge_response(&models, &[]); + assert_eq!( + nudge.models_with_no_default_switch, + vec!["local-model-hash"] + ); + } +} diff --git a/server/src/cursor/services/server_config.rs b/server/src/cursor/services/server_config.rs new file mode 100644 index 0000000..53cb1e4 --- /dev/null +++ b/server/src/cursor/services/server_config.rs @@ -0,0 +1,51 @@ +//! Keeps Cursor Agent traffic on the endpoint selected with `agent -e`. +use axum::{ + body::Body, + http::{header, HeaderValue, Response, StatusCode}, +}; +use prost::Message; + +use crate::Result; + +const HTTP2_CONFIG_FORCE_ALL_DISABLED: i32 = 1; + +#[derive(Clone, PartialEq, Message)] +struct ServerConfigResponse { + #[prost(string, tag = "6")] + config_version: String, + #[prost(int32, tag = "7")] + http2_config: i32, + #[prost(bool, optional, tag = "28")] + cli_sandbox_default_enabled: Option, +} + +pub async fn get() -> Result> { + let payload = server_config().encode_to_vec(); + let mut response = Response::new(Body::from(payload)); + *response.status_mut() = StatusCode::OK; + response.headers_mut().insert( + header::CONTENT_TYPE, + HeaderValue::from_static("application/proto"), + ); + Ok(response) +} + +fn server_config() -> ServerConfigResponse { + ServerConfigResponse { + config_version: "cursor_byok_local_agent_v1".into(), + http2_config: HTTP2_CONFIG_FORCE_ALL_DISABLED, + cli_sandbox_default_enabled: Some(true), + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn forces_agent_cli_to_use_the_selected_legacy_endpoint() { + let config = server_config(); + assert_eq!(config.http2_config, HTTP2_CONFIG_FORCE_ALL_DISABLED); + assert_eq!(config.cli_sandbox_default_enabled, Some(true)); + } +} diff --git a/server/src/cursor/tools/tool_call_result/exec/output.rs b/server/src/cursor/tools/tool_call_result/exec/output.rs index 316085c..2441a5a 100644 --- a/server/src/cursor/tools/tool_call_result/exec/output.rs +++ b/server/src/cursor/tools/tool_call_result/exec/output.rs @@ -459,6 +459,7 @@ mod tests { name: "ReadLints".into(), arguments_text: String::new(), arguments: json!({ "paths": paths }), + argument_error: None, } } diff --git a/server/src/local_app/proxy.rs b/server/src/local_app/proxy.rs index 2a2eb82..70ba5fc 100644 --- a/server/src/local_app/proxy.rs +++ b/server/src/local_app/proxy.rs @@ -152,9 +152,15 @@ fn is_local_path(path: &str) -> bool { path, "/agent.v1.AgentService/RunSSE" | "/aiserver.v1.BidiService/BidiAppend" + | "/aiserver.v1.AiService/GetServerConfig" + | "/aiserver.v1.ServerConfigService/GetServerConfig" | "/aiserver.v1.AiService/AvailableModels" | "/agent.v1.AgentService/GetUsableModels" | "/aiserver.v1.AiService/GetUsableModels" + | "/agent.v1.AgentService/GetDefaultModelForCli" + | "/aiserver.v1.AiService/GetDefaultModelForCli" + | "/aiserver.v1.AiService/GetDefaultModel" + | "/aiserver.v1.AiService/GetDefaultModelNudgeData" | "/aiserver.v1.AuthService/GetEmail" | "/aiserver.v1.DashboardService/GetMe" | "/aiserver.v1.DashboardService/GetTeams" @@ -175,3 +181,22 @@ fn is_local_path(path: &str) -> bool { fn should_route_locally(path: &str, tab_mode: TabMode) -> bool { is_local_path(path) || (is_tab_path(path) && tab_mode != TabMode::Direct) } + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn cursor_cli_transport_and_model_metadata_routes_stay_local() { + for path in [ + "/aiserver.v1.AiService/GetServerConfig", + "/aiserver.v1.ServerConfigService/GetServerConfig", + "/agent.v1.AgentService/GetDefaultModelForCli", + "/aiserver.v1.AiService/GetDefaultModelForCli", + "/aiserver.v1.AiService/GetDefaultModel", + "/aiserver.v1.AiService/GetDefaultModelNudgeData", + ] { + assert!(is_local_path(path), "{path} must not reach Cursor upstream"); + } + } +} diff --git a/server/tests/connect_wire.rs b/server/tests/connect_wire.rs index a370f36..3a4f0a6 100644 --- a/server/tests/connect_wire.rs +++ b/server/tests/connect_wire.rs @@ -21,6 +21,7 @@ use cursor_server::{ proto::{agent::v1 as pb, aiserver::v1 as ai}, }, cursor::transport::TransportRegistry, + network::NetworkClients, }; use flate2::{write::GzEncoder, Compression}; use prost::Message; @@ -102,6 +103,7 @@ async fn bidi_append_gzip_body_is_decompressed_before_protobuf_decode() { .as_path(), ) .unwrap(); + let clients = NetworkClients::new(store.clone()); let registry = TransportRegistry::new( store, Arc::new(fake_provider::FakeProvider::default()), @@ -118,7 +120,7 @@ async fn bidi_append_gzip_body_is_decompressed_before_protobuf_decode() { encoder.write_all(&wire).unwrap(); let compressed = encoder.finish().unwrap(); - let response = cursor::router(registry) + let response = cursor::router(registry, clients) .unwrap() .oneshot( Request::post("/aiserver.v1.BidiService/BidiAppend")