feat(cursor): cli 接入本地模型路由与配置

- feat(服务配置): 本地响应配置并禁用 HTTP/2 传输
- feat(模型目录): 提供默认模型接口及本地路由凭据
- fix(待办状态): 基于检查点合并增量待办并保留未变更项
- fix(模型参数): 忽略未知参数以兼容新版 Cursor 请求
- test(本地路由): 覆盖模型元数据路由与凭据行为
This commit is contained in:
leokun
2026-09-03 16:01:25 +08:00
parent 543f618fee
commit 924b5e5926
10 changed files with 489 additions and 11 deletions
+27 -1
View File
@@ -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),
+55 -3
View File
@@ -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<BlobId>, Option<BlobId>)> {
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::<Result<Vec<_>>>()?
} 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<serde_json::Value> {
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(
+24 -5
View File
@@ -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<bo
))),
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn ignores_unknown_cursor_model_parameters() {
let requested = pb::RequestedModel {
model_id: "test-model".into(),
parameters: vec![pb::requested_model::ModelParameterValue {
id: "optimize_for".into(),
value: "quality".into(),
}],
..Default::default()
};
let model = from_requested(&requested, None).expect("unknown parameter should be ignored");
assert_eq!(model.model_id, "test-model");
assert_eq!(model.latency, ModelLatency::Standard);
assert!(!model.reasoning.enabled);
}
}
+86 -1
View File
@@ -11,7 +11,13 @@ pub struct DerivedState {
}
pub fn fold_derived_state(messages: &[CanonicalMessage]) -> 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::<String, (String, Value)>::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<CanonicalMessage> {
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,
},
]
}
}
+1
View File
@@ -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;
+216
View File
@@ -187,6 +187,32 @@ struct UsableModelsAddition {
models: Vec<agent::ModelDetails>,
}
#[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<String>,
#[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<TransportRegistry>,
) -> Result<Response<Body>> {
let models = registry.store().models().await?;
let plugin_models = configured_plugin_models(&registry).await;
Ok(local_response(
agent::GetDefaultModelForCliResponse {
model: default_model_details(&models, &plugin_models),
}
.encode_to_vec(),
))
}
pub async fn default_model(State(registry): State<TransportRegistry>) -> Result<Response<Body>> {
let models = registry.store().models().await?;
let plugin_models = configured_plugin_models(&registry).await;
Ok(local_response(
default_model_response(&models, &plugin_models).encode_to_vec(),
))
}
pub async fn default_model_nudge(
State(registry): State<TransportRegistry>,
) -> Result<Response<Body>> {
let models = registry.store().models().await?;
let plugin_models = configured_plugin_models(&registry).await;
Ok(local_response(
default_model_nudge_response(&models, &plugin_models).encode_to_vec(),
))
}
async fn configured_plugin_models(registry: &TransportRegistry) -> Vec<PluginModelDescriptor> {
match registry.plugins() {
Some(plugins) => plugins.configured_models().await,
None => Vec::new(),
}
}
fn default_model_details(
models: &[ModelConfig],
plugin_models: &[PluginModelDescriptor],
) -> Option<agent::ModelDetails> {
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<u8>) -> Result<Response<Body>> {
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"]
);
}
}
@@ -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<bool>,
}
pub async fn get() -> Result<Response<Body>> {
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));
}
}
@@ -459,6 +459,7 @@ mod tests {
name: "ReadLints".into(),
arguments_text: String::new(),
arguments: json!({ "paths": paths }),
argument_error: None,
}
}
+25
View File
@@ -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");
}
}
}
+3 -1
View File
@@ -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")