feat: 新增 Commit 设置及本地提交信息生成

- 新增 CommitSettingsCard 组件,支持选择生成模型与编辑提示词
- 新增 /settings/commit GET/PUT 接口及 CommitSettings 持久化
- 实现 WriteGitCommitMessage RPC 本地生成,空 model_id 时直连转发
- 新增 NetworkService/IsConnected 探针响应,防止流式生成被中断
- 添加 commit prompt 模板及 proto 消息定义
- 补充 zh-CN / en-US 国际化词条
This commit is contained in:
ProtectCookies
2026-09-03 10:31:41 +08:00
parent 8fdcdd7f84
commit 42811a27f5
19 changed files with 1292 additions and 63 deletions
+66
View File
@@ -0,0 +1,66 @@
# Git 提交信息生成指南
## 角色与目标
你是一个 git 提交信息生成器。给你一段 git diff,你只输出提交信息本身,不要有任何解释、前言、引号或额外文本。你的整个回复会被直接传给 `git commit`。
## 通用规则(生成标题时)
- 使用现在时(present tense),精准描述本次 diff 的关键改动。
- 关注「改了什么」,而不是罗列文件名。
- 要具体:包含具体细节(包名、版本、功能点),避免笼统表述。
- 排除任何不必要的内容(如翻译说明)。
- 标题最长不超过 50 个字符。
- 提交信息语言:简体中文(无论 diff 中的内容是什么语言,输出一律用简体中文)。
- 只输出提交信息文本本身,不要用引号或其他格式包裹,不要加解释或前言。
## 输出格式(按 type 选择其一)
| type | 格式模板 |
| ------------------ | ------------------------------------------------------------ |
| plain | `<commit message>` |
| conventional | `<type>[optional (<scope>)]: <commit message>`(主题须以小写字母开头) |
| conventional+body | `<type>[optional (<scope>)]: <commit message subject>`(主题须以小写字母开头,body 单独生成) |
| gitmoji | `:emoji: <commit message>` |
| subject+body | `<commit message subject>`(body 单独生成) |
输出必须严格符合所选 type 对应的格式。
## Conventional 类型选择
从下面的「类型-描述」中选择一个最贴合本次 diff 的类型。重要:类型必须全小写(例如 `feat`,而不是 `Feat` 或 `FEAT`)。
```json
{
"docs": "仅文档变更",
"style": "不影响代码含义的变更(空白、格式、缺失分号等)",
"refactor": "改善代码结构但不改变功能的变更(重命名、重构类/方法、抽取函数等)",
"perf": "提升性能的代码变更",
"test": "新增缺失的测试或修正已有测试",
"build": "影响构建系统或外部依赖的变更",
"ci": "对 CI 配置文件和脚本的变更",
"chore": "不修改 src 或 test 文件的其他变更",
"revert": "回退某个之前的提交",
"feat": "新功能",
"fix": "缺陷修复"
}
```
- `conventional`:直接按上表选择类型并输出完整主题行。
- `conventional+body`:只输出 conventional 主题行,body 会单独生成。
## 描述(body)生成规则
当已有提交标题、需要生成描述时,给你标题与 diff,你只输出提交描述正文:
- 简洁:使用 3–6 条要点(每条一行短句),或 2–4 句短句,不要长段落。
- 用现在时聚焦「改了什么、为什么」。
- 每行最多 72 个字符;当某条要点换行时,续行缩进 2 个空格,与要点文字对齐。
- 不要重复标题,不要元评论(如「本次提交……」)。
- 语言:简体中文。
- 只输出提交描述正文,不要有其他内容。
- 修改点要列举清除明白,不能笼统的说「更新功能,修改资源」这种概括性描述。
- 提交必须要有前缀,提交信息不允许有emoji,如果存在多个提交功能,比如美化和bug修改,分开写多个 例如:
- bugfix(*具体修改项*): 修复xxxbug问题
- chore(*具体修改项*): 修改美术资源
+26 -1
View File
@@ -19,7 +19,7 @@ use crate::{
connect,
proto::{agent::v1 as agent, aiserver::v1 as ai},
},
services::{account, analytics, knowledge, model_catalog, tab},
services::{account, analytics, commit_message, knowledge, model_catalog, tab},
transport::{TransportParent, TransportRegistry},
},
Result,
@@ -56,6 +56,14 @@ fn router_with_proxy(
"/aiserver.v1.AiService/GetUsableModels",
post(model_catalog::usable_models),
)
.route(
"/aiserver.v1.AiService/WriteGitCommitMessage",
post(commit_message::write_git_commit_message),
)
.route(
"/aiserver.v1.NetworkService/IsConnected",
post(is_connected),
)
.route(
"/aiserver.v1.AuthService/GetEmail",
post(account::get_email),
@@ -113,6 +121,23 @@ async fn health() -> StatusCode {
StatusCode::NO_CONTENT
}
/// `NetworkService/IsConnected` probe. Cursor's always-local extension checks
/// connectivity roughly 10s after any slow request starts; a 404/error here is
/// treated as "network disconnected" and aborts in-flight work (e.g. commit
/// message generation) even while the model is still streaming. Always answer
/// connected with an empty `IsConnectedResponse` so local BYOK generation is
/// never cancelled by this probe.
async fn is_connected() -> Result<Response<Body>> {
let payload = connect::encode_message(&ai::IsConnectedResponse {})?;
let mut response = Response::new(Body::from(payload));
*response.status_mut() = StatusCode::OK;
response.headers_mut().insert(
header::CONTENT_TYPE,
HeaderValue::from_static("application/proto"),
);
Ok(response)
}
async fn run_sse_handler(
State(registry): State<TransportRegistry>,
Extension(proxy): Extension<CursorProxy>,
+4
View File
@@ -200,6 +200,10 @@ pub fn api_router(service: ControlService) -> Router {
"/__byok-api__/api/settings/desktop",
get(settings::get_desktop).put(settings::update_desktop),
)
.route(
"/__byok-api__/api/settings/commit",
get(settings::get_commit).put(settings::update_commit),
)
.route(
"/__byok-api__/api/harness/cursor/status",
get(harness::status),
+10 -2
View File
@@ -28,8 +28,8 @@ use crate::{
plugin::{PluginDescriptor, PluginRegistry, PluginRuntime, PluginRuntimeStatus},
provider::{is_valid_response_event, ModelEvent, Provider},
store::{
DesktopSettings, PortSettings, ProxySettings, ProxySettingsInput, StatisticsStorage, Store,
TabSettings,
CommitSettings, DesktopSettings, PortSettings, ProxySettings, ProxySettingsInput,
StatisticsStorage, Store, TabSettings,
},
Error, Result,
};
@@ -703,6 +703,14 @@ impl ControlService {
pub async fn set_desktop_settings(&self, settings: DesktopSettings) -> Result<()> {
self.store.set_desktop_settings(settings).await
}
pub async fn commit_settings(&self) -> Result<CommitSettings> {
self.store.commit_settings().await
}
pub async fn set_commit_settings(&self, settings: CommitSettings) -> Result<CommitSettings> {
self.store.set_commit_settings(settings).await
}
}
fn official_call(trace: CursorRunTraceSummary) -> CallSummary {
+36 -3
View File
@@ -1,11 +1,11 @@
//! Implements settings management endpoints.
use crate::Result;
use axum::{extract::State, Json};
use serde::Deserialize;
use serde::{Deserialize, Serialize};
use crate::store::{
DesktopSettings, PortSettings, ProxySettings, ProxySettingsInput, StatisticsStorage,
StatisticsStorageScope, TabSettings,
CommitSettings, DesktopSettings, PortSettings, ProxySettings, ProxySettingsInput,
StatisticsStorage, StatisticsStorageScope, TabSettings, DEFAULT_COMMIT_PROMPT,
};
use super::{ControlService, ObservabilitySettings};
@@ -87,3 +87,36 @@ pub async fn update_desktop(
service.set_desktop_settings(settings).await?;
get_desktop(State(service)).await
}
/// 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".
#[derive(Serialize)]
pub struct CommitSettingsView {
pub model_id: String,
pub prompt: String,
pub default_prompt: &'static str,
}
impl From<CommitSettings> for CommitSettingsView {
fn from(settings: CommitSettings) -> Self {
Self {
model_id: settings.model_id,
prompt: settings.prompt,
default_prompt: DEFAULT_COMMIT_PROMPT.trim(),
}
}
}
pub async fn get_commit(State(service): State<ControlService>) -> Result<Json<CommitSettingsView>> {
let settings = service.commit_settings().await?;
Ok(Json(CommitSettingsView::from(settings)))
}
pub async fn update_commit(
State(service): State<ControlService>,
Json(settings): Json<CommitSettings>,
) -> Result<Json<CommitSettingsView>> {
let saved = service.set_commit_settings(settings).await?;
Ok(Json(CommitSettingsView::from(saved)))
}
+35
View File
@@ -26,9 +26,44 @@ pub mod aiserver {
pub data_binary: Vec<u8>,
}
/// Commit message generation request. Only the fields the local
/// generator consumes are decoded; credentials and heavyweight context
/// fields are intentionally left to prost's unknown-field skipping.
#[derive(Clone, PartialEq, ::prost::Message)]
pub struct WriteGitCommitMessageRequest {
#[prost(string, repeated, tag = "1")]
pub diffs: Vec<String>,
#[prost(string, repeated, tag = "2")]
pub previous_commit_messages: Vec<String>,
#[prost(message, optional, tag = "3")]
pub explicit_context: Option<ExplicitContext>,
}
#[derive(Clone, PartialEq, ::prost::Message)]
pub struct ExplicitContext {
#[prost(string, tag = "1")]
pub context: String,
#[prost(string, optional, tag = "2")]
pub repo_context: Option<String>,
}
#[derive(Clone, PartialEq, ::prost::Message)]
pub struct WriteGitCommitMessageResponse {
#[prost(string, tag = "1")]
pub commit_message: String,
}
#[derive(Clone, Copy, PartialEq, ::prost::Message)]
pub struct BidiAppendResponse {}
/// `NetworkService/IsConnected` reply. The Cursor extension probes this
/// ~10s after any slow request starts; a non-OK result is treated as
/// "network disconnected" and aborts in-flight work (e.g. commit message
/// generation) even while the model is still streaming, so it always
/// answers as connected.
#[derive(Clone, Copy, PartialEq, ::prost::Message)]
pub struct IsConnectedResponse {}
#[derive(Clone, PartialEq, ::prost::Message)]
pub struct CustomErrorDetails {
#[prost(string, tag = "1")]
@@ -0,0 +1,406 @@
//! Generates Git commit messages locally for Cursor's SCM action.
//!
//! 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
//! message, the customizable commit prompt is the system prompt, and the raw
//! completion is cleaned before being returned.
use std::{
sync::Arc,
time::{Duration, Instant},
};
use axum::{
body::{to_bytes, Body},
extract::{Extension, State},
http::{header, HeaderValue, Request, Response, StatusCode},
};
use futures_util::StreamExt;
use prost::Message;
use tokio_util::sync::CancellationToken;
use crate::{
api::cursor::proxy::{self, CursorProxy},
cursor::{
protocol::{connect, proto::aiserver::v1 as ai},
transport::TransportRegistry,
},
model::{
ContentPart, ModelInvocation, ModelRequest, ModelSpec, ProjectedContent, ProjectedMessage,
PromptSpec, Role,
},
provider::{ModelEvent, Provider},
store::CommitSettings,
Error, Result,
};
const DIFF_TOTAL_LIMIT: usize = 40_000;
const DIFF_SINGLE_LIMIT: usize = 16_000;
const PREVIOUS_COMMIT_LIMIT: usize = 12;
const EXPLICIT_CONTEXT_LIMIT: usize = 20_000;
const GENERATION_TIMEOUT: Duration = Duration::from_secs(180);
pub async fn write_git_commit_message(
State(registry): State<TransportRegistry>,
Extension(upstream): Extension<CursorProxy>,
request: Request<Body>,
) -> Result<Response<Body>> {
let settings = registry.store().commit_settings().await?;
if settings.is_direct() {
return forward_direct(&registry, upstream, request).await;
}
generate_local(&registry, request, settings).await
}
async fn forward_direct(
registry: &TransportRegistry,
upstream: CursorProxy,
request: Request<Body>,
) -> Result<Response<Body>> {
let settings = registry.store().tab_settings().await?;
match settings.service_url() {
Some(service_url) => proxy::forward_to_service(&upstream, request, service_url).await,
None => proxy::forward(Extension(upstream), request).await,
}
}
async fn generate_local(
registry: &TransportRegistry,
request: Request<Body>,
settings: CommitSettings,
) -> Result<Response<Body>> {
let (parts, body) = request.into_parts();
let connect_timeout_ms = parts
.headers
.get("connect-timeout-ms")
.and_then(|value| value.to_str().ok())
.and_then(|value| value.parse::<u64>().ok());
tracing::info!(?connect_timeout_ms, "write git commit message received");
let body = to_bytes(body, usize::MAX)
.await
.map_err(|error| Error::Protocol(format!("cannot read request body: {error}")))?;
let request: ai::WriteGitCommitMessageRequest = connect::decode_unary(&body)?;
let diffs = truncate_diffs(&request.diffs, DIFF_TOTAL_LIMIT, DIFF_SINGLE_LIMIT);
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 invocation = build_invocation(
&settings,
&model.model_hash,
build_user_content(&request, &diffs),
);
let provider = registry.conversations().dependencies().provider.clone();
let generated = generate(
provider,
invocation,
connect_timeout_ms.map(Duration::from_millis),
)
.await?;
let commit_message = clean_generated_commit_message(&generated);
if commit_message.is_empty() {
return Err(Error::Provider("generated commit message is empty".into()));
}
let payload = ai::WriteGitCommitMessageResponse { commit_message }.encode_to_vec();
let mut response = Response::new(Body::from(payload));
*response.status_mut() = StatusCode::OK;
response.headers_mut().insert(
header::CONTENT_TYPE,
HeaderValue::from_static("application/proto"),
);
Ok(response)
}
fn build_invocation(
settings: &CommitSettings,
model_hash: &str,
user_content: String,
) -> ModelInvocation {
let call_id = format!("commit-message-{}", uuid::Uuid::new_v4());
ModelInvocation {
call_id: call_id.clone(),
run_id: call_id.clone(),
conversation_id: call_id,
provider_call_index: 0,
request: ModelRequest {
prompt: PromptSpec {
instructions: settings.effective_prompt().to_owned(),
tools: Vec::new(),
},
model: ModelSpec::new(model_hash.to_owned()),
history: vec![ProjectedMessage {
message_id: "commit-message".into(),
role: Role::User,
content: ProjectedContent::Parts(vec![ContentPart::Text { text: user_content }]),
}],
},
}
}
async fn generate(
provider: Arc<dyn Provider>,
invocation: ModelInvocation,
client_timeout: Option<Duration>,
) -> Result<String> {
let cancellation = CancellationToken::new();
let stream = provider.stream(invocation, cancellation.clone());
let mut accumulated = String::new();
let soft_deadline = client_timeout
.map(|timeout| timeout.saturating_sub(Duration::from_millis(700)))
.filter(|deadline| !deadline.is_zero());
let deadline = Instant::now()
+ soft_deadline
.unwrap_or(GENERATION_TIMEOUT)
.min(GENERATION_TIMEOUT);
let mut deadline_hit = false;
let completed = tokio::time::timeout(GENERATION_TIMEOUT, async {
futures_util::pin_mut!(stream);
let mut finished = false;
loop {
let wait = deadline.saturating_duration_since(Instant::now());
if wait.is_zero() {
deadline_hit = true;
break;
}
match tokio::time::timeout(wait, stream.next()).await {
Err(_) => {
deadline_hit = true;
break;
}
Ok(None) => break,
Ok(Some(event)) => match event? {
ModelEvent::TextDelta(delta) => accumulated.push_str(&delta),
ModelEvent::ToolCallStart { .. } => {
return Err(Error::Provider(
"commit message generation must not invoke tools".into(),
));
}
ModelEvent::Done(_) => {
finished = true;
break;
}
_ => {}
},
}
}
if !finished && !deadline_hit {
return Err(Error::Provider(
"provider stream ended without Done during commit message generation".into(),
));
}
Ok(())
})
.await;
match completed {
Ok(result) => {
result?;
if accumulated.trim().is_empty() {
return Err(Error::Provider("generated commit message is empty".into()));
}
Ok(accumulated)
}
Err(_) => {
cancellation.cancel();
Err(Error::Provider(
"commit message generation timed out".into(),
))
}
}
}
fn build_user_content(request: &ai::WriteGitCommitMessageRequest, diffs: &[String]) -> String {
let mut sections = vec!["Generate a Git commit message for the following changes.".to_owned()];
let previous =
truncate_previous_commits(&request.previous_commit_messages, PREVIOUS_COMMIT_LIMIT);
if !previous.is_empty() {
sections.push(format!("Recent commit messages:\n{}", previous.join("\n")));
}
if let Some(context) = &request.explicit_context {
let context_json = explicit_context_json(context);
if !context_json.is_empty() {
sections.push(format!("Explicit context:\n{context_json}"));
}
}
let diff_sections: Vec<String> = diffs
.iter()
.enumerate()
.map(|(index, diff)| format!("--- Diff {} ---\n{}", index + 1, diff))
.collect();
sections.push(format!("Diffs:\n{}", diff_sections.join("\n\n")));
sections.join("\n\n")
}
fn explicit_context_json(context: &ai::ExplicitContext) -> String {
let context_text = context.context.trim();
let repo_context = context
.repo_context
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty());
if context_text.is_empty() && repo_context.is_none() {
return String::new();
}
let mut fields = serde_json::Map::new();
if !context_text.is_empty() {
fields.insert(
"context".into(),
serde_json::Value::String(context_text.to_owned()),
);
}
if let Some(repo_context) = repo_context {
fields.insert(
"repo_context".into(),
serde_json::Value::String(repo_context.to_owned()),
);
}
truncate_text(
&serde_json::to_string(&serde_json::Value::Object(fields)).unwrap_or_default(),
EXPLICIT_CONTEXT_LIMIT,
)
}
fn truncate_diffs(input: &[String], total_limit: usize, single_limit: usize) -> Vec<String> {
let mut result = Vec::new();
let mut remaining = total_limit;
for raw in input {
let diff = raw.trim();
if diff.is_empty() || remaining == 0 {
continue;
}
let truncated = truncate_text(diff, remaining.min(single_limit));
if truncated.is_empty() {
continue;
}
remaining = remaining.saturating_sub(truncated.chars().count());
result.push(truncated);
}
result
}
fn truncate_previous_commits(input: &[String], limit: usize) -> Vec<String> {
let mut result = Vec::new();
for raw in input {
let value = raw.trim();
if value.is_empty() {
continue;
}
result.push(format!("- {value}"));
if result.len() >= limit {
break;
}
}
result
}
fn truncate_text(value: &str, limit: usize) -> String {
let trimmed = value.trim();
if trimmed.is_empty() || limit == 0 {
return String::new();
}
if trimmed.chars().count() <= limit {
return trimmed.to_owned();
}
let truncated: String = trimmed.chars().take(limit).collect();
format!("{}\n...[truncated]", truncated.trim_end())
}
fn clean_generated_commit_message(value: &str) -> String {
let mut result = strip_code_fence(value.trim());
const PREFIXES: [&str; 3] = ["commit message:", "git commit message:", "message:"];
loop {
let lower = result.trim().to_ascii_lowercase();
let Some(prefix) = PREFIXES.iter().find(|prefix| lower.starts_with(*prefix)) else {
break;
};
result = result.trim()[prefix.len()..].to_owned();
}
let result = result.trim().to_owned();
if result.lines().all(|line| line.trim().is_empty()) {
String::new()
} else {
result
}
}
fn strip_code_fence(value: &str) -> String {
let trimmed = value.trim();
if !trimmed.starts_with("```") {
return trimmed.to_owned();
}
let mut lines = trimmed.lines();
if lines.next().is_none() {
return trimmed.to_owned();
}
let mut body: Vec<&str> = lines.collect();
if body
.last()
.is_some_and(|line| line.trim_start().starts_with("```"))
{
body.pop();
}
body.join("\n").trim().to_owned()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn empty_model_id_is_direct() {
assert!(CommitSettings::default().is_direct());
assert!(CommitSettings {
model_id: " ".into(),
prompt: String::new(),
}
.is_direct());
assert!(!CommitSettings {
model_id: "abc".into(),
prompt: String::new(),
}
.is_direct());
}
#[test]
fn truncation_limits_apply_per_diff_and_in_total() {
let diffs = vec!["a".repeat(20_000), "b".repeat(20_000), "c".repeat(20_000)];
let truncated = truncate_diffs(&diffs, DIFF_TOTAL_LIMIT, DIFF_SINGLE_LIMIT);
assert_eq!(truncated.len(), 3);
let total: usize = truncated.iter().map(|diff| diff.chars().count()).sum();
assert!(total <= DIFF_TOTAL_LIMIT + 3 * "\n...[truncated]".len());
}
#[test]
fn cleaning_strips_fences_and_prefixes() {
let raw = "```\nCommit message: fix: 修复登录超时问题\n```";
assert_eq!(clean_generated_commit_message(raw), "fix: 修复登录超时问题");
}
#[test]
fn cleaning_returns_empty_for_blank_output() {
assert_eq!(clean_generated_commit_message(" \n "), "");
}
#[test]
fn explicit_context_drops_empty_fields() {
let empty = explicit_context_json(&ai::ExplicitContext {
context: " ".into(),
repo_context: None,
});
assert_eq!(empty, "");
let filled = explicit_context_json(&ai::ExplicitContext {
context: "背景".into(),
repo_context: Some("repo".into()),
});
assert_eq!(filled, "{\"context\":\"背景\",\"repo_context\":\"repo\"}");
}
}
+1
View File
@@ -3,6 +3,7 @@
pub mod account;
pub mod analytics;
pub mod blob_sync;
pub mod commit_message;
pub mod context_sync;
pub mod knowledge;
pub mod model_catalog;
+1 -2
View File
@@ -9,7 +9,7 @@ use axum::{
use crate::{api::cursor::proxy, cursor::transport::TransportRegistry, Result};
pub const TAB_PATHS: [&str; 17] = [
pub const TAB_PATHS: [&str; 16] = [
"/aiserver.v1.AiService/StreamCpp",
"/aiserver.v1.AiService/StreamNextCursorPrediction",
"/aiserver.v1.AiService/GetCppEditClassification",
@@ -19,7 +19,6 @@ pub const TAB_PATHS: [&str; 17] = [
"/aiserver.v1.AiService/CppAppend",
"/aiserver.v1.AiService/CppEditHistoryAppend",
"/aiserver.v1.AiService/ReportAiCodeChangeMetrics",
"/aiserver.v1.AiService/WriteGitCommitMessage",
"/aiserver.v1.AiService/WriteGitBranchName",
"/aiserver.v1.CppService/AvailableModels",
"/aiserver.v1.CppService/RecordCppFate",
+2
View File
@@ -165,6 +165,8 @@ fn is_local_path(path: &str) -> bool {
| "/aiserver.v1.AiService/KnowledgeBaseList"
| "/aiserver.v1.AiService/KnowledgeBaseUpdate"
| "/aiserver.v1.AiService/KnowledgeBaseRemove"
| "/aiserver.v1.AiService/WriteGitCommitMessage"
| "/aiserver.v1.NetworkService/IsConnected"
| "/aiserver.v1.AnalyticsService/BootstrapStatsig"
| "/auth/full_stripe_profile"
)
+62
View File
@@ -10,6 +10,10 @@ const PROXY_SETTINGS_KEY: &str = "outbound_proxy";
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";
/// Embedded default system prompt for commit message generation.
pub const DEFAULT_COMMIT_PROMPT: &str = include_str!("../../prompt/cursor/commit/prompt.md");
pub const PUBLIC_TAB_SERVICE_URL: &str = "https://tab.leokun.cn";
@@ -79,6 +83,34 @@ impl TabSettings {
}
}
/// 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.
#[derive(Clone, Debug, Default, Deserialize, PartialEq, Eq, Serialize)]
pub struct CommitSettings {
#[serde(default)]
pub model_id: String,
#[serde(default)]
pub prompt: String,
}
impl CommitSettings {
pub fn is_direct(&self) -> bool {
self.model_id.trim().is_empty()
}
pub fn effective_prompt(&self) -> &str {
let trimmed = self.prompt.trim();
if trimmed.is_empty() {
DEFAULT_COMMIT_PROMPT.trim()
} else {
trimmed
}
}
}
#[derive(Clone, Debug, Default, Deserialize, PartialEq, Eq, Serialize)]
pub struct ProxySettingsInput {
pub mode: ProxyMode,
@@ -299,4 +331,34 @@ impl Store {
.await?;
Ok(())
}
pub async fn commit_settings(&self) -> Result<CommitSettings> {
let value = sqlx::query_scalar::<_, String>(
"SELECT value_json FROM service_settings WHERE setting_key = ?",
)
.bind(COMMIT_SETTINGS_KEY)
.fetch_optional(&self.pool)
.await?;
value
.map(|value| serde_json::from_str(&value).map_err(Into::into))
.unwrap_or_else(|| Ok(CommitSettings::default()))
}
pub async fn set_commit_settings(&self, settings: CommitSettings) -> Result<CommitSettings> {
let settings = CommitSettings {
model_id: settings.model_id.trim().to_owned(),
prompt: settings.prompt.trim().to_owned(),
};
let value_json = serde_json::to_string(&settings)?;
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(COMMIT_SETTINGS_KEY)
.bind(value_json)
.bind(now_ms())
.execute(&self.pool)
.await?;
Ok(settings)
}
}
+246
View File
@@ -0,0 +1,246 @@
//! Verifies the local WriteGitCommitMessage engine end to end on the wire.
#[path = "support/fake_provider.rs"]
mod fake_provider;
#[path = "support/fixtures.rs"]
mod fixtures;
use std::sync::Arc;
use axum::{
body::{to_bytes, Body},
http::{header, Request, StatusCode},
};
use cursor_server::{
api::cursor,
cursor::{
prompting::{PromptAssets, PromptCompiler},
protocol::{connect, proto::aiserver::v1 as ai},
transport::TransportRegistry,
},
model::{ContentPart, ModelConfigInput, ModelType, ProjectedContent, OPENAI_CHAT_ENDPOINT},
network::NetworkClients,
provider::{FinishReason, ModelEvent},
store::{CommitSettings, DEFAULT_COMMIT_PROMPT},
};
use tower::ServiceExt;
async fn commit_router(
store: cursor_server::store::Store,
provider: fake_provider::FakeProvider,
) -> axum::Router {
let assets = PromptAssets::load(
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
.join("prompt/cursor")
.as_path(),
)
.unwrap();
let clients = NetworkClients::new(store.clone());
let registry = TransportRegistry::new(store, Arc::new(provider), PromptCompiler::new(assets));
cursor::router(registry, clients).unwrap()
}
fn model_input(model_id: &str) -> ModelConfigInput {
ModelConfigInput {
sort_order: 1,
display_name: "Qwen Flash".into(),
group_name: None,
model_type: ModelType::OpenAi,
base_url: "https://example.com/v1".into(),
use_full_url: false,
api_key: "test-key".into(),
tooltip_data: "模型介绍".into(),
model_id: model_id.into(),
reasoning_effort: None,
openai_endpoint: OPENAI_CHAT_ENDPOINT.into(),
openai_extra_params_enabled: false,
openai_extra_params: serde_json::json!({}),
custom_headers_enabled: false,
custom_headers: serde_json::json!({}),
anthropic_extra_params_enabled: false,
anthropic_extra_params: serde_json::json!({}),
context_window_tokens: None,
max_completion_tokens: None,
anthropic_max_tokens: None,
anthropic_thinking_effort: None,
thinking_budget_tokens: None,
}
}
async fn post_commit_message(
router: axum::Router,
request: ai::WriteGitCommitMessageRequest,
) -> axum::response::Response {
let body = connect::encode_message(&request).unwrap();
router
.oneshot(
Request::post("/aiserver.v1.AiService/WriteGitCommitMessage")
.header(header::CONTENT_TYPE, "application/proto")
.body(Body::from(body))
.unwrap(),
)
.await
.unwrap()
}
fn diff_request(diff: &str) -> ai::WriteGitCommitMessageRequest {
ai::WriteGitCommitMessageRequest {
diffs: vec![diff.into()],
previous_commit_messages: vec!["feat: 上一次提交".into()],
explicit_context: None,
}
}
#[tokio::test]
async fn commit_message_is_generated_through_configured_model() {
let (_directory, store) = fixtures::temp_store().await;
let created = store
.create_model(&model_input("qwen/qwen3-flash"))
.await
.unwrap();
store
.set_commit_settings(CommitSettings {
model_id: created.model_hash.clone(),
prompt: String::new(),
})
.await
.unwrap();
let provider = fake_provider::FakeProvider::default();
provider.push(vec![
ModelEvent::TextStart,
ModelEvent::TextDelta("```\nCommit message: feat: 新增提交引擎\n```".into()),
ModelEvent::TextEnd,
ModelEvent::Done(FinishReason::Stop),
]);
let router = commit_router(store, provider.clone()).await;
let response = post_commit_message(router, diff_request("diff --git a/engine.rs")).await;
assert_eq!(response.status(), StatusCode::OK);
let body = to_bytes(response.into_body(), 64 * 1024).await.unwrap();
let decoded: ai::WriteGitCommitMessageResponse = prost::Message::decode(&body[..]).unwrap();
assert_eq!(decoded.commit_message, "feat: 新增提交引擎");
let requests = provider.requests();
assert_eq!(requests.len(), 1);
assert_eq!(
requests[0].prompt.instructions,
DEFAULT_COMMIT_PROMPT.trim()
);
let ProjectedContent::Parts(parts) = &requests[0].history[0].content else {
panic!("expected user text parts");
};
let ContentPart::Text { text } = &parts[0] else {
panic!("expected text part");
};
assert!(text.contains("diff --git a/engine.rs"));
assert!(text.contains("- feat: 上一次提交"));
}
#[tokio::test]
async fn custom_prompt_and_model_from_commit_settings_are_used() {
let (_directory, store) = fixtures::temp_store().await;
let created = store
.create_model(&model_input("qwen/qwen3-coder"))
.await
.unwrap();
store
.set_commit_settings(CommitSettings {
model_id: created.model_hash,
prompt: "自定义提交提示词".into(),
})
.await
.unwrap();
let provider = fake_provider::FakeProvider::default();
provider.push(vec![
ModelEvent::TextDelta("chore: 清理旧代码".into()),
ModelEvent::Done(FinishReason::Stop),
]);
let router = commit_router(store, provider.clone()).await;
let response = post_commit_message(router, diff_request("diff --git a/old.rs")).await;
assert_eq!(response.status(), StatusCode::OK);
let body = to_bytes(response.into_body(), 64 * 1024).await.unwrap();
let decoded: ai::WriteGitCommitMessageResponse = prost::Message::decode(&body[..]).unwrap();
assert_eq!(decoded.commit_message, "chore: 清理旧代码");
assert_eq!(
provider.requests()[0].prompt.instructions,
"自定义提交提示词"
);
}
#[tokio::test]
async fn empty_diffs_are_rejected_when_generating() {
let (_directory, store) = fixtures::temp_store().await;
let created = store
.create_model(&model_input("qwen/qwen3-flash"))
.await
.unwrap();
store
.set_commit_settings(CommitSettings {
model_id: created.model_hash,
prompt: String::new(),
})
.await
.unwrap();
let provider = fake_provider::FakeProvider::default();
let router = commit_router(store, provider).await;
let response = post_commit_message(router, ai::WriteGitCommitMessageRequest::default()).await;
assert_eq!(response.status(), StatusCode::BAD_REQUEST);
let body = to_bytes(response.into_body(), 4096).await.unwrap();
let text = std::str::from_utf8(&body).unwrap();
assert!(text.contains("diffs are required"));
}
#[tokio::test]
async fn tool_call_events_are_rejected() {
let (_directory, store) = fixtures::temp_store().await;
let created = store
.create_model(&model_input("qwen/qwen3-flash"))
.await
.unwrap();
store
.set_commit_settings(CommitSettings {
model_id: created.model_hash,
prompt: String::new(),
})
.await
.unwrap();
let provider = fake_provider::FakeProvider::default();
provider.push(vec![ModelEvent::ToolCallStart {
index: 0,
call_id: "call-1".into(),
name: "shell".into(),
}]);
let router = commit_router(store, provider).await;
let response = post_commit_message(router, diff_request("diff --git a/x.rs")).await;
assert_eq!(response.status(), StatusCode::BAD_GATEWAY);
let body = to_bytes(response.into_body(), 4096).await.unwrap();
let text = std::str::from_utf8(&body).unwrap();
assert!(text.contains("must not invoke tools"));
}
#[tokio::test]
async fn unconfigured_model_is_rejected() {
let (_directory, store) = fixtures::temp_store().await;
store
.set_commit_settings(CommitSettings {
model_id: "missing-hash".into(),
prompt: String::new(),
})
.await
.unwrap();
let provider = fake_provider::FakeProvider::default();
let router = commit_router(store, provider).await;
let response = post_commit_message(router, diff_request("diff --git a/x.rs")).await;
assert_eq!(response.status(), StatusCode::BAD_GATEWAY);
let body = to_bytes(response.into_body(), 4096).await.unwrap();
let text = std::str::from_utf8(&body).unwrap();
assert!(text.contains("missing-hash"));
}