mirror of
https://wget.la/https://github.com/leookun/cursor-byok
synced 2026-10-03 18:23:51 +08:00
- 新增 CommitSettingsCard 组件,支持选择生成模型与编辑提示词 - 新增 /settings/commit GET/PUT 接口及 CommitSettings 持久化 - 实现 WriteGitCommitMessage RPC 本地生成,空 model_id 时直连转发 - 新增 NetworkService/IsConnected 探针响应,防止流式生成被中断 - 添加 commit prompt 模板及 proto 消息定义 - 补充 zh-CN / en-US 国际化词条
247 lines
8.2 KiB
Rust
247 lines
8.2 KiB
Rust
//! 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"));
|
|
}
|