mirror of
https://wget.la/https://github.com/leookun/cursor-byok
synced 2026-10-08 15:43:10 +08:00
feat(cursor): cli 接入本地模型路由与配置
- feat(服务配置): 本地响应配置并禁用 HTTP/2 传输 - feat(模型目录): 提供默认模型接口及本地路由凭据 - fix(待办状态): 基于检查点合并增量待办并保留未变更项 - fix(模型参数): 忽略未知参数以兼容新版 Cursor 请求 - test(本地路由): 覆盖模型元数据路由与凭据行为
This commit is contained in:
@@ -19,7 +19,9 @@ use crate::{
|
|||||||
connect,
|
connect,
|
||||||
proto::{agent::v1 as agent, aiserver::v1 as ai},
|
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},
|
transport::{TransportParent, TransportRegistry},
|
||||||
},
|
},
|
||||||
Result,
|
Result,
|
||||||
@@ -44,6 +46,14 @@ fn router_with_proxy(
|
|||||||
.route("/__byok-api__/healthz", get(health))
|
.route("/__byok-api__/healthz", get(health))
|
||||||
.route("/agent.v1.AgentService/RunSSE", post(run_sse_handler))
|
.route("/agent.v1.AgentService/RunSSE", post(run_sse_handler))
|
||||||
.route("/aiserver.v1.BidiService/BidiAppend", post(bidi_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(
|
.route(
|
||||||
"/aiserver.v1.AiService/AvailableModels",
|
"/aiserver.v1.AiService/AvailableModels",
|
||||||
post(model_catalog::available_models),
|
post(model_catalog::available_models),
|
||||||
@@ -64,6 +74,22 @@ fn router_with_proxy(
|
|||||||
"/aiserver.v1.NetworkService/IsConnected",
|
"/aiserver.v1.NetworkService/IsConnected",
|
||||||
post(is_connected),
|
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(
|
.route(
|
||||||
"/aiserver.v1.AuthService/GetEmail",
|
"/aiserver.v1.AuthService/GetEmail",
|
||||||
post(account::get_email),
|
post(account::get_email),
|
||||||
|
|||||||
@@ -4,7 +4,10 @@ use std::collections::HashMap;
|
|||||||
use prost::Message;
|
use prost::Message;
|
||||||
|
|
||||||
use crate::{
|
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},
|
model::{CanonicalMessage, MessageContent},
|
||||||
store::BlobId,
|
store::BlobId,
|
||||||
Error, Result,
|
Error, Result,
|
||||||
@@ -17,7 +20,18 @@ impl CheckpointBuilder {
|
|||||||
&self,
|
&self,
|
||||||
messages: &[CanonicalMessage],
|
messages: &[CanonicalMessage],
|
||||||
) -> Result<(Vec<BlobId>, Option<BlobId>)> {
|
) -> 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
|
let todo_values = state
|
||||||
.todos
|
.todos
|
||||||
.as_ref()
|
.as_ref()
|
||||||
@@ -28,7 +42,15 @@ impl CheckpointBuilder {
|
|||||||
.ok_or_else(|| Error::Protocol("TodoWrite state is missing todos[]".into()))
|
.ok_or_else(|| Error::Protocol("TodoWrite state is missing todos[]".into()))
|
||||||
})
|
})
|
||||||
.transpose()?;
|
.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() {
|
for (index, todo) in todo_values.into_iter().flatten().enumerate() {
|
||||||
let status = match todo
|
let status = match todo
|
||||||
.get("status")
|
.get("status")
|
||||||
@@ -97,6 +119,36 @@ impl CheckpointBuilder {
|
|||||||
};
|
};
|
||||||
Ok((todo_ids, plan_id))
|
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(
|
pub(super) fn update_current_step_state(
|
||||||
|
|||||||
@@ -119,11 +119,7 @@ fn from_requested(
|
|||||||
))
|
))
|
||||||
})?);
|
})?);
|
||||||
}
|
}
|
||||||
other => {
|
_ => {}
|
||||||
return Err(Error::Protocol(format!(
|
|
||||||
"unsupported Cursor model parameter: {other}"
|
|
||||||
)))
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
Ok(spec)
|
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);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -11,7 +11,13 @@ pub struct DerivedState {
|
|||||||
}
|
}
|
||||||
|
|
||||||
pub fn fold_derived_state(messages: &[CanonicalMessage]) -> 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();
|
let mut calls = std::collections::HashMap::<String, (String, Value)>::new();
|
||||||
for message in messages {
|
for message in messages {
|
||||||
match &message.content {
|
match &message.content {
|
||||||
@@ -81,3 +87,82 @@ fn normalize(value: &str) -> String {
|
|||||||
.flat_map(char::to_lowercase)
|
.flat_map(char::to_lowercase)
|
||||||
.collect()
|
.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,
|
||||||
|
},
|
||||||
|
]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -8,5 +8,6 @@ pub mod context_sync;
|
|||||||
pub mod knowledge;
|
pub mod knowledge;
|
||||||
pub mod model_catalog;
|
pub mod model_catalog;
|
||||||
pub mod observability;
|
pub mod observability;
|
||||||
|
pub mod server_config;
|
||||||
pub mod tab;
|
pub mod tab;
|
||||||
pub mod usage;
|
pub mod usage;
|
||||||
|
|||||||
@@ -187,6 +187,32 @@ struct UsableModelsAddition {
|
|||||||
models: Vec<agent::ModelDetails>,
|
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] = [
|
const CONTEXTS: [(&str, &str); 5] = [
|
||||||
("200k", "200K"),
|
("200k", "200K"),
|
||||||
("356k", "356K"),
|
("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(®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<TransportRegistry>) -> Result<Response<Body>> {
|
||||||
|
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<TransportRegistry>,
|
||||||
|
) -> Result<Response<Body>> {
|
||||||
|
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<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>> {
|
fn merge_response(upstream: proxy::BufferedResponse, extra: Vec<u8>) -> Result<Response<Body>> {
|
||||||
if !upstream.status.is_success() {
|
if !upstream.status.is_success() {
|
||||||
tracing::warn!(status = %upstream.status, "Cursor model catalog upstream rejected request; using local catalog");
|
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 {
|
fn usable_plugin_model(model: &PluginModelDescriptor) -> agent::ModelDetails {
|
||||||
agent::ModelDetails {
|
agent::ModelDetails {
|
||||||
model_id: model.id.clone(),
|
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: model.display_name.clone(),
|
||||||
display_name_short: model.display_name.clone(),
|
display_name_short: model.display_name.clone(),
|
||||||
thinking_details: Some(agent::ThinkingDetails::default()),
|
thinking_details: Some(agent::ThinkingDetails::default()),
|
||||||
|
credentials: Some(cli_local_model_credentials()),
|
||||||
..Default::default()
|
..Default::default()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -640,6 +762,100 @@ fn usable_model(model: &ModelConfig) -> agent::ModelDetails {
|
|||||||
display_name: model.display_name.clone(),
|
display_name: model.display_name.clone(),
|
||||||
display_name_short: model.display_name.clone(),
|
display_name_short: model.display_name.clone(),
|
||||||
thinking_details: Some(agent::ThinkingDetails::default()),
|
thinking_details: Some(agent::ThinkingDetails::default()),
|
||||||
|
credentials: Some(cli_local_model_credentials()),
|
||||||
..Default::default()
|
..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(),
|
name: "ReadLints".into(),
|
||||||
arguments_text: String::new(),
|
arguments_text: String::new(),
|
||||||
arguments: json!({ "paths": paths }),
|
arguments: json!({ "paths": paths }),
|
||||||
|
argument_error: None,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -152,9 +152,15 @@ fn is_local_path(path: &str) -> bool {
|
|||||||
path,
|
path,
|
||||||
"/agent.v1.AgentService/RunSSE"
|
"/agent.v1.AgentService/RunSSE"
|
||||||
| "/aiserver.v1.BidiService/BidiAppend"
|
| "/aiserver.v1.BidiService/BidiAppend"
|
||||||
|
| "/aiserver.v1.AiService/GetServerConfig"
|
||||||
|
| "/aiserver.v1.ServerConfigService/GetServerConfig"
|
||||||
| "/aiserver.v1.AiService/AvailableModels"
|
| "/aiserver.v1.AiService/AvailableModels"
|
||||||
| "/agent.v1.AgentService/GetUsableModels"
|
| "/agent.v1.AgentService/GetUsableModels"
|
||||||
| "/aiserver.v1.AiService/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.AuthService/GetEmail"
|
||||||
| "/aiserver.v1.DashboardService/GetMe"
|
| "/aiserver.v1.DashboardService/GetMe"
|
||||||
| "/aiserver.v1.DashboardService/GetTeams"
|
| "/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 {
|
fn should_route_locally(path: &str, tab_mode: TabMode) -> bool {
|
||||||
is_local_path(path) || (is_tab_path(path) && tab_mode != TabMode::Direct)
|
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");
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -21,6 +21,7 @@ use cursor_server::{
|
|||||||
proto::{agent::v1 as pb, aiserver::v1 as ai},
|
proto::{agent::v1 as pb, aiserver::v1 as ai},
|
||||||
},
|
},
|
||||||
cursor::transport::TransportRegistry,
|
cursor::transport::TransportRegistry,
|
||||||
|
network::NetworkClients,
|
||||||
};
|
};
|
||||||
use flate2::{write::GzEncoder, Compression};
|
use flate2::{write::GzEncoder, Compression};
|
||||||
use prost::Message;
|
use prost::Message;
|
||||||
@@ -102,6 +103,7 @@ async fn bidi_append_gzip_body_is_decompressed_before_protobuf_decode() {
|
|||||||
.as_path(),
|
.as_path(),
|
||||||
)
|
)
|
||||||
.unwrap();
|
.unwrap();
|
||||||
|
let clients = NetworkClients::new(store.clone());
|
||||||
let registry = TransportRegistry::new(
|
let registry = TransportRegistry::new(
|
||||||
store,
|
store,
|
||||||
Arc::new(fake_provider::FakeProvider::default()),
|
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();
|
encoder.write_all(&wire).unwrap();
|
||||||
let compressed = encoder.finish().unwrap();
|
let compressed = encoder.finish().unwrap();
|
||||||
|
|
||||||
let response = cursor::router(registry)
|
let response = cursor::router(registry, clients)
|
||||||
.unwrap()
|
.unwrap()
|
||||||
.oneshot(
|
.oneshot(
|
||||||
Request::post("/aiserver.v1.BidiService/BidiAppend")
|
Request::post("/aiserver.v1.BidiService/BidiAppend")
|
||||||
|
|||||||
Reference in New Issue
Block a user