feat: add group name functionality to models

- Introduced a new `group_name` field in the model configuration to allow for custom provider-group display names.
- Updated the `CursorModelCards`, `CursorModelEditor`, and `CursorSettingsPage` components to support group settings.
- Enhanced the UI to include group settings options, allowing users to modify group names and associated configurations.
- Added localization strings for new group settings features in both English and Chinese.
- Implemented a database migration to add the `group_name` column to the model configurations.
This commit is contained in:
leookun
2026-08-30 23:28:05 +08:00
parent 76baa3b0e7
commit e7a1cca4c6
32 changed files with 1950 additions and 137 deletions
@@ -0,0 +1,4 @@
-- Custom provider-group display name shared by models with the same upstream host.
-- NULL means no custom name; the UI falls back to the base_url hostname and the
-- Cursor model picker badge falls back to the model type label.
ALTER TABLE model_configs ADD COLUMN group_name TEXT;
+27 -3
View File
@@ -19,7 +19,9 @@ use crate::{
connect,
proto::{agent::v1 as agent, aiserver::v1 as ai},
},
services::{account, analytics, model_catalog, observability::CursorTraceRecorder, tab},
services::{
account, analytics, knowledge, model_catalog, observability::CursorTraceRecorder, tab,
},
transport::{TransportParent, TransportRegistry},
},
Result,
@@ -27,10 +29,15 @@ use crate::{
pub fn router(registry: TransportRegistry) -> Result<Router> {
let proxy = CursorProxy::cursor(registry.store().clone())?;
Ok(router_with_proxy(registry, proxy))
let knowledge = knowledge::KnowledgeService::managed()?;
Ok(router_with_proxy(registry, proxy, knowledge))
}
fn router_with_proxy(registry: TransportRegistry, proxy: CursorProxy) -> Router {
fn router_with_proxy(
registry: TransportRegistry,
proxy: CursorProxy,
knowledge_service: knowledge::KnowledgeService,
) -> Router {
let web_cache = registry.web_cache().router();
Router::new()
.route("/__byok-api__/healthz", get(health))
@@ -69,6 +76,22 @@ fn router_with_proxy(registry: TransportRegistry, proxy: CursorProxy) -> Router
"/aiserver.v1.DashboardService/GetUsageLimitStatusAndActiveGrants",
post(account::usage_limit_status),
)
.route(
"/aiserver.v1.AiService/KnowledgeBaseAdd",
post(knowledge::add),
)
.route(
"/aiserver.v1.AiService/KnowledgeBaseList",
post(knowledge::list),
)
.route(
"/aiserver.v1.AiService/KnowledgeBaseUpdate",
post(knowledge::update),
)
.route(
"/aiserver.v1.AiService/KnowledgeBaseRemove",
post(knowledge::remove),
)
.route(
analytics::BOOTSTRAP_STATSIG_PATH,
post(analytics::bootstrap_statsig),
@@ -80,6 +103,7 @@ fn router_with_proxy(registry: TransportRegistry, proxy: CursorProxy) -> Router
.fallback(proxy::forward)
.method_not_allowed_fallback(proxy::forward)
.layer(Extension(proxy))
.layer(Extension(knowledge_service))
.with_state(registry)
.merge(web_cache)
}
+1
View File
@@ -56,6 +56,7 @@ impl App {
compiler,
WebCache::managed()?,
plugins.clone(),
crate::config::managed_data_dir()?.join("rules"),
);
let control =
control::ControlService::new(store.clone(), provider, plugin_runtime, plugins)?;
+76
View File
@@ -146,6 +146,37 @@ async fn decode_part<T: Message + Default>(
.map_err(|error| Error::Protocol(format!("invalid {name} context Blob: {error}")))
}
/// 把本地 md 规则目录(rules 服务的存储)合并进请求上下文,
/// 使 BYOK 运行在 IDE 未携带这些规则时也能消费它们。
/// 与 IDE 已发规则按内容去重;读取失败只告警,不影响运行。
pub fn merge_local_rules(context: &mut pb::RequestContext, rules_dir: &Path) {
let records = match crate::cursor::services::knowledge::RuleStore::open(rules_dir.into())
.and_then(|store| store.list())
{
Ok(records) => records,
Err(error) => {
tracing::warn!(%error, "cannot read local rules; continuing without them");
return;
}
};
let existing = context
.rules
.iter()
.chain(context.non_file_rules.iter())
.map(|rule| rule.content.trim().to_owned())
.chain(context.cloud_rule.iter().map(|rule| rule.trim().to_owned()))
.collect::<HashSet<_>>();
for record in records {
if record.knowledge.trim().is_empty() || existing.contains(record.knowledge.trim()) {
continue;
}
context.non_file_rules.push(pb::CursorRule {
content: record.knowledge,
..Default::default()
});
}
}
pub fn request_context(request: &pb::AgentRunRequest) -> Option<&pb::RequestContext> {
let action = request.action.as_ref()?;
action
@@ -563,3 +594,48 @@ fn xml(value: &str) -> String {
.replace('<', "&lt;")
.replace('>', "&gt;")
}
#[cfg(test)]
mod tests {
use super::*;
fn rule(content: &str) -> pb::CursorRule {
pb::CursorRule {
content: content.into(),
..Default::default()
}
}
#[test]
fn merge_local_rules_appends_and_dedupes_by_content() {
let directory = tempfile::tempdir().unwrap();
std::fs::write(directory.path().join("a.md"), "shared rule").unwrap();
std::fs::write(directory.path().join("b.md"), "local only rule").unwrap();
std::fs::write(directory.path().join("c.md"), " \n").unwrap();
let mut context = pb::RequestContext {
non_file_rules: vec![rule(" shared rule ")],
..Default::default()
};
merge_local_rules(&mut context, directory.path());
let contents = context
.non_file_rules
.iter()
.map(|rule| rule.content.as_str())
.collect::<Vec<_>>();
assert_eq!(
contents,
[" shared rule ", "local only rule"],
"IDE-sent duplicate is kept once and blank local rules are skipped"
);
}
#[test]
fn merge_local_rules_survives_a_missing_directory() {
let directory = tempfile::tempdir().unwrap();
let mut context = pb::RequestContext::default();
merge_local_rules(&mut context, &directory.path().join("nested/rules"));
assert!(context.non_file_rules.is_empty());
}
}
+7 -1
View File
@@ -51,6 +51,7 @@ pub(crate) struct PrepareDependencies<'a> {
pub checkpoint: &'a CheckpointBuilder,
pub blob_sync: &'a BlobSynchronizer,
pub context_sync: &'a RequestContextSynchronizer,
pub local_rules_dir: Option<&'a std::path::Path>,
}
pub(crate) async fn prepare(
@@ -64,6 +65,7 @@ pub(crate) async fn prepare(
checkpoint,
blob_sync,
context_sync,
local_rules_dir,
} = dependencies;
checkpoint
.import_prefetched(&request.pre_fetched_blobs)
@@ -119,7 +121,11 @@ pub(crate) async fn prepare(
.artifact("history_projection", "byok_server", &encoded, summary)
.await;
}
let request_context = context::hydrate(request, context_sync).await?;
let mut request_context = context::hydrate(request, context_sync).await?;
if let Some(rules_dir) = local_rules_dir {
context::merge_local_rules(&mut request_context, rules_dir);
}
let request_context = request_context;
let ActionProjection {
mode: mode_number,
mut turn_user,
@@ -26,6 +26,8 @@ pub(crate) struct ConversationDependencies {
pub provider: Arc<dyn Provider>,
pub compiler: PromptCompiler,
pub web_cache: WebCache,
/// 本地 rules 服务的 md 存储目录;编译请求上下文时合并其中的规则。
pub local_rules_dir: Option<std::path::PathBuf>,
}
struct RegistryInner {
@@ -47,6 +49,7 @@ impl ConversationRegistry {
provider: Arc<dyn Provider>,
compiler: PromptCompiler,
web_cache: WebCache,
local_rules_dir: Option<std::path::PathBuf>,
) -> Self {
Self {
inner: Arc::new(RegistryInner {
@@ -58,6 +61,7 @@ impl ConversationRegistry {
provider,
compiler,
web_cache,
local_rules_dir,
},
}),
}
@@ -471,6 +471,7 @@ fn spawn_run_request(
checkpoint: &checkpoint,
blob_sync: &blob_sync,
context_sync: &context_sync,
local_rules_dir: dependencies.local_rules_dir.as_deref(),
},
) => prepared,
};
+356
View File
@@ -0,0 +1,356 @@
//! Serves Cursor user rules: upstream-first with an offline markdown cache.
//!
//! 每个请求先回放离线日志再尝试上游;上游成功时把结果写穿到本地镜像,
//! 上游不可达时降级为本地 md 存储并记录日志等待回放。
mod store;
mod sync;
use std::sync::Arc;
use axum::{
body::{to_bytes, Body, Bytes},
extract::Extension,
http::{header, Request, Response},
};
use prost::Message;
use crate::{api::cursor::proxy, config, cursor::protocol::connect, Result};
pub(crate) use store::{RuleRecord, RuleStore};
#[derive(Clone, PartialEq, Message)]
pub(crate) struct KnowledgeBaseAddRequest {
#[prost(string, tag = "1")]
knowledge: String,
#[prost(string, tag = "2")]
title: String,
#[prost(string, tag = "3")]
git_origin: String,
#[prost(string, optional, tag = "4")]
composer_id: Option<String>,
}
#[derive(Clone, PartialEq, Message)]
pub(crate) struct KnowledgeBaseAddResponse {
#[prost(bool, tag = "1")]
success: bool,
#[prost(string, tag = "2")]
id: String,
}
#[derive(Clone, PartialEq, Message)]
pub(crate) struct KnowledgeBaseListRequest {
#[prost(int32, optional, tag = "1")]
limit: Option<i32>,
#[prost(string, optional, tag = "2")]
git_origin: Option<String>,
}
#[derive(Clone, PartialEq, Message)]
pub(crate) struct KnowledgeBaseListResponse {
#[prost(bool, tag = "1")]
success: bool,
#[prost(message, repeated, tag = "2")]
all_results: Vec<KnowledgeBaseListItem>,
}
#[derive(Clone, PartialEq, Message)]
pub(crate) struct KnowledgeBaseListItem {
#[prost(string, tag = "1")]
id: String,
#[prost(string, tag = "2")]
knowledge: String,
#[prost(string, tag = "3")]
title: String,
#[prost(string, tag = "4")]
created_at: String,
#[prost(bool, tag = "5")]
is_generated: bool,
}
#[derive(Clone, PartialEq, Message)]
pub(crate) struct KnowledgeBaseUpdateRequest {
#[prost(string, tag = "1")]
id: String,
#[prost(string, tag = "2")]
knowledge: String,
#[prost(string, tag = "3")]
title: String,
}
#[derive(Clone, PartialEq, Message)]
pub(crate) struct KnowledgeBaseUpdateResponse {
#[prost(bool, tag = "1")]
success: bool,
}
#[derive(Clone, PartialEq, Message)]
pub(crate) struct KnowledgeBaseRemoveRequest {
#[prost(string, tag = "1")]
id: String,
}
#[derive(Clone, PartialEq, Message)]
pub(crate) struct KnowledgeBaseRemoveResponse {
#[prost(bool, tag = "1")]
success: bool,
}
/// 规则存储与并发锁;经 axum Extension 注入四个 handler。
#[derive(Clone)]
pub struct KnowledgeService {
inner: Arc<Inner>,
}
struct Inner {
store: RuleStore,
lock: tokio::sync::Mutex<()>,
}
impl KnowledgeService {
pub fn managed() -> Result<Self> {
Self::with_root(config::managed_data_dir()?.join("rules"))
}
/// 指定存储根目录构造;managed() 与集成测试共用。
pub fn with_root(root: std::path::PathBuf) -> Result<Self> {
Ok(Self {
inner: Arc::new(Inner {
store: RuleStore::open(root)?,
lock: tokio::sync::Mutex::new(()),
}),
})
}
}
pub async fn add(
Extension(upstream): Extension<proxy::CursorProxy>,
Extension(service): Extension<KnowledgeService>,
request: Request<Body>,
) -> Result<Response<Body>> {
let (parts, body) = buffered(request).await?;
let message: KnowledgeBaseAddRequest = connect::decode_unary(&body)?;
let _guard = service.inner.lock.lock().await;
let store = &service.inner.store;
if sync::replay(&upstream, &parts.headers, store).await? {
match proxy::forward_buffered(&upstream, Request::from_parts(parts, Body::from(body))).await
{
Ok(response) if response.status.is_success() => {
if let Ok(reply) = connect::decode_unary::<KnowledgeBaseAddResponse>(&response.body)
{
if reply.success && !reply.id.is_empty() {
store.upsert(&RuleRecord {
id: reply.id,
knowledge: message.knowledge,
title: message.title,
created_at: now(),
is_generated: false,
git_origin: message.git_origin,
})?;
}
}
return Ok(response.into_response());
}
Ok(response) => {
tracing::warn!(status = %response.status, "rules upstream rejected add; storing locally");
}
Err(error) => {
tracing::warn!(%error, "rules upstream unavailable for add; storing locally");
}
}
}
let id = format!("{}{}", store::LOCAL_ID_PREFIX, uuid::Uuid::new_v4());
store.upsert(&RuleRecord {
id: id.clone(),
knowledge: message.knowledge,
title: message.title,
created_at: now(),
is_generated: false,
git_origin: message.git_origin,
})?;
store.record_add(&id)?;
proto(KnowledgeBaseAddResponse { success: true, id })
}
pub async fn list(
Extension(upstream): Extension<proxy::CursorProxy>,
Extension(service): Extension<KnowledgeService>,
request: Request<Body>,
) -> Result<Response<Body>> {
let (parts, body) = buffered(request).await?;
let message: KnowledgeBaseListRequest = connect::decode_unary(&body)?;
let _guard = service.inner.lock.lock().await;
let store = &service.inner.store;
let git_origin = message.git_origin.unwrap_or_default();
if sync::replay(&upstream, &parts.headers, store).await? {
match proxy::forward_buffered(&upstream, Request::from_parts(parts, Body::from(body))).await
{
Ok(response) if response.status.is_success() => {
if let Ok(reply) =
connect::decode_unary::<KnowledgeBaseListResponse>(&response.body)
{
// 带 git_origin 过滤的列表只是子集,整体覆盖会误删其他规则。
if reply.success && git_origin.is_empty() {
sync::mirror(store, reply.all_results)?;
}
}
return Ok(response.into_response());
}
Ok(response) => {
tracing::warn!(status = %response.status, "rules upstream rejected list; serving local cache");
}
Err(error) => {
tracing::warn!(%error, "rules upstream unavailable for list; serving local cache");
}
}
}
let mut records = store.list()?;
if !git_origin.is_empty() {
records.retain(|record| record.git_origin == git_origin);
}
if let Some(limit) = message.limit {
if limit >= 0 {
records.truncate(limit as usize);
}
}
proto(KnowledgeBaseListResponse {
success: true,
all_results: records
.into_iter()
.map(|record| KnowledgeBaseListItem {
id: record.id,
knowledge: record.knowledge,
title: record.title,
created_at: record.created_at,
is_generated: record.is_generated,
})
.collect(),
})
}
pub async fn update(
Extension(upstream): Extension<proxy::CursorProxy>,
Extension(service): Extension<KnowledgeService>,
request: Request<Body>,
) -> Result<Response<Body>> {
let (parts, body) = buffered(request).await?;
let message: KnowledgeBaseUpdateRequest = connect::decode_unary(&body)?;
let _guard = service.inner.lock.lock().await;
let store = &service.inner.store;
if sync::replay(&upstream, &parts.headers, store).await? {
match proxy::forward_buffered(&upstream, Request::from_parts(parts, Body::from(body))).await
{
Ok(response) if response.status.is_success() => {
if let Ok(reply) =
connect::decode_unary::<KnowledgeBaseUpdateResponse>(&response.body)
{
if reply.success {
let existing = store.get(&message.id)?;
store.upsert(&RuleRecord {
id: message.id,
knowledge: message.knowledge,
title: message.title,
created_at: existing
.as_ref()
.map_or_else(now, |record| record.created_at.clone()),
is_generated: existing
.as_ref()
.is_some_and(|record| record.is_generated),
git_origin: existing
.map(|record| record.git_origin)
.unwrap_or_default(),
})?;
}
}
return Ok(response.into_response());
}
Ok(response) => {
tracing::warn!(status = %response.status, "rules upstream rejected update; storing locally");
}
Err(error) => {
tracing::warn!(%error, "rules upstream unavailable for update; storing locally");
}
}
}
let Some(mut record) = store.get(&message.id)? else {
return proto(KnowledgeBaseUpdateResponse { success: false });
};
record.knowledge = message.knowledge;
record.title = message.title;
store.upsert(&record)?;
store.record_update(&message.id)?;
proto(KnowledgeBaseUpdateResponse { success: true })
}
pub async fn remove(
Extension(upstream): Extension<proxy::CursorProxy>,
Extension(service): Extension<KnowledgeService>,
request: Request<Body>,
) -> Result<Response<Body>> {
let (parts, body) = buffered(request).await?;
let message: KnowledgeBaseRemoveRequest = connect::decode_unary(&body)?;
let _guard = service.inner.lock.lock().await;
let store = &service.inner.store;
if sync::replay(&upstream, &parts.headers, store).await? {
match proxy::forward_buffered(&upstream, Request::from_parts(parts, Body::from(body))).await
{
Ok(response) if response.status.is_success() => {
if let Ok(reply) =
connect::decode_unary::<KnowledgeBaseRemoveResponse>(&response.body)
{
if reply.success {
store.remove(&message.id)?;
}
}
return Ok(response.into_response());
}
Ok(response) => {
tracing::warn!(status = %response.status, "rules upstream rejected remove; removing locally");
}
Err(error) => {
tracing::warn!(%error, "rules upstream unavailable for remove; removing locally");
}
}
}
store.remove(&message.id)?;
store.record_remove(&message.id)?;
proto(KnowledgeBaseRemoveResponse { success: true })
}
async fn buffered(request: Request<Body>) -> Result<(axum::http::request::Parts, Bytes)> {
let (parts, body) = request.into_parts();
let body = to_bytes(body, usize::MAX)
.await
.map_err(|error| crate::Error::Protocol(format!("cannot read request body: {error}")))?;
Ok((parts, body))
}
fn now() -> String {
chrono::Utc::now().to_rfc3339_opts(chrono::SecondsFormat::Millis, true)
}
fn proto(message: impl Message) -> Result<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,
axum::http::HeaderValue::from_static("application/proto"),
);
response.headers_mut().insert(
header::CONTENT_LENGTH,
length
.to_string()
.parse()
.expect("body length is always a valid header value"),
);
Ok(response)
}
@@ -0,0 +1,487 @@
//! Persists rules as markdown files with a JSON metadata sidecar.
use std::{
collections::BTreeMap,
path::{Path, PathBuf},
};
use serde::{Deserialize, Serialize};
use crate::{Error, Result};
const META_FILE: &str = "meta.json";
const RULE_EXTENSION: &str = "md";
pub const LOCAL_ID_PREFIX: &str = "local-";
/// 一条规则的完整视图:knowledge 来自 md 文件,其余字段来自 meta.json。
#[derive(Clone, Debug, PartialEq)]
pub struct RuleRecord {
pub id: String,
pub knowledge: String,
pub title: String,
pub created_at: String,
pub is_generated: bool,
pub git_origin: String,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum JournalOp {
Add,
Update,
Remove,
}
/// 离线期间未同步到上游的一次变更。
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub struct JournalEntry {
pub op: JournalOp,
pub id: String,
}
#[derive(Default, Serialize, Deserialize)]
struct Meta {
#[serde(default)]
rules: BTreeMap<String, RuleMeta>,
#[serde(default)]
journal: Vec<JournalEntry>,
}
#[derive(Clone, Default, Serialize, Deserialize)]
struct RuleMeta {
#[serde(default)]
title: String,
#[serde(default)]
created_at: String,
#[serde(default)]
is_generated: bool,
#[serde(default)]
git_origin: String,
}
/// md 文件为核心的规则存储;调用方需自行串行化并发访问。
pub struct RuleStore {
root: PathBuf,
}
impl RuleStore {
pub fn open(root: PathBuf) -> Result<Self> {
std::fs::create_dir_all(&root)?;
Ok(Self { root })
}
pub fn list(&self) -> Result<Vec<RuleRecord>> {
let meta = self.read_meta();
let mut records = Vec::new();
for entry in std::fs::read_dir(&self.root)? {
let path = entry?.path();
if path.extension().and_then(|value| value.to_str()) != Some(RULE_EXTENSION) {
continue;
}
let Some(id) = path.file_stem().and_then(|value| value.to_str()) else {
continue;
};
if validate_id(id).is_err() {
continue;
}
let knowledge = std::fs::read_to_string(&path)?;
records.push(assemble(id, knowledge, meta.rules.get(id), &path));
}
records.sort_by(|left, right| {
timestamp(&right.created_at)
.cmp(&timestamp(&left.created_at))
.then_with(|| left.id.cmp(&right.id))
});
Ok(records)
}
pub fn get(&self, id: &str) -> Result<Option<RuleRecord>> {
validate_id(id)?;
let path = self.rule_path(id);
let knowledge = match std::fs::read_to_string(&path) {
Ok(knowledge) => knowledge,
Err(error) if error.kind() == std::io::ErrorKind::NotFound => return Ok(None),
Err(error) => return Err(error.into()),
};
let meta = self.read_meta();
Ok(Some(assemble(id, knowledge, meta.rules.get(id), &path)))
}
pub fn upsert(&self, record: &RuleRecord) -> Result<()> {
validate_id(&record.id)?;
write_atomic(&self.rule_path(&record.id), record.knowledge.as_bytes())?;
let mut meta = self.read_meta();
meta.rules.insert(record.id.clone(), rule_meta(record));
self.write_meta(&meta)
}
pub fn remove(&self, id: &str) -> Result<()> {
validate_id(id)?;
remove_file_if_exists(&self.rule_path(id))?;
let mut meta = self.read_meta();
if meta.rules.remove(id).is_some() {
self.write_meta(&meta)?;
}
Ok(())
}
/// 离线新增的规则在上游落地后,把本地临时 id 换成上游分配的真实 id。
pub fn promote(&self, old_id: &str, new_id: &str) -> Result<()> {
validate_id(old_id)?;
validate_id(new_id)?;
let source = self.rule_path(old_id);
let target = self.rule_path(new_id);
#[cfg(windows)]
remove_file_if_exists(&target)?;
std::fs::rename(&source, &target)?;
let mut meta = self.read_meta();
if let Some(rule) = meta.rules.remove(old_id) {
meta.rules.insert(new_id.into(), rule);
}
for entry in &mut meta.journal {
if entry.id == old_id {
entry.id = new_id.into();
}
}
self.write_meta(&meta)
}
/// 用上游的完整列表覆盖本地镜像;仅应在日志为空(已全部回放)时调用。
pub fn replace_all(&self, records: &[RuleRecord]) -> Result<()> {
let mut meta = self.read_meta();
meta.rules.clear();
for record in records {
validate_id(&record.id)?;
write_atomic(&self.rule_path(&record.id), record.knowledge.as_bytes())?;
meta.rules.insert(record.id.clone(), rule_meta(record));
}
for entry in std::fs::read_dir(&self.root)? {
let path = entry?.path();
if path.extension().and_then(|value| value.to_str()) != Some(RULE_EXTENSION) {
continue;
}
let keep = path
.file_stem()
.and_then(|value| value.to_str())
.is_some_and(|id| meta.rules.contains_key(id));
if !keep {
remove_file_if_exists(&path)?;
}
}
self.write_meta(&meta)
}
pub fn journal_front(&self) -> Result<Option<JournalEntry>> {
Ok(self.read_meta().journal.first().cloned())
}
pub fn pop_journal(&self) -> Result<()> {
let mut meta = self.read_meta();
if !meta.journal.is_empty() {
meta.journal.remove(0);
self.write_meta(&meta)?;
}
Ok(())
}
pub fn record_add(&self, id: &str) -> Result<()> {
let mut meta = self.read_meta();
meta.journal.push(JournalEntry {
op: JournalOp::Add,
id: id.into(),
});
self.write_meta(&meta)
}
pub fn record_update(&self, id: &str) -> Result<()> {
let mut meta = self.read_meta();
if journal_contains(&meta.journal, id, JournalOp::Add) {
// 回放 add 时会读取最新内容,无需单独的 update 日志。
return Ok(());
}
let op = if id.starts_with(LOCAL_ID_PREFIX) {
// 本地临时 id 没有对应的 add 日志(如镜像覆盖后的残留),按新增回放。
JournalOp::Add
} else {
JournalOp::Update
};
if !journal_contains(&meta.journal, id, op) {
meta.journal.push(JournalEntry { op, id: id.into() });
self.write_meta(&meta)?;
}
Ok(())
}
pub fn record_remove(&self, id: &str) -> Result<()> {
let mut meta = self.read_meta();
let never_synced = journal_contains(&meta.journal, id, JournalOp::Add);
meta.journal.retain(|entry| entry.id != id);
if !never_synced && !id.starts_with(LOCAL_ID_PREFIX) {
meta.journal.push(JournalEntry {
op: JournalOp::Remove,
id: id.into(),
});
}
self.write_meta(&meta)
}
fn rule_path(&self, id: &str) -> PathBuf {
self.root.join(format!("{id}.{RULE_EXTENSION}"))
}
fn meta_path(&self) -> PathBuf {
self.root.join(META_FILE)
}
fn read_meta(&self) -> Meta {
match std::fs::read(self.meta_path()) {
Ok(bytes) => serde_json::from_slice(&bytes).unwrap_or_else(|error| {
tracing::warn!(%error, "rules meta.json is corrupt; starting from empty metadata");
Meta::default()
}),
Err(_) => Meta::default(),
}
}
fn write_meta(&self, meta: &Meta) -> Result<()> {
write_atomic(&self.meta_path(), &serde_json::to_vec_pretty(meta)?)
}
}
fn assemble(id: &str, knowledge: String, meta: Option<&RuleMeta>, path: &Path) -> RuleRecord {
match meta {
Some(meta) => RuleRecord {
id: id.into(),
knowledge,
title: meta.title.clone(),
created_at: meta.created_at.clone(),
is_generated: meta.is_generated,
git_origin: meta.git_origin.clone(),
},
// 用户手放的 md 文件没有元数据,用文件名当标题、修改时间当创建时间。
None => RuleRecord {
id: id.into(),
knowledge,
title: id.into(),
created_at: file_modified_at(path),
is_generated: false,
git_origin: String::new(),
},
}
}
fn rule_meta(record: &RuleRecord) -> RuleMeta {
RuleMeta {
title: record.title.clone(),
created_at: record.created_at.clone(),
is_generated: record.is_generated,
git_origin: record.git_origin.clone(),
}
}
fn journal_contains(journal: &[JournalEntry], id: &str, op: JournalOp) -> bool {
journal.iter().any(|entry| entry.id == id && entry.op == op)
}
fn timestamp(created_at: &str) -> i64 {
chrono::DateTime::parse_from_rfc3339(created_at)
.map(|time| time.timestamp_millis())
.unwrap_or(0)
}
fn file_modified_at(path: &Path) -> String {
let modified = std::fs::metadata(path)
.and_then(|meta| meta.modified())
.unwrap_or_else(|_| std::time::SystemTime::now());
chrono::DateTime::<chrono::Utc>::from(modified)
.to_rfc3339_opts(chrono::SecondsFormat::Millis, true)
}
fn validate_id(id: &str) -> Result<()> {
if id.is_empty()
|| id.len() > 128
|| !id
.bytes()
.all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'-' | b'_'))
{
return Err(Error::Protocol(format!("invalid rule id: {id:?}")));
}
Ok(())
}
fn remove_file_if_exists(path: &Path) -> Result<()> {
match std::fs::remove_file(path) {
Ok(()) => Ok(()),
Err(error) if error.kind() == std::io::ErrorKind::NotFound => Ok(()),
Err(error) => Err(error.into()),
}
}
fn write_atomic(path: &Path, bytes: &[u8]) -> Result<()> {
use std::io::Write;
let directory = path.parent().expect("rule path has a parent");
let temporary = directory.join(format!(".{}.tmp", uuid::Uuid::new_v4()));
let mut file = std::fs::File::create(&temporary)?;
file.write_all(bytes)?;
file.sync_all()?;
drop(file);
#[cfg(windows)]
remove_file_if_exists(path)?;
std::fs::rename(&temporary, path).inspect_err(|_| {
let _ = std::fs::remove_file(&temporary);
})?;
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
fn record(id: &str, knowledge: &str, created_at: &str) -> RuleRecord {
RuleRecord {
id: id.into(),
knowledge: knowledge.into(),
title: format!("title-{id}"),
created_at: created_at.into(),
is_generated: false,
git_origin: String::new(),
}
}
fn journal(store: &RuleStore) -> Vec<JournalEntry> {
store.read_meta().journal
}
#[test]
fn upserts_lists_and_removes_rules() {
let root = tempfile::tempdir().unwrap();
let store = RuleStore::open(root.path().join("rules")).unwrap();
store
.upsert(&record("100", "older", "2026-01-01T00:00:00.000Z"))
.unwrap();
store
.upsert(&record("200", "newer", "2026-02-01T00:00:00.000Z"))
.unwrap();
let listed = store.list().unwrap();
assert_eq!(
listed
.iter()
.map(|rule| rule.id.as_str())
.collect::<Vec<_>>(),
["200", "100"],
"list is sorted by created_at descending"
);
assert_eq!(listed[0].knowledge, "newer");
assert_eq!(listed[0].title, "title-200");
store.remove("200").unwrap();
assert!(store.get("200").unwrap().is_none());
assert_eq!(store.list().unwrap().len(), 1);
}
#[test]
fn rejects_path_traversal_ids() {
let root = tempfile::tempdir().unwrap();
let store = RuleStore::open(root.path().join("rules")).unwrap();
assert!(store.get("../escape").is_err());
assert!(store.get("a/b").is_err());
assert!(store.get("").is_err());
}
#[test]
fn compacts_offline_journal() {
let root = tempfile::tempdir().unwrap();
let store = RuleStore::open(root.path().join("rules")).unwrap();
// 离线新增后再更新:回放 add 即可携带最新内容,不产生 update 日志。
store
.upsert(&record("local-a", "v1", "2026-01-01T00:00:00.000Z"))
.unwrap();
store.record_add("local-a").unwrap();
store.record_update("local-a").unwrap();
assert_eq!(
journal(&store),
vec![JournalEntry {
op: JournalOp::Add,
id: "local-a".into()
}]
);
// 离线新增后又删除:上游从未见过它,日志清空。
store.record_remove("local-a").unwrap();
assert!(journal(&store).is_empty());
// 更新上游已有规则:多次更新合并为一条;删除后 update 日志被顶替。
store.record_update("42").unwrap();
store.record_update("42").unwrap();
assert_eq!(
journal(&store),
vec![JournalEntry {
op: JournalOp::Update,
id: "42".into()
}]
);
store.record_remove("42").unwrap();
assert_eq!(
journal(&store),
vec![JournalEntry {
op: JournalOp::Remove,
id: "42".into()
}]
);
}
#[test]
fn promote_renames_rule_and_journal_ids() {
let root = tempfile::tempdir().unwrap();
let store = RuleStore::open(root.path().join("rules")).unwrap();
store
.upsert(&record("local-a", "content", "2026-01-01T00:00:00.000Z"))
.unwrap();
store.record_add("local-a").unwrap();
store.promote("local-a", "17353272").unwrap();
assert!(store.get("local-a").unwrap().is_none());
let promoted = store.get("17353272").unwrap().unwrap();
assert_eq!(promoted.knowledge, "content");
assert_eq!(promoted.title, "title-local-a");
assert_eq!(journal(&store)[0].id, "17353272");
}
#[test]
fn replace_all_mirrors_upstream_state() {
let root = tempfile::tempdir().unwrap();
let store = RuleStore::open(root.path().join("rules")).unwrap();
store
.upsert(&record("stale", "gone soon", "2026-01-01T00:00:00.000Z"))
.unwrap();
store
.replace_all(&[record(
"17353272",
"from upstream",
"2026-02-01T00:00:00.000Z",
)])
.unwrap();
let listed = store.list().unwrap();
assert_eq!(listed.len(), 1);
assert_eq!(listed[0].id, "17353272");
assert_eq!(listed[0].knowledge, "from upstream");
assert!(store.get("stale").unwrap().is_none());
}
#[test]
fn lists_hand_written_markdown_without_metadata() {
let root = tempfile::tempdir().unwrap();
let store = RuleStore::open(root.path().join("rules")).unwrap();
std::fs::write(root.path().join("rules/manual_rule.md"), "hand written").unwrap();
let listed = store.list().unwrap();
assert_eq!(listed.len(), 1);
assert_eq!(listed[0].id, "manual_rule");
assert_eq!(listed[0].title, "manual_rule");
assert_eq!(listed[0].knowledge, "hand written");
}
}
@@ -0,0 +1,184 @@
//! Replays the offline journal to upstream and mirrors upstream list state.
use axum::{
body::{Body, Bytes},
http::{header, HeaderMap, HeaderValue, Method, Request},
};
use prost::Message;
use crate::{api::cursor::proxy, cursor::protocol::connect, Result};
use super::{
store::{JournalOp, RuleRecord, RuleStore},
KnowledgeBaseAddRequest, KnowledgeBaseAddResponse, KnowledgeBaseListItem,
KnowledgeBaseRemoveRequest, KnowledgeBaseRemoveResponse, KnowledgeBaseUpdateRequest,
KnowledgeBaseUpdateResponse,
};
const ADD_PATH: &str = "/aiserver.v1.AiService/KnowledgeBaseAdd";
const UPDATE_PATH: &str = "/aiserver.v1.AiService/KnowledgeBaseUpdate";
const REMOVE_PATH: &str = "/aiserver.v1.AiService/KnowledgeBaseRemove";
/// 逐条把离线日志推送到上游。返回 true 表示日志已清空(上游可用),
/// false 表示上游不可达,剩余日志保留、调用方应降级到本地。
pub async fn replay(
upstream: &proxy::CursorProxy,
headers: &HeaderMap,
store: &RuleStore,
) -> Result<bool> {
while let Some(entry) = store.journal_front()? {
let advanced = match entry.op {
JournalOp::Add => replay_add(upstream, headers, store, &entry.id).await?,
JournalOp::Update => replay_update(upstream, headers, store, &entry.id).await?,
JournalOp::Remove => replay_remove(upstream, headers, store, &entry.id).await?,
};
if !advanced {
return Ok(false);
}
}
Ok(true)
}
/// 用上游返回的完整列表覆盖本地镜像。仅应在日志已清空时调用。
pub fn mirror(store: &RuleStore, items: Vec<KnowledgeBaseListItem>) -> Result<()> {
let records = items
.into_iter()
.map(|item| RuleRecord {
id: item.id,
knowledge: item.knowledge,
title: item.title,
created_at: item.created_at,
is_generated: item.is_generated,
git_origin: String::new(),
})
.collect::<Vec<_>>();
store.replace_all(&records)
}
async fn replay_add(
upstream: &proxy::CursorProxy,
headers: &HeaderMap,
store: &RuleStore,
id: &str,
) -> Result<bool> {
let Some(record) = store.get(id)? else {
// 规则文件已不在(被手动删除等),日志作废。
store.pop_journal()?;
return Ok(true);
};
let message = KnowledgeBaseAddRequest {
knowledge: record.knowledge,
title: record.title,
git_origin: record.git_origin,
composer_id: None,
};
let Some(body) = send(upstream, headers, ADD_PATH, &message).await else {
return Ok(false);
};
let Ok(reply) = connect::decode_unary::<KnowledgeBaseAddResponse>(&body) else {
return Ok(false);
};
if !reply.success || reply.id.is_empty() {
tracing::warn!(
id,
"rules upstream declined replayed add; dropping journal entry"
);
store.pop_journal()?;
return Ok(true);
}
store.promote(id, &reply.id)?;
store.pop_journal()?;
tracing::info!(
local_id = id,
upstream_id = reply.id,
"replayed offline rule add to upstream"
);
Ok(true)
}
async fn replay_update(
upstream: &proxy::CursorProxy,
headers: &HeaderMap,
store: &RuleStore,
id: &str,
) -> Result<bool> {
let Some(record) = store.get(id)? else {
store.pop_journal()?;
return Ok(true);
};
let message = KnowledgeBaseUpdateRequest {
id: id.into(),
knowledge: record.knowledge,
title: record.title,
};
let Some(body) = send(upstream, headers, UPDATE_PATH, &message).await else {
return Ok(false);
};
let Ok(reply) = connect::decode_unary::<KnowledgeBaseUpdateResponse>(&body) else {
return Ok(false);
};
if !reply.success {
tracing::warn!(
id,
"rules upstream declined replayed update; dropping journal entry"
);
}
store.pop_journal()?;
Ok(true)
}
async fn replay_remove(
upstream: &proxy::CursorProxy,
headers: &HeaderMap,
store: &RuleStore,
id: &str,
) -> Result<bool> {
let message = KnowledgeBaseRemoveRequest { id: id.into() };
let Some(body) = send(upstream, headers, REMOVE_PATH, &message).await else {
return Ok(false);
};
let Ok(reply) = connect::decode_unary::<KnowledgeBaseRemoveResponse>(&body) else {
return Ok(false);
};
if !reply.success {
tracing::warn!(
id,
"rules upstream declined replayed remove; dropping journal entry"
);
}
store.pop_journal()?;
Ok(true)
}
/// 以当前请求的头为模板向上游发起一次 unary RPC。
/// 成功(2xx)返回响应体;不可达或被拒绝返回 None,由调用方保留日志。
async fn send(
upstream: &proxy::CursorProxy,
template: &HeaderMap,
path: &str,
message: &impl Message,
) -> Option<Bytes> {
let mut headers = template.clone();
// 模板里的上游 URL 头指向原始 RPC 路径,必须移除才能命中回放路径。
headers.remove(proxy::UPSTREAM_URL_HEADER);
headers.remove(header::CONTENT_LENGTH);
headers.insert(
header::CONTENT_TYPE,
HeaderValue::from_static("application/proto"),
);
let mut request = Request::new(Body::from(message.encode_to_vec()));
*request.method_mut() = Method::POST;
*request.uri_mut() = path.parse().expect("replay path is a valid URI");
*request.headers_mut() = headers;
match proxy::forward_buffered(upstream, request).await {
Ok(response) if response.status.is_success() => Some(response.body),
Ok(response) => {
tracing::warn!(path, status = %response.status, "rules journal replay rejected by upstream");
None
}
Err(error) => {
tracing::warn!(path, %error, "rules journal replay cannot reach upstream");
None
}
}
}
+1
View File
@@ -4,6 +4,7 @@ pub mod account;
pub mod analytics;
pub mod blob_sync;
pub mod context_sync;
pub mod knowledge;
pub mod model_catalog;
pub mod observability;
pub mod tab;
+15 -6
View File
@@ -10,7 +10,7 @@ use prost::Message;
use crate::{
api::cursor::proxy::{self, CursorProxy},
cursor::{protocol::proto::agent::v1 as agent, transport::TransportRegistry},
model::{format_token_count, parse_token_count, ModelConfig, ModelType},
model::{format_token_count, parse_token_count, ModelConfig},
plugin::PluginModelDescriptor,
Error, Result,
};
@@ -378,16 +378,25 @@ fn available_model(model: &ModelConfig) -> AvailableModel {
display_name: "Cursor".into(),
}),
model_picker_badges: vec![ModelPickerBadge {
label: match model.model_type {
ModelType::OpenAi => "OpenAI".into(),
ModelType::Anthropic => "Anthropic".into(),
},
label: model
.group_name
.clone()
.unwrap_or_else(|| provider_host(&model.base_url)),
variant: 1,
dismiss_on_selection: false,
}],
}
}
/// 徽章回退标签:base_url 的主机名。入库时已校验为带主机的 HTTP(S) URL,
/// 解析失败仅是理论分支,此时原样返回 base_url。
fn provider_host(base_url: &str) -> String {
reqwest::Url::parse(base_url.trim())
.ok()
.and_then(|url| url.host_str().map(str::to_lowercase))
.unwrap_or_else(|| base_url.trim().into())
}
fn model_parameters(
contexts: &[(String, String)],
thinking: bool,
@@ -611,7 +620,7 @@ fn available_plugin_model(model: &PluginModelDescriptor) -> AvailableModel {
display_name: model.provider_type.clone(),
}),
model_picker_badges: vec![ModelPickerBadge {
label: model.provider_type.clone(),
label: model.plugin_name.clone(),
variant: 1,
dismiss_on_selection: false,
}],
+29 -2
View File
@@ -50,7 +50,24 @@ impl TransportRegistry {
compiler: PromptCompiler,
web_cache: WebCache,
) -> Self {
Self::build(store, provider, compiler, web_cache, None)
Self::build(store, provider, compiler, web_cache, None, None)
}
/// 附带本地 rules 目录的构造;编译请求上下文时会合并该目录下的 md 规则。
pub fn with_local_rules(
store: Store,
provider: Arc<dyn Provider>,
compiler: PromptCompiler,
local_rules_dir: std::path::PathBuf,
) -> Self {
Self::build(
store,
provider,
compiler,
WebCache::default(),
None,
Some(local_rules_dir),
)
}
pub fn with_plugins(
@@ -59,8 +76,16 @@ impl TransportRegistry {
compiler: PromptCompiler,
web_cache: WebCache,
plugins: PluginRegistry,
local_rules_dir: std::path::PathBuf,
) -> Self {
Self::build(store, provider, compiler, web_cache, Some(plugins))
Self::build(
store,
provider,
compiler,
web_cache,
Some(plugins),
Some(local_rules_dir),
)
}
fn build(
@@ -69,6 +94,7 @@ impl TransportRegistry {
compiler: PromptCompiler,
web_cache: WebCache,
plugins: Option<PluginRegistry>,
local_rules_dir: Option<std::path::PathBuf>,
) -> Self {
Self {
inner: Arc::new(RegistryInner {
@@ -80,6 +106,7 @@ impl TransportRegistry {
provider,
compiler,
web_cache.clone(),
local_rules_dir,
),
store,
web_cache,
+4
View File
@@ -161,6 +161,10 @@ fn is_local_path(path: &str) -> bool {
| "/aiserver.v1.DashboardService/GetUserProfile"
| "/aiserver.v1.DashboardService/GetCurrentPeriodUsage"
| "/aiserver.v1.DashboardService/GetUsageLimitStatusAndActiveGrants"
| "/aiserver.v1.AiService/KnowledgeBaseAdd"
| "/aiserver.v1.AiService/KnowledgeBaseList"
| "/aiserver.v1.AiService/KnowledgeBaseUpdate"
| "/aiserver.v1.AiService/KnowledgeBaseRemove"
| "/aiserver.v1.AnalyticsService/BootstrapStatsig"
| "/auth/full_stripe_profile"
)
+11
View File
@@ -87,6 +87,9 @@ pub struct ModelConfigInput {
#[serde(default)]
pub sort_order: i64,
pub display_name: String,
/// 供应商分组的自定义显示名;同一 base_url 主机下的模型共享。
#[serde(default)]
pub group_name: Option<String>,
#[serde(rename = "type")]
pub model_type: ModelType,
pub base_url: String,
@@ -124,6 +127,7 @@ pub struct ModelConfig {
pub model_hash: String,
pub sort_order: i64,
pub display_name: String,
pub group_name: Option<String>,
#[serde(rename = "type")]
pub model_type: ModelType,
pub base_url: String,
@@ -204,6 +208,12 @@ impl ModelConfig {
pub fn normalize_model_input(input: &ModelConfigInput) -> Result<ModelConfigInput> {
let display_name = required(&input.display_name, "model display name")?;
let group_name = input
.group_name
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
.map(String::from);
let base_url = normalize_request_url(&input.base_url)?;
let api_key = required(&input.api_key, "model API key")?;
let tooltip_data = required(&input.tooltip_data, "model tooltip")?;
@@ -230,6 +240,7 @@ pub fn normalize_model_input(input: &ModelConfigInput) -> Result<ModelConfigInpu
let normalized = ModelConfigInput {
sort_order: input.sort_order.max(0),
display_name,
group_name,
model_type: input.model_type,
base_url,
use_full_url: input.use_full_url,
+1
View File
@@ -423,6 +423,7 @@ mod tests {
.create_model(&ModelConfigInput {
sort_order: 0,
display_name: "Test Model".into(),
group_name: None,
model_type: ModelType::OpenAi,
base_url: "https://example.com/v1/chat/completions".into(),
use_full_url: true,
+1
View File
@@ -158,6 +158,7 @@ fn model_input(model: LegacyModel) -> Result<ModelConfigInput> {
Ok(ModelConfigInput {
sort_order: model.sort,
display_name: model.display_name.clone(),
group_name: None,
model_type,
base_url,
use_full_url,
+1 -1
View File
@@ -446,7 +446,7 @@ mod tests {
.unwrap();
assert_eq!(checksum_after, checksum_before);
assert_eq!(versions, vec![1, 2, 3, 4, 5, 6, 7]);
assert_eq!(versions, vec![1, 2, 3, 4, 5, 6, 7, 8]);
assert_eq!(checkpoint_table_exists, 1);
}
}
+72 -4
View File
@@ -11,7 +11,7 @@ use crate::{
use super::{now_ms, Store};
const MODEL_COLUMNS: &str = r#"
model_hash, sort_order, display_name, model_type, base_url, use_full_url, api_key, tooltip_data,
model_hash, sort_order, display_name, group_name, model_type, base_url, use_full_url, api_key, tooltip_data,
model_id, reasoning_effort, openai_endpoint, openai_extra_params_enabled,
openai_extra_params_json, custom_headers_enabled, custom_headers_json,
anthropic_extra_params_enabled, anthropic_extra_params_json, context_window_tokens,
@@ -123,7 +123,7 @@ impl Store {
}
let result = sqlx::query(
r#"UPDATE model_configs SET
model_hash = ?, sort_order = ?, display_name = ?, model_type = ?, base_url = ?,
model_hash = ?, sort_order = ?, display_name = ?, group_name = ?, model_type = ?, base_url = ?,
use_full_url = ?, api_key = ?, tooltip_data = ?, model_id = ?, reasoning_effort = ?,
openai_endpoint = ?, openai_extra_params_enabled = ?, openai_extra_params_json = ?,
custom_headers_enabled = ?, custom_headers_json = ?,
@@ -135,6 +135,7 @@ impl Store {
.bind(&next_hash)
.bind(input.sort_order)
.bind(&input.display_name)
.bind(&input.group_name)
.bind(input.model_type.as_str())
.bind(&input.base_url)
.bind(input.use_full_url)
@@ -242,13 +243,13 @@ async fn insert_model_with_conflict(
) -> Result<bool> {
let mut statement = String::from(
r#"INSERT INTO model_configs(
model_hash, sort_order, display_name, model_type, base_url, use_full_url, api_key, tooltip_data,
model_hash, sort_order, display_name, group_name, model_type, base_url, use_full_url, api_key, tooltip_data,
model_id, reasoning_effort, openai_endpoint, openai_extra_params_enabled,
openai_extra_params_json, custom_headers_enabled, custom_headers_json,
anthropic_extra_params_enabled, anthropic_extra_params_json, context_window_tokens,
max_completion_tokens, anthropic_max_tokens, anthropic_thinking_effort,
thinking_budget_tokens, created_at_ms, updated_at_ms
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)"#,
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)"#,
);
if ignore_existing {
statement.push_str(" ON CONFLICT(model_hash) DO NOTHING");
@@ -257,6 +258,7 @@ async fn insert_model_with_conflict(
.bind(hash)
.bind(input.sort_order)
.bind(&input.display_name)
.bind(&input.group_name)
.bind(input.model_type.as_str())
.bind(&input.base_url)
.bind(input.use_full_url)
@@ -288,6 +290,7 @@ fn model_from_row(row: sqlx::sqlite::SqliteRow) -> Result<ModelConfig> {
model_hash: row.try_get("model_hash")?,
sort_order: row.try_get("sort_order")?,
display_name: row.try_get("display_name")?,
group_name: row.try_get("group_name")?,
model_type: ModelType::from_str(row.try_get("model_type")?)?,
base_url: row.try_get("base_url")?,
use_full_url: row.try_get("use_full_url")?,
@@ -320,6 +323,71 @@ fn model_from_row(row: sqlx::sqlite::SqliteRow) -> Result<ModelConfig> {
})
}
#[cfg(test)]
mod tests {
use super::*;
fn model_input(group_name: Option<&str>) -> ModelConfigInput {
ModelConfigInput {
sort_order: 0,
display_name: "Test Model".into(),
group_name: group_name.map(String::from),
model_type: ModelType::OpenAi,
base_url: "https://example.com/v1/chat/completions".into(),
use_full_url: true,
api_key: "test-key".into(),
tooltip_data: "Test Model".into(),
model_id: "test-model".into(),
reasoning_effort: None,
openai_endpoint: crate::model::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,
}
}
/// 分组名是纯展示字段:入库时去除首尾空白、空串归一为 NULL,
/// 更新分组名不得改变模型身份哈希。
#[tokio::test]
async fn group_name_round_trips_without_changing_model_identity() {
let directory = tempfile::tempdir().unwrap();
let store = Store::connect(&format!(
"sqlite://{}",
directory.path().join("test.db").display()
))
.await
.unwrap();
let created = store
.create_model(&model_input(Some(" My Group ")))
.await
.unwrap();
assert_eq!(created.group_name.as_deref(), Some("My Group"));
let renamed = store
.update_model(&created.model_hash, &model_input(Some("Renamed")))
.await
.unwrap();
assert_eq!(renamed.model_hash, created.model_hash);
assert_eq!(renamed.group_name.as_deref(), Some("Renamed"));
let cleared = store
.update_model(&created.model_hash, &model_input(Some(" ")))
.await
.unwrap();
assert_eq!(cleared.model_hash, created.model_hash);
assert_eq!(cleared.group_name, None);
}
}
fn optional_u64(row: &sqlx::sqlite::SqliteRow, column: &str) -> Result<Option<u64>> {
row.try_get::<Option<i64>, _>(column)?
.map(|value| {
+1
View File
@@ -27,6 +27,7 @@ async fn summarize_replaces_model_history_and_preserves_cursor_history() {
.create_model(&ModelConfigInput {
sort_order: 0,
display_name: "Test Model".into(),
group_name: None,
model_type: ModelType::OpenAi,
base_url: "https://example.com/v1/chat/completions".into(),
use_full_url: true,
+1
View File
@@ -828,6 +828,7 @@ async fn injected_user_context_interrupts_automatic_compaction() {
.create_model(&ModelConfigInput {
sort_order: 0,
display_name: "Test Model".into(),
group_name: None,
model_type: ModelType::OpenAi,
base_url: "https://example.com/v1/chat/completions".into(),
use_full_url: true,
+213
View File
@@ -0,0 +1,213 @@
//! Verifies KnowledgeBase rules CRUD falls back to local markdown storage
//! when the Cursor upstream is unreachable or rejects the request.
#[path = "support/fixtures.rs"]
mod fixtures;
use axum::{
body::{to_bytes, Body},
extract::Extension,
http::{header, Request, Response},
};
use cursor_server::{
api::cursor::proxy::CursorProxy,
cursor::services::knowledge::{self, KnowledgeService},
};
use prost::Message;
// 测试侧的镜像消息定义,同时充当 wire 兼容性检查。
#[derive(Clone, PartialEq, Message)]
struct AddRequest {
#[prost(string, tag = "1")]
knowledge: String,
#[prost(string, tag = "2")]
title: String,
#[prost(string, tag = "3")]
git_origin: String,
}
#[derive(Clone, PartialEq, Message)]
struct AddResponse {
#[prost(bool, tag = "1")]
success: bool,
#[prost(string, tag = "2")]
id: String,
}
#[derive(Clone, PartialEq, Message)]
struct ListRequest {
#[prost(int32, optional, tag = "1")]
limit: Option<i32>,
}
#[derive(Clone, PartialEq, Message)]
struct ListResponse {
#[prost(bool, tag = "1")]
success: bool,
#[prost(message, repeated, tag = "2")]
all_results: Vec<ListItem>,
}
#[derive(Clone, PartialEq, Message)]
struct ListItem {
#[prost(string, tag = "1")]
id: String,
#[prost(string, tag = "2")]
knowledge: String,
#[prost(string, tag = "3")]
title: String,
#[prost(string, tag = "4")]
created_at: String,
#[prost(bool, tag = "5")]
is_generated: bool,
}
#[derive(Clone, PartialEq, Message)]
struct UpdateRequest {
#[prost(string, tag = "1")]
id: String,
#[prost(string, tag = "2")]
knowledge: String,
#[prost(string, tag = "3")]
title: String,
}
#[derive(Clone, PartialEq, Message)]
struct UpdateResponse {
#[prost(bool, tag = "1")]
success: bool,
}
#[derive(Clone, PartialEq, Message)]
struct RemoveRequest {
#[prost(string, tag = "1")]
id: String,
}
#[derive(Clone, PartialEq, Message)]
struct RemoveResponse {
#[prost(bool, tag = "1")]
success: bool,
}
fn proto_request(message: &impl Message) -> Request<Body> {
Request::post("/test")
.header(header::CONTENT_TYPE, "application/proto")
.body(Body::from(message.encode_to_vec()))
.unwrap()
}
async fn decode<M: Message + Default>(response: Response<Body>) -> M {
let body = to_bytes(response.into_body(), usize::MAX).await.unwrap();
M::decode(body.as_ref()).unwrap()
}
/// 无凭据请求上游必然失败(网络错误或 401),四个接口全部走本地降级,
/// 覆盖 md 持久化、离线日志压缩与增删改查闭环。
#[tokio::test]
async fn offline_crud_round_trip_persists_markdown() {
let (_store_dir, store) = fixtures::temp_store().await;
let upstream = CursorProxy::cursor(store).unwrap();
let rules_dir = tempfile::tempdir().unwrap();
let rules_root = rules_dir.path().join("rules");
let service = KnowledgeService::with_root(rules_root.clone()).unwrap();
// Add:得到本地临时 id,md 文件落盘。
let response = knowledge::add(
Extension(upstream.clone()),
Extension(service.clone()),
proto_request(&AddRequest {
knowledge: "always answer in haiku".into(),
title: "haiku rule".into(),
git_origin: String::new(),
}),
)
.await
.unwrap();
let added: AddResponse = decode(response).await;
assert!(added.success);
assert!(added.id.starts_with("local-"), "offline add uses a local id");
let markdown = rules_root.join(format!("{}.md", added.id));
assert_eq!(
std::fs::read_to_string(&markdown).unwrap(),
"always answer in haiku"
);
// List:本地缓存返回刚写入的规则。
let response = knowledge::list(
Extension(upstream.clone()),
Extension(service.clone()),
proto_request(&ListRequest { limit: Some(100) }),
)
.await
.unwrap();
let listed: ListResponse = decode(response).await;
assert!(listed.success);
assert_eq!(listed.all_results.len(), 1);
assert_eq!(listed.all_results[0].id, added.id);
assert_eq!(listed.all_results[0].title, "haiku rule");
// Update:内容与标题都更新到 md 与元数据。
let response = knowledge::update(
Extension(upstream.clone()),
Extension(service.clone()),
proto_request(&UpdateRequest {
id: added.id.clone(),
knowledge: "always answer in sonnets".into(),
title: "sonnet rule".into(),
}),
)
.await
.unwrap();
let updated: UpdateResponse = decode(response).await;
assert!(updated.success);
assert_eq!(
std::fs::read_to_string(&markdown).unwrap(),
"always answer in sonnets"
);
// Remove:文件删除,列表为空。
let response = knowledge::remove(
Extension(upstream.clone()),
Extension(service.clone()),
proto_request(&RemoveRequest {
id: added.id.clone(),
}),
)
.await
.unwrap();
let removed: RemoveResponse = decode(response).await;
assert!(removed.success);
assert!(!markdown.exists());
let response = knowledge::list(
Extension(upstream),
Extension(service),
proto_request(&ListRequest { limit: Some(100) }),
)
.await
.unwrap();
let listed: ListResponse = decode(response).await;
assert!(listed.all_results.is_empty());
}
#[tokio::test]
async fn updating_missing_rule_reports_failure() {
let (_store_dir, store) = fixtures::temp_store().await;
let upstream = CursorProxy::cursor(store).unwrap();
let rules_dir = tempfile::tempdir().unwrap();
let service = KnowledgeService::with_root(rules_dir.path().join("rules")).unwrap();
let response = knowledge::update(
Extension(upstream),
Extension(service),
proto_request(&UpdateRequest {
id: "17353272".into(),
knowledge: "anything".into(),
title: "anything".into(),
}),
)
.await
.unwrap();
let updated: UpdateResponse = decode(response).await;
assert!(!updated.success);
}
+133
View File
@@ -0,0 +1,133 @@
//! Verifies local markdown rules are merged into the request-context message.
#[path = "support/fake_provider.rs"]
mod fake_provider;
#[path = "support/fixtures.rs"]
mod fixtures;
use std::sync::Arc;
use cursor_server::{
cursor::{
prompting::{PromptAssets, PromptCompiler},
protocol::connect,
protocol::proto::agent::v1 as pb,
TransportCommand, TransportRegistry,
},
model::{ContentPart, ProjectedContent},
provider::{FinishReason, ModelEvent},
};
#[tokio::test]
async fn local_markdown_rules_land_in_the_request_context_message() {
let (_store_dir, store) = fixtures::temp_store().await;
let provider = fake_provider::FakeProvider::default();
provider.push(vec![
ModelEvent::Start {
model_call_id: "call-1".into(),
},
ModelEvent::TextStart,
ModelEvent::TextDelta("ok".into()),
ModelEvent::TextEnd,
ModelEvent::Done(FinishReason::Stop),
]);
let assets = PromptAssets::load(
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
.join("prompt/cursor")
.as_path(),
)
.unwrap();
let rules_dir = tempfile::tempdir().unwrap();
let rules_root = rules_dir.path().join("rules");
std::fs::create_dir_all(&rules_root).unwrap();
std::fs::write(rules_root.join("17353272.md"), "Always answer in haiku.").unwrap();
let registry = TransportRegistry::with_local_rules(
store,
Arc::new(provider.clone()),
PromptCompiler::new(assets),
rules_root,
);
let handle = registry.get_or_create("rules-request").await.unwrap();
let mut output = handle.subscribe();
handle
.command(TransportCommand::Append {
seqno: 0,
message: Box::new(user_run()),
})
.await
.unwrap();
loop {
let frame = tokio::time::timeout(std::time::Duration::from_secs(5), output.recv())
.await
.expect("run finishes within timeout")
.expect("output stays open until EndStream");
let ended = connect::decode_frames(&frame)
.unwrap()
.iter()
.any(|(flags, _)| flags & connect::END_STREAM_FLAG != 0);
if ended {
break;
}
}
let requests = provider.requests();
assert_eq!(requests.len(), 1);
let context_texts = requests[0]
.history
.iter()
.filter(|message| message.message_id.starts_with("request-context:"))
.map(|message| {
let ProjectedContent::Parts(parts) = &message.content else {
panic!("request context message must be parts")
};
let [ContentPart::Text { text }] = parts.as_slice() else {
panic!("request context message must be one text part")
};
text.clone()
})
.collect::<Vec<_>>();
assert_eq!(
context_texts.len(),
1,
"exactly one request-context message is projected"
);
assert!(
context_texts[0].contains("<user_rule>\nAlways answer in haiku.\n</user_rule>"),
"local markdown rule must appear as a user rule: {}",
context_texts[0]
);
registry.shutdown().await;
}
fn user_run() -> pb::AgentClientMessage {
pb::AgentClientMessage {
message: Some(pb::agent_client_message::Message::RunRequest(
pb::AgentRunRequest {
action: Some(pb::ConversationAction {
action: Some(pb::conversation_action::Action::UserMessageAction(
pb::UserMessageAction {
user_message: Some(pb::UserMessage {
text: "hello".into(),
message_id: "rules-user".into(),
mode: pb::AgentMode::Agent as i32,
..Default::default()
}),
..Default::default()
},
)),
..Default::default()
}),
conversation_id: Some("rules-conversation".into()),
run_id: Some("rules-request".into()),
requested_model: Some(pb::RequestedModel {
model_id: "test-model".into(),
..Default::default()
}),
..Default::default()
},
)),
}
}