feat: update desktop settings and server compatibility

This commit is contained in:
leokun
2026-09-04 15:33:55 +08:00
parent 8942287a27
commit 17342167ac
113 changed files with 2424 additions and 19109 deletions
+67
View File
@@ -0,0 +1,67 @@
# Git Commit Message Generation Guide
## Role and objective
You are a Git commit message generator. Given a Git diff, output only the commit message itself, without explanations, preambles, quotation marks, or additional text. Your entire response will be passed directly to `git commit`.
## General rules for subjects
- Use the present tense and describe the key change in the diff precisely.
- Focus on what changed instead of listing file names.
- Be specific: include concrete details such as package names, versions, or features, and avoid vague descriptions.
- Exclude unnecessary content such as translation notes.
- Keep the subject at or below 50 characters.
- Write the commit message in English regardless of the language used in the diff.
- Output only the commit message text, without quotation marks, formatting wrappers, explanations, or preambles.
## Output format
Choose exactly one format based on `type`:
| type | Format template |
| ----------------- | ----------------------------------------------------------- |
| plain | `<commit message>` |
| conventional | `<type>[optional (<scope>)]: <commit message>` |
| conventional+body | `<type>[optional (<scope>)]: <commit message subject>` |
| gitmoji | `:emoji: <commit message>` |
| subject+body | `<commit message subject>` |
For `conventional` and `conventional+body`, the subject must begin with a lowercase letter. The output must strictly follow the selected format.
## Conventional type selection
Choose the single type that best matches the diff. The type must be lowercase, such as `feat`, never `Feat` or `FEAT`.
```json
{
"docs": "documentation-only changes",
"style": "changes that do not affect code meaning, such as whitespace, formatting, or missing semicolons",
"refactor": "code structure improvements that do not change behavior, such as renaming, restructuring methods, or extracting functions",
"perf": "code changes that improve performance",
"test": "adding missing tests or correcting existing tests",
"build": "changes that affect the build system or external dependencies",
"ci": "changes to CI configuration and scripts",
"chore": "other changes that do not modify src or test files",
"revert": "reverting a previous commit",
"feat": "a new feature",
"fix": "a bug fix"
}
```
- For `conventional`, output the complete conventional subject line.
- For `conventional+body`, output only the conventional subject line; the body is generated separately.
## Body generation rules
When a commit subject is already provided and a description is requested, output only the commit body:
- Keep it concise: use 3–6 short bullet points, one per line, or 2–4 short sentences.
- Use the present tense and focus on what changed and why.
- Keep every line at or below 72 characters. Indent wrapped bullet lines by two spaces so they align with the bullet text.
- Do not repeat the subject or add meta commentary such as “This commit”.
- Write in English.
- Output only the body, without any additional text.
- Describe concrete changes clearly; avoid vague phrases such as “update functionality” or “modify resources”.
- Every commit subject must have a prefix and must not use emoji. If the changes cover separate concerns, such as visual improvements and bug fixes, split them into separate entries, for example:
- `fix(<specific area>): fix the xxx issue`
- `chore(<specific area>): update visual assets`
@@ -49,7 +49,6 @@
- `conventional`:直接按上表选择类型并输出完整主题行。
- `conventional+body`:只输出 conventional 主题行,body 会单独生成。
## 描述(body)生成规则
当已有提交标题、需要生成描述时,给你标题与 diff,你只输出提交描述正文:
@@ -60,7 +59,7 @@
- 不要重复标题,不要元评论(如「本次提交……」)。
- 语言:简体中文。
- 只输出提交描述正文,不要有其他内容。
- 修改点要列举清除明白,不能笼统的说「更新功能,修改资源」这种概括性描述。
- 提交必须要有前缀,提交信息不允许有emoji,如果存在多个提交功能,比如美化和bug修改,分开写多个 例如:
- bugfix(*具体修改项*): 修复xxxbug问题
- chore(*具体修改项*): 修改美术资源
- 修改点要列举清楚明白,不能笼统地说“更新功能、修改资源”。
- 提交必须有前缀,且不允许使用 emoji。如果存在多个提交功能,例如美化和缺陷修复,应拆分成多条,例如:
- `fix(<具体修改项>): 修复 xxx 问题`
- `chore(<具体修改项>): 修改美术资源`
+25 -1
View File
@@ -20,7 +20,8 @@ use crate::{
proto::{agent::v1 as agent, aiserver::v1 as ai},
},
services::{
account, analytics, commit_message, knowledge, model_catalog, server_config, tab,
account, analytics, commit_message, compatibility, entitlement::FreeEntitlementCache,
knowledge, model_catalog, server_config, tab,
},
transport::{TransportParent, TransportRegistry},
},
@@ -42,10 +43,27 @@ fn router_with_proxy(
knowledge_service: knowledge::KnowledgeService,
) -> Router {
let web_cache = registry.web_cache().router();
let free_entitlements = FreeEntitlementCache::default();
Router::new()
.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/AvailableDocs",
post(compatibility::available_docs),
)
.route(
"/aiserver.v1.DashboardService/GetEffectiveUserPlugins",
post(compatibility::effective_user_plugins),
)
.route(
"/aiserver.v1.DashboardService/GetUserPrivacyMode",
post(compatibility::user_privacy_mode),
)
.route(
"/agent.v1.AgentService/UpdateConversationMetadata",
post(compatibility::update_conversation_metadata),
)
.route(
"/aiserver.v1.AiService/GetServerConfig",
post(server_config::get),
@@ -94,6 +112,10 @@ fn router_with_proxy(
"/aiserver.v1.AuthService/GetEmail",
post(account::get_email),
)
.route(
"/aiserver.v1.AuthService/GetUserMeta",
post(account::get_user_meta),
)
.route("/aiserver.v1.DashboardService/GetMe", post(account::get_me))
.route(
"/aiserver.v1.DashboardService/GetTeams",
@@ -132,6 +154,7 @@ fn router_with_proxy(
post(analytics::bootstrap_statsig),
)
.route("/auth/full_stripe_profile", get(account::stripe_profile))
.route("/auth/stripe_profile", get(account::stripe_profile))
.merge(tab::router())
.route_layer(DefaultBodyLimit::disable())
.route_layer(RequestDecompressionLayer::new())
@@ -139,6 +162,7 @@ fn router_with_proxy(
.method_not_allowed_fallback(proxy::forward)
.layer(Extension(proxy))
.layer(Extension(knowledge_service))
.layer(Extension(free_entitlements))
.with_state(registry)
.merge(web_cache)
}
+261 -2
View File
@@ -1,15 +1,25 @@
//! Implements advertisement configuration endpoints.
//! Advertisement service contract and desktop HTTP handler.
use std::{
collections::{BTreeSet, HashMap},
path::{Path as FilePath, PathBuf},
};
use axum::{
body::Body,
extract::{Path, State},
http::{HeaderMap, StatusCode},
http::{header, HeaderMap, HeaderValue, Response, StatusCode},
Json,
};
use bytes::BytesMut;
use futures_util::StreamExt;
use serde::{Deserialize, Serialize};
use sha2::{Digest, Sha256};
use url::Url;
use uuid::Uuid;
use crate::{Error, Result};
use crate::{config, Error, Result};
use super::ControlService;
@@ -23,6 +33,9 @@ pub(super) const OS_HEADER: &str = "X-Cursor-Assistant-OS";
pub(super) const APP_VERSION_HEADER: &str = "X-Cursor-Assistant-Version";
pub(super) const DISABLED_AD_IDS_HEADER: &str = "disable-ad-ids";
pub(super) const LANGUAGE_HEADER: &str = "accept-language";
const ADS_IMAGE_ROUTE: &str = "/__byok-api__/api/ads/images";
const MAX_AD_IMAGE_BYTES: usize = 10 * 1024 * 1024;
const IMAGE_EXTENSIONS: &[&str] = &["png", "jpg", "gif", "webp", "avif"];
#[derive(Clone, Debug, Deserialize, Serialize)]
pub struct AdRuntime {
@@ -104,6 +117,137 @@ impl AdRuntime {
}
Ok(self)
}
pub(super) async fn cache_images(&mut self, client: &reqwest::Client) {
let cache_dir = match config::managed_data_dir() {
Ok(path) => path.join("ads"),
Err(error) => {
tracing::warn!(%error, "failed to resolve advertisement image cache directory");
return;
}
};
let urls = self
.slots
.iter()
.flat_map(|slot| [&slot.target.image_url, &slot.content.image_url])
.cloned()
.collect::<BTreeSet<_>>();
let downloads = futures_util::future::join_all(urls.iter().map(|url| {
let cache_dir = &cache_dir;
async move {
let result = cache_image(client, cache_dir, url).await;
(url, result)
}
}))
.await;
let mut cached_urls = HashMap::new();
for (url, result) in downloads {
match result {
Ok(cached_url) => {
cached_urls.insert(url.as_str(), cached_url);
}
Err(error) => {
tracing::warn!(%error, image_url = %url, "failed to cache advertisement image")
}
}
}
for slot in &mut self.slots {
if let Some(url) = cached_urls.get(slot.target.image_url.as_str()) {
slot.target.image_url.clone_from(url);
}
if let Some(url) = cached_urls.get(slot.content.image_url.as_str()) {
slot.content.image_url.clone_from(url);
}
}
}
}
async fn cache_image(client: &reqwest::Client, cache_dir: &FilePath, url: &str) -> Result<String> {
tokio::fs::create_dir_all(cache_dir).await?;
let hash = hex::encode(Sha256::digest(url.as_bytes()));
if let Some(file_name) = cached_file_name(cache_dir, &hash).await {
return Ok(format!("{ADS_IMAGE_ROUTE}/{file_name}"));
}
let response = client
.get(url)
.timeout(std::time::Duration::from_secs(10))
.send()
.await?;
if !response.status().is_success() {
return Err(Error::Provider(format!(
"advertisement image download failed ({})",
response.status()
)));
}
let content_type = response
.headers()
.get(header::CONTENT_TYPE)
.and_then(|value| value.to_str().ok())
.and_then(image_extension)
.ok_or_else(|| {
Error::Provider("advertisement image has an unsupported content type".into())
})?;
if response
.content_length()
.is_some_and(|length| length > MAX_AD_IMAGE_BYTES as u64)
{
return Err(Error::Provider("advertisement image exceeds 10 MiB".into()));
}
let mut bytes = BytesMut::new();
let mut stream = response.bytes_stream();
while let Some(chunk) = stream.next().await {
let chunk = chunk?;
if bytes.len() + chunk.len() > MAX_AD_IMAGE_BYTES {
return Err(Error::Provider("advertisement image exceeds 10 MiB".into()));
}
bytes.extend_from_slice(&chunk);
}
let file_name = format!("{hash}.{content_type}");
let destination = cache_dir.join(&file_name);
let temporary = cache_dir.join(format!(".{file_name}.{}.tmp", Uuid::new_v4()));
tokio::fs::write(&temporary, &bytes).await?;
if let Err(error) = tokio::fs::rename(&temporary, &destination).await {
if !destination.exists() {
let _ = tokio::fs::remove_file(&temporary).await;
return Err(error.into());
}
let _ = tokio::fs::remove_file(&temporary).await;
}
Ok(format!("{ADS_IMAGE_ROUTE}/{file_name}"))
}
async fn cached_file_name(cache_dir: &FilePath, hash: &str) -> Option<String> {
for extension in IMAGE_EXTENSIONS {
let file_name = format!("{hash}.{extension}");
if tokio::fs::metadata(cache_dir.join(&file_name))
.await
.is_ok()
{
return Some(file_name);
}
}
None
}
fn image_extension(content_type: &str) -> Option<&'static str> {
match content_type
.split(';')
.next()
.unwrap_or_default()
.trim()
.to_ascii_lowercase()
.as_str()
{
"image/png" => Some("png"),
"image/jpeg" => Some("jpg"),
"image/gif" => Some("gif"),
"image/webp" => Some("webp"),
"image/avif" => Some("avif"),
_ => None,
}
}
fn validate_http_url(value: &str, field: &str) -> Result<()> {
@@ -117,6 +261,54 @@ fn validate_http_url(value: &str, field: &str) -> Result<()> {
Ok(())
}
pub async fn image(
Path(file_name): Path<String>,
) -> std::result::Result<Response<Body>, StatusCode> {
let (hash, extension) = file_name.rsplit_once('.').ok_or(StatusCode::NOT_FOUND)?;
if hash.len() != 64
|| !hash.bytes().all(|byte| byte.is_ascii_hexdigit())
|| !IMAGE_EXTENSIONS.contains(&extension)
{
return Err(StatusCode::NOT_FOUND);
}
let path: PathBuf = config::managed_data_dir()
.map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?
.join("ads")
.join(&file_name);
let bytes = tokio::fs::read(path).await.map_err(|error| {
if error.kind() == std::io::ErrorKind::NotFound {
StatusCode::NOT_FOUND
} else {
StatusCode::INTERNAL_SERVER_ERROR
}
})?;
let mut response = Response::new(Body::from(bytes));
response.headers_mut().insert(
header::CONTENT_TYPE,
HeaderValue::from_static(content_type_for_extension(extension)),
);
response.headers_mut().insert(
header::CACHE_CONTROL,
HeaderValue::from_static("public, max-age=31536000, immutable"),
);
response.headers_mut().insert(
header::X_CONTENT_TYPE_OPTIONS,
HeaderValue::from_static("nosniff"),
);
Ok(response)
}
fn content_type_for_extension(extension: &str) -> &'static str {
match extension {
"png" => "image/png",
"jpg" => "image/jpeg",
"gif" => "image/gif",
"webp" => "image/webp",
"avif" => "image/avif",
_ => "application/octet-stream",
}
}
pub async fn get(
State(service): State<ControlService>,
headers: HeaderMap,
@@ -147,3 +339,70 @@ pub async fn dismiss(
service.dismiss_ad(&ad_id, &input).await?;
Ok(StatusCode::NO_CONTENT)
}
#[cfg(test)]
mod tests {
use std::sync::{
atomic::{AtomicUsize, Ordering},
Arc,
};
use axum::{
body::Body,
http::{header, Response},
routing::get,
Router,
};
use super::*;
#[test]
fn recognizes_supported_image_content_types() {
assert_eq!(image_extension("image/png"), Some("png"));
assert_eq!(image_extension("image/jpeg; charset=binary"), Some("jpg"));
assert_eq!(image_extension("IMAGE/WEBP"), Some("webp"));
assert_eq!(image_extension("image/svg+xml"), None);
assert_eq!(image_extension("text/html"), None);
}
#[tokio::test]
async fn downloads_an_ad_image_once_and_reuses_the_cache() {
let requests = Arc::new(AtomicUsize::new(0));
let request_counter = requests.clone();
let app = Router::new().route(
"/ad.png",
get(move || {
request_counter.fetch_add(1, Ordering::SeqCst);
async {
Response::builder()
.header(header::CONTENT_TYPE, "image/png")
.body(Body::from(&b"cached image"[..]))
.unwrap()
}
}),
);
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let address = listener.local_addr().unwrap();
let server = tokio::spawn(async move { axum::serve(listener, app).await.unwrap() });
let root = tempfile::tempdir().unwrap();
let client = reqwest::Client::new();
let remote_url = format!("http://{address}/ad.png");
let first = cache_image(&client, root.path(), &remote_url)
.await
.unwrap();
let second = cache_image(&client, root.path(), &remote_url)
.await
.unwrap();
assert_eq!(first, second);
assert!(first.starts_with(ADS_IMAGE_ROUTE));
assert_eq!(requests.load(Ordering::SeqCst), 1);
let file_name = first.rsplit('/').next().unwrap();
assert_eq!(
tokio::fs::read(root.path().join(file_name)).await.unwrap(),
b"cached image"
);
server.abort();
}
}
+1
View File
@@ -112,6 +112,7 @@ fn proxy_error(error: impl std::fmt::Display) -> Response<Body> {
pub fn api_router(service: ControlService) -> Router {
Router::new()
.route("/__byok-api__/api/ads", get(ads::get))
.route("/__byok-api__/api/ads/images/{file_name}", get(ads::image))
.route(
"/__byok-api__/api/ads/{ad_id}/dismissals",
post(ads::dismiss),
+3 -1
View File
@@ -283,7 +283,9 @@ impl ControlService {
message.chars().take(200).collect::<String>()
)));
}
response.json::<AdRuntime>().await?.into_menu_slots()
let mut runtime = response.json::<AdRuntime>().await?.into_menu_slots()?;
runtime.cache_images(&client).await;
Ok(runtime)
}
pub(super) async fn dismiss_ad(&self, ad_id: &str, input: &AdDismissalInput) -> Result<()> {
+48 -10
View File
@@ -1,11 +1,15 @@
//! Implements settings management endpoints.
use crate::Result;
use axum::{extract::State, Json};
use axum::{
extract::State,
http::{header, HeaderMap},
Json,
};
use serde::{Deserialize, Serialize};
use crate::store::{
CommitSettings, DesktopSettings, PortSettings, ProxySettings, ProxySettingsInput,
StatisticsStorage, StatisticsStorageScope, TabSettings, DEFAULT_COMMIT_PROMPT,
CommitPromptLocale, CommitSettings, DesktopSettings, PortSettings, ProxySettings,
ProxySettingsInput, StatisticsStorage, StatisticsStorageScope, TabSettings,
};
use super::{ControlService, ObservabilitySettings};
@@ -90,27 +94,35 @@ pub async fn update_desktop(
/// Settings view for commit message generation. Empty `model_id` means 直连
/// (forward the original Cursor RPC). A non-empty value is a configured
/// Cursor model hash. Empty `prompt` means "use the built-in default".
/// built-in or plugin model identifier. Empty `prompt` means "use the built-in default".
#[derive(Serialize)]
pub struct CommitSettingsView {
pub model_id: String,
pub prompt: String,
pub prompt_locale: CommitPromptLocale,
pub default_prompt: &'static str,
}
impl From<CommitSettings> for CommitSettingsView {
fn from(settings: CommitSettings) -> Self {
impl CommitSettingsView {
fn new(settings: CommitSettings, default_locale: CommitPromptLocale) -> Self {
Self {
model_id: settings.model_id,
prompt: settings.prompt,
default_prompt: DEFAULT_COMMIT_PROMPT.trim(),
prompt_locale: settings.prompt_locale,
default_prompt: default_locale.default_prompt(),
}
}
}
pub async fn get_commit(State(service): State<ControlService>) -> Result<Json<CommitSettingsView>> {
pub async fn get_commit(
State(service): State<ControlService>,
headers: HeaderMap,
) -> Result<Json<CommitSettingsView>> {
let settings = service.commit_settings().await?;
Ok(Json(CommitSettingsView::from(settings)))
Ok(Json(CommitSettingsView::new(
settings,
requested_commit_locale(&headers),
)))
}
pub async fn update_commit(
@@ -118,5 +130,31 @@ pub async fn update_commit(
Json(settings): Json<CommitSettings>,
) -> Result<Json<CommitSettingsView>> {
let saved = service.set_commit_settings(settings).await?;
Ok(Json(CommitSettingsView::from(saved)))
let default_locale = saved.prompt_locale;
Ok(Json(CommitSettingsView::new(saved, default_locale)))
}
fn requested_commit_locale(headers: &HeaderMap) -> CommitPromptLocale {
match headers
.get(header::ACCEPT_LANGUAGE)
.and_then(|value| value.to_str().ok())
{
Some(value) if value.eq_ignore_ascii_case("zh-CN") => CommitPromptLocale::ZhCn,
_ => CommitPromptLocale::EnUs,
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn commit_default_prompt_locale_comes_from_interface_language() {
let mut headers = HeaderMap::new();
headers.insert(header::ACCEPT_LANGUAGE, "zh-CN".parse().unwrap());
assert_eq!(requested_commit_locale(&headers), CommitPromptLocale::ZhCn);
headers.insert(header::ACCEPT_LANGUAGE, "en-US".parse().unwrap());
assert_eq!(requested_commit_locale(&headers), CommitPromptLocale::EnUs);
}
}
+275 -80
View File
@@ -1,15 +1,17 @@
//! Implements Cursor account information services.
use axum::{
body::{Body, Bytes},
body::{to_bytes, Body},
extract::Extension,
http::{header, Request, Response},
http::{header, HeaderValue, Request, Response},
};
use prost::Message;
use serde_json::{Map, Value};
use serde_json::Value;
use crate::{api::cursor::proxy, Result};
use crate::{api::cursor::proxy, local_app, Result};
const LOCAL_AUTH_ID: &str = "local_ultra";
use super::entitlement::FreeEntitlementCache;
const LOCAL_AUTH_ID: &str = "cursor-local-user";
const LOCAL_EMAIL: &str = "cursor@ai.com";
const LOCAL_ULTRA_PLAN_INCLUDED_CENTS: i32 = 20_000;
@@ -21,6 +23,20 @@ struct GetEmailResponse {
sign_up_type: i32,
}
#[derive(Clone, PartialEq, Message)]
struct GetUserMetaResponse {
#[prost(string, tag = "1")]
email: String,
#[prost(int32, tag = "2")]
sign_up_type: i32,
#[prost(int64, tag = "3")]
user_id: i64,
#[prost(string, optional, tag = "4")]
workos_id: Option<String>,
#[prost(string, optional, tag = "5")]
profile_picture_url: Option<String>,
}
#[derive(Clone, PartialEq, Message)]
struct GetMeResponse {
#[prost(string, tag = "1")]
@@ -41,6 +57,8 @@ struct GetMeResponse {
email_domain_type: Option<String>,
#[prost(string, optional, tag = "12")]
country: Option<String>,
#[prost(string, optional, tag = "13")]
profile_picture_url: Option<String>,
}
#[derive(Clone, PartialEq, Message)]
@@ -134,7 +152,7 @@ pub async fn get_email(
Extension(upstream): Extension<proxy::CursorProxy>,
request: Request<Body>,
) -> Result<Response<Body>> {
forward_or(upstream, request, || {
local_or_forward(upstream, request, || {
proto(GetEmailResponse {
email: LOCAL_EMAIL.into(),
sign_up_type: 3,
@@ -143,11 +161,27 @@ pub async fn get_email(
.await
}
pub async fn get_user_meta(
Extension(upstream): Extension<proxy::CursorProxy>,
request: Request<Body>,
) -> Result<Response<Body>> {
local_or_forward(upstream, request, || {
proto(GetUserMetaResponse {
email: LOCAL_EMAIL.into(),
sign_up_type: 3,
user_id: 1,
workos_id: None,
profile_picture_url: None,
})
})
.await
}
pub async fn get_me(
Extension(upstream): Extension<proxy::CursorProxy>,
request: Request<Body>,
) -> Result<Response<Body>> {
forward_or(upstream, request, || {
local_or_forward(upstream, request, || {
proto(GetMeResponse {
auth_id: LOCAL_AUTH_ID.into(),
user_id: 1,
@@ -158,6 +192,7 @@ pub async fn get_me(
is_enterprise_user: Some(false),
email_domain_type: Some("personal".into()),
country: Some("US".into()),
profile_picture_url: None,
})
})
.await
@@ -167,14 +202,14 @@ pub async fn get_teams(
Extension(upstream): Extension<proxy::CursorProxy>,
request: Request<Body>,
) -> Result<Response<Body>> {
forward_or(upstream, request, || proto(Empty {})).await
local_or_forward(upstream, request, || proto(Empty {})).await
}
pub async fn get_user_profile(
Extension(upstream): Extension<proxy::CursorProxy>,
request: Request<Body>,
) -> Result<Response<Body>> {
forward_or(upstream, request, || {
local_or_forward(upstream, request, || {
proto(GetUserProfileResponse {
public_visibility_allowed: Some(true),
max_visibility: Some("PUBLIC".into()),
@@ -183,85 +218,170 @@ pub async fn get_user_profile(
.await
}
pub async fn current_period_usage() -> Result<Response<Body>> {
let now = chrono::Utc::now();
proto(GetCurrentPeriodUsageResponse {
billing_cycle_start: (now - chrono::Duration::days(30)).timestamp_millis(),
billing_cycle_end: (now + chrono::Duration::days(10 * 365)).timestamp_millis(),
plan_usage: Some(PlanUsage {
total_spend: 0,
included_spend: LOCAL_ULTRA_PLAN_INCLUDED_CENTS,
remaining: LOCAL_ULTRA_PLAN_INCLUDED_CENTS,
limit: LOCAL_ULTRA_PLAN_INCLUDED_CENTS,
remaining_bonus: Some(false),
bonus_tooltip: Some("Ultra local account mock is active.".into()),
auto_spend: Some(0),
api_spend: Some(0),
auto_percent_used: Some(0.0),
api_percent_used: Some(0.0),
total_percent_used: Some(0.0),
}),
spend_limit_usage: Some(SpendLimitUsage {
limit_type: "user".into(),
}),
display_threshold: Some(99_999_999),
enabled: true,
display_message: "Ultra plan active".into(),
auto_model_selected_display_message: Some("Ultra plan active".into()),
named_model_selected_display_message: Some("Ultra plan active".into()),
pub async fn current_period_usage(
Extension(upstream): Extension<proxy::CursorProxy>,
Extension(free_entitlements): Extension<FreeEntitlementCache>,
request: Request<Body>,
) -> Result<Response<Body>> {
local_or_confirmed_free_or_forward(upstream, &free_entitlements, request, || {
let now = chrono::Utc::now();
proto(GetCurrentPeriodUsageResponse {
billing_cycle_start: (now - chrono::Duration::days(30)).timestamp_millis(),
billing_cycle_end: (now + chrono::Duration::days(10 * 365)).timestamp_millis(),
plan_usage: Some(PlanUsage {
total_spend: 0,
included_spend: LOCAL_ULTRA_PLAN_INCLUDED_CENTS,
remaining: LOCAL_ULTRA_PLAN_INCLUDED_CENTS,
limit: LOCAL_ULTRA_PLAN_INCLUDED_CENTS,
remaining_bonus: Some(false),
bonus_tooltip: Some("Ultra local account mock is active.".into()),
auto_spend: Some(0),
api_spend: Some(0),
auto_percent_used: Some(0.0),
api_percent_used: Some(0.0),
total_percent_used: Some(0.0),
}),
spend_limit_usage: Some(SpendLimitUsage {
limit_type: "user".into(),
}),
display_threshold: Some(99_999_999),
enabled: true,
display_message: "Ultra plan active".into(),
auto_model_selected_display_message: Some("Ultra plan active".into()),
named_model_selected_display_message: Some("Ultra plan active".into()),
})
})
.await
}
pub async fn usage_limit_status() -> Result<Response<Body>> {
proto(GetUsageLimitStatusAndActiveGrantsResponse {
usage_limit_policy_status: Some(UsageLimitPolicyStatus {
is_in_slow_pool: false,
features: Default::default(),
can_configure_spend_limit: true,
has_pending_request: false,
allowed_model_ids: Vec::new(),
allowed_model_tags: Vec::new(),
}),
pub async fn usage_limit_status(
Extension(upstream): Extension<proxy::CursorProxy>,
Extension(free_entitlements): Extension<FreeEntitlementCache>,
request: Request<Body>,
) -> Result<Response<Body>> {
local_or_confirmed_free_or_forward(upstream, &free_entitlements, request, || {
proto(GetUsageLimitStatusAndActiveGrantsResponse {
usage_limit_policy_status: Some(UsageLimitPolicyStatus {
is_in_slow_pool: false,
features: Default::default(),
can_configure_spend_limit: true,
has_pending_request: false,
allowed_model_ids: Vec::new(),
allowed_model_tags: Vec::new(),
}),
})
})
.await
}
pub async fn stripe_profile(
Extension(upstream): Extension<proxy::CursorProxy>,
Extension(free_entitlements): Extension<FreeEntitlementCache>,
request: Request<Body>,
) -> Result<Response<Body>> {
match proxy::forward_buffered(&upstream, request).await {
Ok(response) if response.status.is_success() => {
let mut profile = serde_json::from_slice::<Map<String, Value>>(&response.body)?;
ultra(&mut profile);
Ok(response.with_body(Bytes::from(serde_json::to_vec(&profile)?)))
}
Ok(response) => {
tracing::warn!(status = %response.status, "Cursor account upstream rejected profile; using local Ultra identity");
json(ultra_profile())
}
Err(error) => {
tracing::warn!(%error, "Cursor account upstream unavailable; using local Ultra identity");
json(ultra_profile())
}
if local_app::request_uses_local_cursor_token(request.headers()) {
return local_stripe_profile(request.headers().get(header::ORIGIN).cloned());
}
let request_headers = request.headers().clone();
let origin = request.headers().get(header::ORIGIN).cloned();
let upstream_response = match proxy::forward_buffered(&upstream, request).await {
Ok(response) => response,
Err(error) if free_entitlements.is_confirmed_free(&request_headers) => {
tracing::warn!(%error, "using cached Free entitlement after Stripe upstream failure");
return local_stripe_profile(origin);
}
Err(error) => return Err(error),
};
if !upstream_response.status.is_success() {
if should_fallback_to_cached_free(
upstream_response.status,
&free_entitlements,
&request_headers,
) {
tracing::warn!(
status = %upstream_response.status,
"using cached Free entitlement after Stripe upstream failure"
);
return local_stripe_profile(origin);
}
return Ok(upstream_response.into_response());
}
let Some(membership_type) = membership_type(&upstream_response.body) else {
return Ok(upstream_response.into_response());
};
let observed = free_entitlements.observe_membership(&request_headers, &membership_type);
if observed && membership_type.eq_ignore_ascii_case("free") {
return local_stripe_profile(origin);
}
Ok(upstream_response.into_response())
}
async fn forward_or(
fn should_fallback_to_cached_free(
status: axum::http::StatusCode,
free_entitlements: &FreeEntitlementCache,
headers: &axum::http::HeaderMap,
) -> bool {
status.is_server_error() && free_entitlements.is_confirmed_free(headers)
}
fn membership_type(body: &[u8]) -> Option<String> {
let profile: Value = serde_json::from_slice(body).ok()?;
let membership_type = profile.get("membershipType")?.as_str()?.trim();
(!membership_type.is_empty()).then(|| membership_type.to_owned())
}
fn local_stripe_profile(origin: Option<HeaderValue>) -> Result<Response<Body>> {
let mut response = json(ultra_profile())?;
if let Some(origin) = origin {
response
.headers_mut()
.insert(header::ACCESS_CONTROL_ALLOW_ORIGIN, origin);
response.headers_mut().insert(
header::ACCESS_CONTROL_ALLOW_CREDENTIALS,
HeaderValue::from_static("true"),
);
response
.headers_mut()
.insert(header::VARY, HeaderValue::from_static("Origin"));
}
Ok(response)
}
async fn local_or_confirmed_free_or_forward(
upstream: proxy::CursorProxy,
free_entitlements: &FreeEntitlementCache,
request: Request<Body>,
local: impl FnOnce() -> Result<Response<Body>>,
) -> Result<Response<Body>> {
if local_app::request_uses_local_cursor_token(request.headers())
|| free_entitlements.is_confirmed_free(request.headers())
{
consume_body(request).await?;
return local();
}
proxy::forward(Extension(upstream), request).await
}
async fn local_or_forward(
upstream: proxy::CursorProxy,
request: Request<Body>,
fallback: impl FnOnce() -> Result<Response<Body>>,
local: impl FnOnce() -> Result<Response<Body>>,
) -> Result<Response<Body>> {
match proxy::forward_buffered(&upstream, request).await {
Ok(response) if response.status.is_success() => Ok(response.into_response()),
Ok(response) => {
tracing::debug!(status = %response.status, "Cursor identity upstream rejected request; using local identity");
fallback()
}
Err(error) => {
tracing::warn!(%error, "Cursor identity upstream unavailable; using local identity");
fallback()
}
if local_app::request_uses_local_cursor_token(request.headers()) {
consume_body(request).await?;
return local();
}
proxy::forward(Extension(upstream), request).await
}
async fn consume_body(request: Request<Body>) -> Result<()> {
to_bytes(request.into_body(), usize::MAX)
.await
.map_err(|error| crate::Error::Protocol(format!("cannot read request body: {error}")))?;
Ok(())
}
fn proto(message: impl Message) -> Result<Response<Body>> {
@@ -289,15 +409,6 @@ fn response(content_type: &'static str, body: Vec<u8>) -> Result<Response<Body>>
Ok(response)
}
fn ultra(profile: &mut Map<String, Value>) {
profile.insert("membershipType".into(), Value::String("ultra".into()));
profile.insert(
"individualMembershipType".into(),
Value::String("ultra".into()),
);
profile.insert("subscriptionStatus".into(), Value::String("active".into()));
}
fn ultra_profile() -> Value {
serde_json::json!({
"membershipType": "ultra",
@@ -310,3 +421,87 @@ fn ultra_profile() -> Value {
"isTeamMember": false
})
}
#[cfg(test)]
mod tests {
use std::sync::{
atomic::{AtomicBool, Ordering},
Arc,
};
use axum::body::Bytes;
use futures_util::stream;
use super::*;
#[tokio::test]
async fn local_account_response_consumes_request_body_before_replying() {
let polled = Arc::new(AtomicBool::new(false));
let observed = polled.clone();
let body = Body::from_stream(stream::once(async move {
observed.store(true, Ordering::SeqCst);
Ok::<_, std::convert::Infallible>(Bytes::from_static(b"request"))
}));
consume_body(Request::new(body)).await.unwrap();
assert!(polled.load(Ordering::SeqCst));
}
#[test]
fn cached_free_fallback_accepts_server_failures_but_not_auth_failures() {
let cache = FreeEntitlementCache::default();
let mut headers = axum::http::HeaderMap::new();
headers.insert(
header::AUTHORIZATION,
HeaderValue::from_static("Bearer official-free-token"),
);
assert!(cache.observe_membership(&headers, "free"));
assert!(should_fallback_to_cached_free(
axum::http::StatusCode::BAD_GATEWAY,
&cache,
&headers
));
assert!(!should_fallback_to_cached_free(
axum::http::StatusCode::UNAUTHORIZED,
&cache,
&headers
));
}
#[test]
fn reads_only_a_non_empty_membership_type() {
assert_eq!(
membership_type(br#"{"membershipType":"free"}"#).as_deref(),
Some("free")
);
assert_eq!(membership_type(br#"{"membershipType":""}"#), None);
assert_eq!(membership_type(br#"{"subscriptionStatus":"active"}"#), None);
assert_eq!(membership_type(b"not-json"), None);
}
#[test]
fn local_stripe_profile_allows_the_cursor_app_origin() {
let origin = HeaderValue::from_static("vscode-file://vscode-app");
let response = local_stripe_profile(Some(origin.clone())).unwrap();
assert_eq!(
response.headers().get(header::ACCESS_CONTROL_ALLOW_ORIGIN),
Some(&origin)
);
assert_eq!(
response
.headers()
.get(header::ACCESS_CONTROL_ALLOW_CREDENTIALS),
Some(&HeaderValue::from_static("true"))
);
assert_eq!(
response.headers().get(header::VARY),
Some(&HeaderValue::from_static("Origin"))
);
assert_eq!(
response.headers().get(header::CONTENT_TYPE),
Some(&HeaderValue::from_static("application/json"))
);
}
}
+26 -17
View File
@@ -2,8 +2,8 @@
//!
//! Cursor sends `aiserver.v1.AiService/WriteGitCommitMessage` with the staged
//! diffs. Empty commit-settings `model_id` keeps the original behaviour and
//! forwards the RPC unchanged (直连). A configured Cursor model hash answers
//! the request locally: truncated diffs + previous commits form the user
//! forwards the RPC unchanged (直连). A configured local model identifier
//! answers the request locally: truncated diffs + previous commits form the user
//! message, the customizable commit prompt is the system prompt, and the raw
//! completion is cleaned before being returned.
use std::{
@@ -30,6 +30,7 @@ use crate::{
ContentPart, ModelInvocation, ModelRequest, ModelSpec, ProjectedContent, ProjectedMessage,
PromptSpec, Role,
},
plugin::ADAPTER_ID_PREFIX,
provider::{ModelEvent, Provider},
store::CommitSettings,
Error, Result,
@@ -85,19 +86,11 @@ async fn generate_local(
if diffs.is_empty() {
return Err(Error::Protocol("diffs are required".into()));
}
let model_hash = settings.model_id.trim();
let model = registry
.store()
.model(model_hash)
.await?
.ok_or_else(|| {
Error::Provider(format!(
"commit model {model_hash} is not configured; select a Cursor model in the commit settings"
))
})?;
let model_id = settings.model_id.trim();
ensure_configured_model(registry, model_id).await?;
let invocation = build_invocation(
&settings,
&model.model_hash,
model_id,
build_user_content(&request, &diffs),
);
let provider = registry.conversations().dependencies().provider.clone();
@@ -121,9 +114,25 @@ async fn generate_local(
Ok(response)
}
async fn ensure_configured_model(registry: &TransportRegistry, model_id: &str) -> Result<()> {
if model_id.starts_with(ADAPTER_ID_PREFIX) {
let plugins = registry.plugins().ok_or_else(|| {
Error::Provider(format!("commit plugin model {model_id} is unavailable"))
})?;
plugins.model_descriptor(model_id).await?;
return Ok(());
}
if registry.store().model(model_id).await?.is_some() {
return Ok(());
}
Err(Error::Provider(format!(
"commit model {model_id} is not configured; select a configured model in the commit settings"
)))
}
fn build_invocation(
settings: &CommitSettings,
model_hash: &str,
model_id: &str,
user_content: String,
) -> ModelInvocation {
let call_id = format!("commit-message-{}", uuid::Uuid::new_v4());
@@ -137,7 +146,7 @@ fn build_invocation(
instructions: settings.effective_prompt().to_owned(),
tools: Vec::new(),
},
model: ModelSpec::new(model_hash.to_owned()),
model: ModelSpec::new(model_id.to_owned()),
history: vec![ProjectedMessage {
message_id: "commit-message".into(),
role: Role::User,
@@ -360,12 +369,12 @@ mod tests {
assert!(CommitSettings::default().is_direct());
assert!(CommitSettings {
model_id: " ".into(),
prompt: String::new(),
..CommitSettings::default()
}
.is_direct());
assert!(!CommitSettings {
model_id: "abc".into(),
prompt: String::new(),
..CommitSettings::default()
}
.is_direct());
}
+119
View File
@@ -0,0 +1,119 @@
//! Routes Cursor metadata calls according to the supplied authentication token.
use axum::{
body::{to_bytes, Body},
extract::{Extension, Request},
http::{header, HeaderValue, Response},
};
use prost::Message;
use crate::{
api::cursor::proxy::{self, CursorProxy},
cursor::protocol::proto::agent::v1 as agent,
local_app, Result,
};
#[derive(Clone, Copy, PartialEq, Message)]
struct EmptyResponse {}
pub async fn available_docs(
Extension(proxy): Extension<CursorProxy>,
request: Request<Body>,
) -> Result<Response<Body>> {
route(&proxy, request, EmptyResponse {}).await
}
pub async fn effective_user_plugins(
Extension(proxy): Extension<CursorProxy>,
request: Request<Body>,
) -> Result<Response<Body>> {
route(&proxy, request, EmptyResponse {}).await
}
pub async fn user_privacy_mode(
Extension(proxy): Extension<CursorProxy>,
request: Request<Body>,
) -> Result<Response<Body>> {
route(&proxy, request, EmptyResponse {}).await
}
pub async fn update_conversation_metadata(
Extension(proxy): Extension<CursorProxy>,
request: Request<Body>,
) -> Result<Response<Body>> {
route(
&proxy,
request,
agent::UpdateConversationMetadataResponse {},
)
.await
}
async fn route<M: Message>(
proxy: &CursorProxy,
request: Request<Body>,
mock: M,
) -> Result<Response<Body>> {
let local = local_app::request_uses_local_cursor_token(request.headers());
if local {
consume_body(request).await?;
return Ok(proto(mock));
}
proxy::forward(Extension(proxy.clone()), request).await
}
async fn consume_body(request: Request<Body>) -> Result<()> {
to_bytes(request.into_body(), usize::MAX)
.await
.map_err(|error| crate::Error::Protocol(format!("cannot read request body: {error}")))?;
Ok(())
}
fn proto(message: impl Message) -> Response<Body> {
let body = message.encode_to_vec();
let length = body.len();
let mut response = Response::new(Body::from(body));
response.headers_mut().insert(
header::CONTENT_TYPE,
HeaderValue::from_static("application/proto"),
);
response.headers_mut().insert(
header::CONTENT_LENGTH,
HeaderValue::from_str(&length.to_string()).expect("body length is valid"),
);
response
}
#[cfg(test)]
mod tests {
use std::sync::{
atomic::{AtomicBool, Ordering},
Arc,
};
use axum::body::{to_bytes, Bytes};
use futures_util::stream;
use super::*;
#[tokio::test]
async fn local_response_consumes_request_body_before_replying() {
let polled = Arc::new(AtomicBool::new(false));
let observed = polled.clone();
let body = Body::from_stream(stream::once(async move {
observed.store(true, Ordering::SeqCst);
Ok::<_, std::convert::Infallible>(Bytes::from_static(b"request"))
}));
consume_body(Request::new(body)).await.unwrap();
assert!(polled.load(Ordering::SeqCst));
}
#[tokio::test]
async fn empty_unary_mock_is_bare_protobuf_without_connect_frame() {
let response = proto(EmptyResponse {});
assert_eq!(response.headers()[header::CONTENT_LENGTH], "0");
let body = to_bytes(response.into_body(), usize::MAX).await.unwrap();
assert!(body.is_empty());
}
}
+143
View File
@@ -0,0 +1,143 @@
//! Caches a recently confirmed official Cursor Free entitlement.
use std::{
sync::Arc,
time::{Duration, Instant},
};
use axum::http::{header, HeaderMap};
use parking_lot::Mutex;
use sha2::{Digest, Sha256};
use crate::local_app;
const FREE_ENTITLEMENT_TTL: Duration = Duration::from_secs(5 * 60);
#[derive(Clone, Default)]
pub struct FreeEntitlementCache {
cached: Arc<Mutex<Option<CachedFreeEntitlement>>>,
}
struct CachedFreeEntitlement {
token_hash: [u8; 32],
expires_at: Instant,
}
impl FreeEntitlementCache {
pub fn is_confirmed_free(&self, headers: &HeaderMap) -> bool {
let Some(token_hash) = official_token_hash(headers) else {
return false;
};
let now = Instant::now();
let mut cached = self.cached.lock();
match cached.as_ref() {
Some(entry) if entry.expires_at > now && entry.token_hash == token_hash => true,
Some(entry) if entry.expires_at <= now => {
*cached = None;
false
}
_ => false,
}
}
pub fn observe_membership(&self, headers: &HeaderMap, membership_type: &str) -> bool {
let Some(token_hash) = official_token_hash(headers) else {
return false;
};
let mut cached = self.cached.lock();
if membership_type.eq_ignore_ascii_case("free") {
*cached = Some(CachedFreeEntitlement {
token_hash,
expires_at: Instant::now() + FREE_ENTITLEMENT_TTL,
});
} else if cached
.as_ref()
.is_some_and(|entry| entry.token_hash == token_hash)
{
*cached = None;
}
true
}
}
fn official_token_hash(headers: &HeaderMap) -> Option<[u8; 32]> {
if local_app::request_uses_local_cursor_token(headers) {
return None;
}
let token = headers
.get(header::AUTHORIZATION)?
.to_str()
.ok()?
.strip_prefix("Bearer ")?;
if token.is_empty() {
return None;
}
Some(Sha256::digest(token.as_bytes()).into())
}
#[cfg(test)]
mod tests {
use super::*;
fn official_headers(token: &str) -> HeaderMap {
let mut headers = HeaderMap::new();
headers.insert(
header::AUTHORIZATION,
format!("Bearer {token}").parse().unwrap(),
);
headers
}
#[test]
fn caches_only_a_confirmed_free_official_token() {
let cache = FreeEntitlementCache::default();
let free = official_headers("official-free-token");
let other = official_headers("another-token");
cache.observe_membership(&free, "free");
assert!(cache.is_confirmed_free(&free));
assert!(!cache.is_confirmed_free(&other));
}
#[test]
fn confirmed_non_free_membership_clears_the_same_token() {
let cache = FreeEntitlementCache::default();
let headers = official_headers("official-token");
cache.observe_membership(&headers, "free");
cache.observe_membership(&headers, "pro");
assert!(!cache.is_confirmed_free(&headers));
}
#[test]
fn local_token_never_enters_the_entitlement_cache() {
let cache = FreeEntitlementCache::default();
let mut headers = HeaderMap::new();
headers.insert(
header::AUTHORIZATION,
crate::local_app::local_cursor_authorization()
.parse()
.unwrap(),
);
assert!(!cache.observe_membership(&headers, "free"));
assert!(!cache.is_confirmed_free(&headers));
assert!(cache.cached.lock().is_none());
}
#[test]
fn expired_confirmation_is_removed() {
let cache = FreeEntitlementCache::default();
let headers = official_headers("official-token");
let token_hash = official_token_hash(&headers).unwrap();
*cache.cached.lock() = Some(CachedFreeEntitlement {
token_hash,
expires_at: Instant::now() - Duration::from_secs(1),
});
assert!(!cache.is_confirmed_free(&headers));
assert!(cache.cached.lock().is_none());
}
}
+2
View File
@@ -4,7 +4,9 @@ pub mod account;
pub mod analytics;
pub mod blob_sync;
pub mod commit_message;
pub mod compatibility;
pub mod context_sync;
pub(crate) mod entitlement;
pub mod knowledge;
pub mod model_catalog;
pub mod observability;
+75 -3
View File
@@ -52,23 +52,24 @@ async fn inject_if_missing_at(path: &Path) -> Result<()> {
.execute(&mut connection)
.await?;
let token = local_token()?;
let account = sqlx::query("SELECT CAST(value AS TEXT) AS value FROM ItemTable WHERE key = ?")
.bind("cursorAuth/accessToken")
.fetch_optional(&mut connection)
.await?;
if account.is_some_and(|row| {
row.try_get::<String, _>("value")
.is_ok_and(|value| !value.trim().is_empty())
.is_ok_and(|value| !value.trim().is_empty() && value != token)
}) {
return Ok(());
}
let token = local_token()?;
let values = [
("cursorAuth/accessToken", token.as_str()),
("cursorAuth/refreshToken", token.as_str()),
("cursorAuth/cachedEmail", EMAIL),
("cursorAuth/cachedSignUpType", SIGN_UP_TYPE),
("cursorAuth/stripeMembershipAuthId", SUBJECT),
("cursorAuth/stripeMembershipType", MEMBERSHIP_TYPE),
("cursorAuth/stripeSubscriptionStatus", SUBSCRIPTION_STATUS),
];
@@ -89,7 +90,17 @@ async fn inject_if_missing_at(path: &Path) -> Result<()> {
Ok(())
}
fn local_token() -> Result<String> {
pub(crate) fn is_local_cursor_authorization(authorization: &str) -> bool {
authorization
.strip_prefix("Bearer ")
.is_some_and(is_local_cursor_token)
}
fn is_local_cursor_token(token: &str) -> bool {
local_token().is_ok_and(|local| local == token)
}
pub(super) fn local_token() -> Result<String> {
let header = URL_SAFE_NO_PAD.encode(br#"{"alg":"HS256","typ":"JWT"}"#);
let payload = URL_SAFE_NO_PAD.encode(serde_json::to_vec(&json!({
"sub": SUBJECT,
@@ -101,3 +112,64 @@ fn local_token() -> Result<String> {
}))?);
Ok(format!("{header}.{payload}.{SUBJECT}"))
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn reinjection_repairs_local_membership_cache() {
let directory = tempfile::tempdir().unwrap();
let path = directory.path().join("state.vscdb");
inject_if_missing_at(&path).await.unwrap();
let mut connection = SqliteConnection::connect(&format!("sqlite:{}", path.display()))
.await
.unwrap();
sqlx::query("UPDATE ItemTable SET value = 'free' WHERE key = ?")
.bind("cursorAuth/stripeMembershipType")
.execute(&mut connection)
.await
.unwrap();
sqlx::query("DELETE FROM ItemTable WHERE key = ?")
.bind("cursorAuth/stripeMembershipAuthId")
.execute(&mut connection)
.await
.unwrap();
drop(connection);
inject_if_missing_at(&path).await.unwrap();
let mut connection = SqliteConnection::connect(&format!("sqlite:{}", path.display()))
.await
.unwrap();
let membership_type: String =
sqlx::query_scalar("SELECT CAST(value AS TEXT) FROM ItemTable WHERE key = ?")
.bind("cursorAuth/stripeMembershipType")
.fetch_one(&mut connection)
.await
.unwrap();
let membership_auth_id: String =
sqlx::query_scalar("SELECT CAST(value AS TEXT) FROM ItemTable WHERE key = ?")
.bind("cursorAuth/stripeMembershipAuthId")
.fetch_one(&mut connection)
.await
.unwrap();
assert_eq!(membership_type, MEMBERSHIP_TYPE);
assert_eq!(membership_auth_id, SUBJECT);
}
#[test]
fn recognizes_only_the_injected_cursor_token() {
let token = local_token().unwrap();
assert!(is_local_cursor_token(&token));
assert!(is_local_cursor_authorization(&format!("Bearer {token}")));
assert!(!is_local_cursor_authorization(&token));
assert!(!is_local_cursor_authorization(
"Bearer official-cursor-token"
));
assert!(!is_local_cursor_token("official-cursor-token"));
assert!(!is_local_cursor_token(""));
}
}
+31 -3
View File
@@ -1,6 +1,7 @@
//! Exposes the local desktop application integration.
mod account;
mod ca;
mod process;
mod proxy;
mod settings;
@@ -21,8 +22,16 @@ pub(crate) fn proxy_host_allowed(host: &str) -> bool {
proxy::is_cursor_host(host)
}
fn integration_prerequisites_ready(ca: &CaState, backend_ready: bool) -> bool {
matches!(ca, CaState::Ready) && backend_ready
pub(crate) fn request_uses_local_cursor_token(headers: &axum::http::HeaderMap) -> bool {
headers
.get(axum::http::header::AUTHORIZATION)
.and_then(|value| value.to_str().ok())
.is_some_and(account::is_local_cursor_authorization)
}
#[cfg(test)]
pub(crate) fn local_cursor_authorization() -> String {
format!("Bearer {}", account::local_token().unwrap())
}
#[derive(Clone, Debug, Serialize)]
@@ -49,6 +58,7 @@ pub struct CursorHarnessStatus {
pub configured_models: usize,
pub enabled_models: usize,
pub integration: IntegrationState,
pub settings_applied: bool,
pub proxy_url: Option<String>,
pub ca_install_command: Option<String>,
}
@@ -103,7 +113,10 @@ impl CursorHarness {
let configured_models = models.len();
let enabled_models = configured_models;
let ca = self.inner.ca.state()?;
if integration_prerequisites_ready(&ca, self.inner.backend_addr.read().is_some()) {
if self.inner.store.cursor_takeover_enabled().await?
&& matches!(ca, CaState::Ready)
&& self.inner.backend_addr.read().is_some()
{
self.enable().await?;
}
let proxy = self.inner.proxy.lock().await;
@@ -124,6 +137,7 @@ impl CursorHarness {
configured_models,
enabled_models,
integration,
settings_applied,
proxy_url,
ca_install_command: self.inner.ca.install_command(),
})
@@ -139,9 +153,23 @@ impl CursorHarness {
}
pub async fn set_enabled(&self, enabled: bool) -> Result<CursorHarnessStatus> {
let settings_applied = {
let proxy = self.inner.proxy.lock().await;
proxy
.url()
.as_deref()
.map(settings::settings_match)
.transpose()?
.unwrap_or(false)
};
if enabled {
if !settings_applied {
process::terminate_cursor().await?;
}
self.inner.store.set_cursor_takeover_enabled(true).await?;
self.enable().await?;
} else {
self.inner.store.set_cursor_takeover_enabled(false).await?;
self.disable().await?;
}
self.status().await
+79
View File
@@ -0,0 +1,79 @@
//! Terminates the Cursor desktop process before an explicit takeover.
use tokio::process::Command;
use crate::{Error, Result};
pub async fn terminate_cursor() -> Result<()> {
terminate_platform_cursor().await
}
#[cfg(target_os = "macos")]
async fn terminate_platform_cursor() -> Result<()> {
terminate_unix_process("Cursor").await
}
#[cfg(target_os = "linux")]
async fn terminate_platform_cursor() -> Result<()> {
terminate_unix_process("cursor").await?;
terminate_unix_process("Cursor").await
}
#[cfg(any(target_os = "macos", target_os = "linux"))]
async fn terminate_unix_process(name: &str) -> Result<()> {
let running = Command::new("pgrep").args(["-x", name]).status().await?;
if !running.success() {
return match running.code() {
Some(1) => Ok(()),
_ => Err(Error::Config(format!(
"failed to inspect the {name} process"
))),
};
}
let terminated = Command::new("pkill").args(["-x", name]).status().await?;
if terminated.success() || terminated.code() == Some(1) {
Ok(())
} else {
Err(Error::Config(format!(
"failed to terminate the {name} process"
)))
}
}
#[cfg(target_os = "windows")]
async fn terminate_platform_cursor() -> Result<()> {
let processes = Command::new("tasklist")
.args(["/FI", "IMAGENAME eq Cursor.exe", "/NH", "/FO", "CSV"])
.output()
.await?;
if !processes.status.success() {
return Err(Error::Config(
"failed to inspect the Cursor.exe process".into(),
));
}
if !String::from_utf8_lossy(&processes.stdout)
.to_ascii_lowercase()
.contains("cursor.exe")
{
return Ok(());
}
let terminated = Command::new("taskkill")
.args(["/F", "/T", "/IM", "Cursor.exe"])
.status()
.await?;
if terminated.success() {
Ok(())
} else {
Err(Error::Config(
"failed to terminate the Cursor.exe process".into(),
))
}
}
#[cfg(not(any(target_os = "macos", target_os = "linux", target_os = "windows")))]
async fn terminate_platform_cursor() -> Result<()> {
Err(Error::Config(format!(
"terminating Cursor is unsupported on {}",
std::env::consts::OS
)))
}
+13
View File
@@ -159,6 +159,10 @@ fn is_local_path(path: &str) -> bool {
path,
"/agent.v1.AgentService/RunSSE"
| "/aiserver.v1.BidiService/BidiAppend"
| "/aiserver.v1.AiService/AvailableDocs"
| "/aiserver.v1.DashboardService/GetEffectiveUserPlugins"
| "/aiserver.v1.DashboardService/GetUserPrivacyMode"
| "/agent.v1.AgentService/UpdateConversationMetadata"
| "/aiserver.v1.AiService/GetServerConfig"
| "/aiserver.v1.ServerConfigService/GetServerConfig"
| "/aiserver.v1.AiService/AvailableModels"
@@ -169,6 +173,7 @@ fn is_local_path(path: &str) -> bool {
| "/aiserver.v1.AiService/GetDefaultModel"
| "/aiserver.v1.AiService/GetDefaultModelNudgeData"
| "/aiserver.v1.AuthService/GetEmail"
| "/aiserver.v1.AuthService/GetUserMeta"
| "/aiserver.v1.DashboardService/GetMe"
| "/aiserver.v1.DashboardService/GetTeams"
| "/aiserver.v1.DashboardService/GetUserProfile"
@@ -182,6 +187,7 @@ fn is_local_path(path: &str) -> bool {
| "/aiserver.v1.NetworkService/IsConnected"
| "/aiserver.v1.AnalyticsService/BootstrapStatsig"
| "/auth/full_stripe_profile"
| "/auth/stripe_profile"
)
}
@@ -202,6 +208,13 @@ mod tests {
"/aiserver.v1.AiService/GetDefaultModelForCli",
"/aiserver.v1.AiService/GetDefaultModel",
"/aiserver.v1.AiService/GetDefaultModelNudgeData",
"/aiserver.v1.AiService/AvailableDocs",
"/aiserver.v1.DashboardService/GetEffectiveUserPlugins",
"/aiserver.v1.DashboardService/GetUserPrivacyMode",
"/aiserver.v1.AuthService/GetUserMeta",
"/agent.v1.AgentService/UpdateConversationMetadata",
"/auth/full_stripe_profile",
"/auth/stripe_profile",
] {
assert!(is_local_path(path), "{path} must not reach Cursor upstream");
}
+1
View File
@@ -6,6 +6,7 @@ mod normalize;
mod openai_chat;
mod openai_responses;
mod recorder;
mod request_template;
mod router;
use std::pin::Pin;
+67
View File
@@ -0,0 +1,67 @@
//! Renders per-invocation placeholders in user-configured provider request values.
const SESSION_ID_PLACEHOLDER: &str = "{{SessionId}}";
pub(super) fn render_json_strings(
value: &serde_json::Value,
conversation_id: &str,
) -> serde_json::Value {
match value {
serde_json::Value::String(value) => {
serde_json::Value::String(value.replace(SESSION_ID_PLACEHOLDER, conversation_id))
}
serde_json::Value::Array(values) => serde_json::Value::Array(
values
.iter()
.map(|value| render_json_strings(value, conversation_id))
.collect(),
),
serde_json::Value::Object(values) => serde_json::Value::Object(
values
.iter()
.map(|(name, value)| (name.clone(), render_json_strings(value, conversation_id)))
.collect(),
),
value => value.clone(),
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn renders_session_id_in_nested_json_string_values() {
let template = serde_json::json!({
"session": "{{SessionId}}",
"nested": {
"label": "conversation={{SessionId}}/{{SessionId}}",
"values": ["{{SessionId}}", 42, true, null]
}
});
assert_eq!(
render_json_strings(&template, "cursor-conversation-id"),
serde_json::json!({
"session": "cursor-conversation-id",
"nested": {
"label": "conversation=cursor-conversation-id/cursor-conversation-id",
"values": ["cursor-conversation-id", 42, true, null]
}
})
);
}
#[test]
fn leaves_json_property_names_and_unrelated_strings_unchanged() {
let template = serde_json::json!({
"{{SessionId}}": "literal",
"other": "{{sessionId}}"
});
assert_eq!(
render_json_strings(&template, "cursor-conversation-id"),
template
);
}
}
+39 -4
View File
@@ -82,7 +82,11 @@ impl Provider for ProviderRouter {
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.extra_params =
super::request_template::render_json_strings(
model.extra_params(),
&invocation.conversation_id,
);
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 guard = recorder.cancel_on_drop();
@@ -90,7 +94,14 @@ impl Provider for ProviderRouter {
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() },
custom_headers: if model.custom_headers_enabled {
custom_headers(
&model.custom_headers,
&invocation.conversation_id,
)?
} else {
reqwest::header::HeaderMap::new()
},
max_output_tokens: model.max_output_tokens(),
request_timeout,
allowed_body_fields: None,
@@ -297,8 +308,12 @@ fn root_error_message(error: &(dyn std::error::Error + 'static)) -> String {
current.to_string()
}
fn custom_headers(value: &serde_json::Value) -> Result<reqwest::header::HeaderMap> {
let object = value
fn custom_headers(
value: &serde_json::Value,
conversation_id: &str,
) -> Result<reqwest::header::HeaderMap> {
let rendered = super::request_template::render_json_strings(value, conversation_id);
let object = rendered
.as_object()
.ok_or_else(|| Error::Config("custom headers must be an object".into()))?;
let mut headers = reqwest::header::HeaderMap::new();
@@ -356,6 +371,26 @@ fn build_inner(
mod tests {
use super::*;
#[test]
fn renders_cursor_conversation_id_in_custom_header_values() {
let template = serde_json::json!({
"x-opencode-session-id": "{{SessionId}}",
"x-label": "cursor/{{SessionId}}"
});
let headers = custom_headers(&template, "cursor-conversation-id").unwrap();
assert_eq!(
headers.get("x-opencode-session-id").unwrap(),
"cursor-conversation-id"
);
assert_eq!(
headers.get("x-label").unwrap(),
"cursor/cursor-conversation-id"
);
assert_eq!(template["x-opencode-session-id"], "{{SessionId}}");
}
#[tokio::test]
async fn pending_provider_event_hits_the_idle_timeout() {
let mut stream: ProviderStream = Box::pin(futures_util::stream::pending());
+84 -6
View File
@@ -11,9 +11,11 @@ const TAB_SETTINGS_KEY: &str = "cursor_tab";
const INSTALLATION_ID_KEY: &str = "installation_id";
const DESKTOP_SETTINGS_KEY: &str = "desktop_lifecycle";
const COMMIT_SETTINGS_KEY: &str = "commit_settings";
const CURSOR_TAKEOVER_ENABLED_KEY: &str = "cursor_takeover_enabled";
/// Embedded default system prompt for commit message generation.
pub const DEFAULT_COMMIT_PROMPT: &str = include_str!("../../prompt/cursor/commit/prompt.md");
/// Embedded default system prompts for commit message generation.
pub const DEFAULT_COMMIT_PROMPT_ZH_CN: &str = include_str!("../../prompt/cursor/commit/zh-CN.md");
pub const DEFAULT_COMMIT_PROMPT_EN_US: &str = include_str!("../../prompt/cursor/commit/en-US.md");
pub const PUBLIC_TAB_SERVICE_URL: &str = "https://tab.leokun.cn";
@@ -83,17 +85,37 @@ impl TabSettings {
}
}
#[derive(Clone, Copy, Debug, Default, Deserialize, PartialEq, Eq, Serialize)]
pub enum CommitPromptLocale {
#[default]
#[serde(rename = "zh-CN")]
ZhCn,
#[serde(rename = "en-US")]
EnUs,
}
impl CommitPromptLocale {
pub fn default_prompt(self) -> &'static str {
match self {
Self::ZhCn => DEFAULT_COMMIT_PROMPT_ZH_CN.trim(),
Self::EnUs => DEFAULT_COMMIT_PROMPT_EN_US.trim(),
}
}
}
/// User preferences for Git commit message generation.
///
/// Empty `model_id` means 直连: forward the original Cursor RPC unchanged.
/// A non-empty value is the `model_hash` of a model configured on the Cursor
/// page, and the request is generated locally through that model.
/// A non-empty value is the stable identifier of a configured built-in or
/// plugin model, and the request is generated locally through that model.
#[derive(Clone, Debug, Default, Deserialize, PartialEq, Eq, Serialize)]
pub struct CommitSettings {
#[serde(default)]
pub model_id: String,
#[serde(default)]
pub prompt: String,
#[serde(default)]
pub prompt_locale: CommitPromptLocale,
}
impl CommitSettings {
@@ -104,7 +126,7 @@ impl CommitSettings {
pub fn effective_prompt(&self) -> &str {
let trimmed = self.prompt.trim();
if trimmed.is_empty() {
DEFAULT_COMMIT_PROMPT.trim()
self.prompt_locale.default_prompt()
} else {
trimmed
}
@@ -139,6 +161,32 @@ pub(crate) struct ProxySettingsSecret {
}
impl Store {
pub(crate) async fn cursor_takeover_enabled(&self) -> Result<bool> {
let value = sqlx::query_scalar::<_, String>(
"SELECT value_json FROM service_settings WHERE setting_key = ?",
)
.bind(CURSOR_TAKEOVER_ENABLED_KEY)
.fetch_optional(&self.pool)
.await?;
value
.map(|value| serde_json::from_str(&value).map_err(Into::into))
.unwrap_or(Ok(true))
}
pub(crate) async fn set_cursor_takeover_enabled(&self, enabled: bool) -> Result<()> {
let value_json = serde_json::to_string(&enabled)?;
let _write = self.writes.lock().await;
sqlx::query(
"INSERT INTO service_settings(setting_key, value_json, updated_at_ms) VALUES (?, ?, ?) ON CONFLICT(setting_key) DO UPDATE SET value_json = excluded.value_json, updated_at_ms = excluded.updated_at_ms",
)
.bind(CURSOR_TAKEOVER_ENABLED_KEY)
.bind(value_json)
.bind(now_ms())
.execute(&self.pool)
.await?;
Ok(())
}
pub(crate) async fn installation_id(&self) -> Result<String> {
let generated = uuid::Uuid::new_v4().to_string();
let _write = self.writes.lock().await;
@@ -348,6 +396,7 @@ impl Store {
let settings = CommitSettings {
model_id: settings.model_id.trim().to_owned(),
prompt: settings.prompt.trim().to_owned(),
prompt_locale: settings.prompt_locale,
};
let value_json = serde_json::to_string(&settings)?;
let _write = self.writes.lock().await;
@@ -365,7 +414,36 @@ impl Store {
#[cfg(test)]
mod tests {
use super::ProxyMode;
use super::{
CommitPromptLocale, CommitSettings, ProxyMode, DEFAULT_COMMIT_PROMPT_EN_US,
DEFAULT_COMMIT_PROMPT_ZH_CN,
};
#[test]
fn default_commit_prompt_follows_its_saved_locale() {
for (prompt_locale, expected) in [
(CommitPromptLocale::ZhCn, DEFAULT_COMMIT_PROMPT_ZH_CN),
(CommitPromptLocale::EnUs, DEFAULT_COMMIT_PROMPT_EN_US),
] {
let settings = CommitSettings {
prompt_locale,
..CommitSettings::default()
};
assert_eq!(settings.effective_prompt(), expected.trim());
}
}
#[test]
fn custom_commit_prompt_does_not_change_with_locale() {
for prompt_locale in [CommitPromptLocale::ZhCn, CommitPromptLocale::EnUs] {
let settings = CommitSettings {
prompt: "custom prompt".into(),
prompt_locale,
..CommitSettings::default()
};
assert_eq!(settings.effective_prompt(), "custom prompt");
}
}
#[test]
fn default_proxy_mode_uses_the_default_wire_value() {
+7 -2
View File
@@ -20,7 +20,7 @@ use cursor_server::{
model::{ContentPart, ModelConfigInput, ModelType, ProjectedContent, OPENAI_CHAT_ENDPOINT},
network::NetworkClients,
provider::{FinishReason, ModelEvent},
store::{CommitSettings, DEFAULT_COMMIT_PROMPT},
store::{CommitPromptLocale, CommitSettings, DEFAULT_COMMIT_PROMPT_ZH_CN},
};
use tower::ServiceExt;
@@ -101,6 +101,7 @@ async fn commit_message_is_generated_through_configured_model() {
.set_commit_settings(CommitSettings {
model_id: created.model_hash.clone(),
prompt: String::new(),
prompt_locale: CommitPromptLocale::ZhCn,
})
.await
.unwrap();
@@ -124,7 +125,7 @@ async fn commit_message_is_generated_through_configured_model() {
assert_eq!(requests.len(), 1);
assert_eq!(
requests[0].prompt.instructions,
DEFAULT_COMMIT_PROMPT.trim()
DEFAULT_COMMIT_PROMPT_ZH_CN.trim()
);
let ProjectedContent::Parts(parts) = &requests[0].history[0].content else {
panic!("expected user text parts");
@@ -147,6 +148,7 @@ async fn custom_prompt_and_model_from_commit_settings_are_used() {
.set_commit_settings(CommitSettings {
model_id: created.model_hash,
prompt: "自定义提交提示词".into(),
prompt_locale: CommitPromptLocale::ZhCn,
})
.await
.unwrap();
@@ -180,6 +182,7 @@ async fn empty_diffs_are_rejected_when_generating() {
.set_commit_settings(CommitSettings {
model_id: created.model_hash,
prompt: String::new(),
prompt_locale: CommitPromptLocale::ZhCn,
})
.await
.unwrap();
@@ -205,6 +208,7 @@ async fn tool_call_events_are_rejected() {
.set_commit_settings(CommitSettings {
model_id: created.model_hash,
prompt: String::new(),
prompt_locale: CommitPromptLocale::ZhCn,
})
.await
.unwrap();
@@ -231,6 +235,7 @@ async fn unconfigured_model_is_rejected() {
.set_commit_settings(CommitSettings {
model_id: "missing-hash".into(),
prompt: String::new(),
prompt_locale: CommitPromptLocale::ZhCn,
})
.await
.unwrap();