mirror of
https://wget.la/https://github.com/leookun/cursor-byok
synced 2026-10-05 20:44:07 +08:00
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:
@@ -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"));
|
||||
}
|
||||
Reference in New Issue
Block a user