mirror of
https://wget.la/https://github.com/leookun/cursor-byok
synced 2026-10-03 18:23:51 +08:00
feat: update desktop settings and server compatibility
This commit is contained in:
@@ -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(<具体修改项>): 修改美术资源`
|
||||
@@ -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
@@ -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();
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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),
|
||||
|
||||
@@ -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<()> {
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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"))
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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());
|
||||
}
|
||||
|
||||
@@ -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());
|
||||
}
|
||||
}
|
||||
@@ -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());
|
||||
}
|
||||
}
|
||||
@@ -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;
|
||||
|
||||
@@ -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(""));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
)))
|
||||
}
|
||||
@@ -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");
|
||||
}
|
||||
|
||||
@@ -6,6 +6,7 @@ mod normalize;
|
||||
mod openai_chat;
|
||||
mod openai_responses;
|
||||
mod recorder;
|
||||
mod request_template;
|
||||
mod router;
|
||||
|
||||
use std::pin::Pin;
|
||||
|
||||
@@ -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
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -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());
|
||||
|
||||
@@ -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() {
|
||||
|
||||
@@ -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();
|
||||
|
||||
Reference in New Issue
Block a user