mirror of
https://wget.la/https://github.com/leookun/cursor-byok
synced 2026-10-03 18:23:51 +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,
|
||||
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),
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
},
|
||||
]
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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(®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>> {
|
||||
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,
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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");
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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")
|
||||
|
||||
Reference in New Issue
Block a user