mirror of
https://wget.la/https://github.com/leookun/cursor-byok
synced 2026-10-04 02:40:50 +08:00
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:
@@ -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;
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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)?;
|
||||
|
||||
@@ -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('<', "<")
|
||||
.replace('>', ">")
|
||||
}
|
||||
|
||||
#[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());
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
};
|
||||
|
||||
@@ -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(×tamp(&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
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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;
|
||||
|
||||
@@ -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,
|
||||
}],
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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"
|
||||
)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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| {
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
@@ -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()
|
||||
},
|
||||
)),
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user