From e7a1cca4c6b9e4ea45e5e0173462f8d3c3a91e40 Mon Sep 17 00:00:00 2001 From: leookun Date: Sun, 30 Aug 2026 23:28:05 +0800 Subject: [PATCH] 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. --- apps/desktop/src/demo/api.ts | 1 + .../src/features/models/CursorModelCards.tsx | 43 +- .../src/features/models/CursorModelEditor.tsx | 1 + .../models/CursorSettings.module.scss | 28 +- .../features/models/CursorSettingsPage.tsx | 62 ++- apps/desktop/src/i18n/generated/catalog.json | 293 +++++++---- apps/desktop/src/i18n/locales/en-US.json | 5 + apps/desktop/src/i18n/locales/zh-CN.json | 5 + apps/desktop/src/shared/api.ts | 2 + .../migrations/0008_add_model_group_name.sql | 4 + server/src/api/cursor/handlers.rs | 30 +- server/src/app.rs | 1 + server/src/cursor/compile/context.rs | 76 +++ server/src/cursor/compile/run.rs | 8 +- server/src/cursor/conversation/registry.rs | 4 + server/src/cursor/conversation/runtime.rs | 1 + server/src/cursor/services/knowledge/mod.rs | 356 +++++++++++++ server/src/cursor/services/knowledge/store.rs | 487 ++++++++++++++++++ server/src/cursor/services/knowledge/sync.rs | 184 +++++++ server/src/cursor/services/mod.rs | 1 + server/src/cursor/services/model_catalog.rs | 21 +- server/src/cursor/transport/registry.rs | 31 +- server/src/local_app/proxy.rs | 4 + server/src/model/configuration.rs | 11 + server/src/provider/recorder.rs | 1 + server/src/store/legacy_config.rs | 1 + server/src/store/migrations.rs | 2 +- server/src/store/models.rs | 76 ++- server/tests/compaction.rs | 1 + server/tests/interrupt.rs | 1 + server/tests/knowledge_rules.rs | 213 ++++++++ server/tests/local_rules_context.rs | 133 +++++ 32 files changed, 1950 insertions(+), 137 deletions(-) create mode 100644 server/migrations/0008_add_model_group_name.sql create mode 100644 server/src/cursor/services/knowledge/mod.rs create mode 100644 server/src/cursor/services/knowledge/store.rs create mode 100644 server/src/cursor/services/knowledge/sync.rs create mode 100644 server/tests/knowledge_rules.rs create mode 100644 server/tests/local_rules_context.rs diff --git a/apps/desktop/src/demo/api.ts b/apps/desktop/src/demo/api.ts index 18c6e2b..5577482 100644 --- a/apps/desktop/src/demo/api.ts +++ b/apps/desktop/src/demo/api.ts @@ -185,6 +185,7 @@ function createModel({ hash, order, name, type, url, modelId, endpoint = "/v1/re model_hash: hash, sort_order: order, display_name: name, + group_name: null, type, base_url: url, use_full_url: false, diff --git a/apps/desktop/src/features/models/CursorModelCards.tsx b/apps/desktop/src/features/models/CursorModelCards.tsx index b1de27a..3049b77 100644 --- a/apps/desktop/src/features/models/CursorModelCards.tsx +++ b/apps/desktop/src/features/models/CursorModelCards.tsx @@ -32,6 +32,7 @@ type CursorModelCardsProps = { onTestPluginModel: (model: PluginModelDescriptor) => void; onPluginSettings: (model: PluginModelDescriptor) => void; onReorder: (modelHashes: string[]) => void; + onGroupSettings: (group: CursorModelGroup) => void; }; type ModelGridProps = Omit & { @@ -60,6 +61,7 @@ export function CursorModelCards(props: CursorModelCardsProps) { key={group.key} label={group.label} icon={group.icon} + onSettings={props.grouping === "provider" ? () => props.onGroupSettings(group) : undefined} > {group.models.map((model) => void; children: ReactNode; }) { const [open, setOpen] = useState(true); return - +
+ + {onSettings && } + +
{open &&
{children}
}
; } @@ -273,8 +287,9 @@ function ModelGrid({ } function providerGroup(model: Model) { - const label = providerDomain(model.base_url); - return { key: label, label, icon: flatColorOrganizationIcon }; + const key = providerDomain(model.base_url); + const label = model.group_name?.trim() || key; + return { key, label, icon: flatColorOrganizationIcon }; } function providerDomain(baseUrl: string) { diff --git a/apps/desktop/src/features/models/CursorModelEditor.tsx b/apps/desktop/src/features/models/CursorModelEditor.tsx index ee0648d..d01d1d9 100644 --- a/apps/desktop/src/features/models/CursorModelEditor.tsx +++ b/apps/desktop/src/features/models/CursorModelEditor.tsx @@ -24,6 +24,7 @@ export const emptyCursorModelDraft = (): CursorModelDraft => ({ model: { sort_order: 0, display_name: "", + group_name: null, type: "openai", base_url: "", use_full_url: false, diff --git a/apps/desktop/src/features/models/CursorSettings.module.scss b/apps/desktop/src/features/models/CursorSettings.module.scss index a96a8d3..0624aad 100644 --- a/apps/desktop/src/features/models/CursorSettings.module.scss +++ b/apps/desktop/src/features/models/CursorSettings.module.scss @@ -73,8 +73,34 @@ .groupCard { padding: 4px 12px 8px; } +.groupHeader { + display: flex; + align-items: center; + gap: 8px; +} +.groupSettings { + padding: 4px 2px; + color: var(--vscode-descriptionForeground); + background: none; + border: none; + font-size: type.$font-size-xs; + white-space: nowrap; + cursor: pointer; + + &:hover { color: var(--vscode-foreground); } +} +.groupChevron { + display: flex; + align-items: center; + padding: 4px 0; + color: var(--vscode-foreground); + background: none; + border: none; + cursor: pointer; +} .groupToggle { - width: 100%; + flex: 1; + min-width: 0; display: flex; align-items: center; gap: 8px; diff --git a/apps/desktop/src/features/models/CursorSettingsPage.tsx b/apps/desktop/src/features/models/CursorSettingsPage.tsx index 33deb24..5497ee6 100644 --- a/apps/desktop/src/features/models/CursorSettingsPage.tsx +++ b/apps/desktop/src/features/models/CursorSettingsPage.tsx @@ -2,13 +2,14 @@ import { useCallback, useEffect, useRef, useState } from "react"; import { useNavigate } from "react-router-dom"; import { api, configuredPluginModels, type Model, type ModelInput } from "../../shared/api"; import { CursorCaGate, CursorCaProvider, CursorModelGate, CursorModelProvider } from "./CursorGates"; -import { CursorModelCards, cursorModelGroups, type CursorModelGrouping } from "./CursorModelCards"; +import { CursorModelCards, cursorModelGroups, type CursorModelGroup, type CursorModelGrouping } from "./CursorModelCards"; import { CursorModelEditor, emptyCursorModelDraft, type CursorModelDraft } from "./CursorModelEditor"; import { CursorModelTestResult, type CursorModelTestState } from "./CursorModelTestResult"; import styles from "./CursorSettings.module.scss"; import { PageContent } from "../../shell/layout/PageContent"; import { LegacyModelImport } from "./LegacyModelImport"; import { ConfirmDialog } from "../../shared/ui/ConfirmDialog"; +import { FormField, SecretTextInput, TextInput } from "../../shared/ui/FormControls"; import controls from "../../shared/ui/Controls.module.scss"; import { Icon } from "../../shared/ui/Icon"; import { Modal } from "../../shared/ui/Modal"; @@ -34,6 +35,11 @@ export function CursorSettingsPage() { const [savingAndTesting, setSavingAndTesting] = useState(false); const [batchTesting, setBatchTesting] = useState(false); const [grouping, setGrouping] = useState("flat"); + const [settingsGroup, setSettingsGroup] = useState(null); + const [groupNameDraft, setGroupNameDraft] = useState(""); + const [groupBaseUrlDraft, setGroupBaseUrlDraft] = useState(""); + const [groupApiKeyDraft, setGroupApiKeyDraft] = useState(""); + const [groupSettingsBusy, setGroupSettingsBusy] = useState(false); const activeModelTests = useRef(new Map()); const caReady = cursorHarness?.ca === "ready"; const pluginModels = configuredPluginModels(plugins); @@ -208,6 +214,39 @@ export function CursorSettingsPage() { }]); if (created) message(t("模型已复制")); }; + const openGroupSettings = (group: CursorModelGroup) => { + setGroupNameDraft(group.models.find((model) => model.group_name?.trim())?.group_name?.trim() ?? ""); + setGroupBaseUrlDraft(sharedValue(group.models.map((model) => model.base_url)) ?? ""); + setGroupApiKeyDraft(sharedValue(group.models.map((model) => model.api_key)) ?? ""); + setSettingsGroup(group); + }; + const saveGroupSettings = async () => { + if (!settingsGroup) return; + const group_name = groupNameDraft.trim() || null; + const base_url = groupBaseUrlDraft.trim(); + const api_key = groupApiKeyDraft.trim(); + setGroupSettingsBusy(true); + try { + for (const model of settingsGroup.models) { + const input: ModelInput = { + ...modelInput(model), + group_name, + ...(base_url ? { base_url } : {}), + ...(api_key ? { api_key } : {}), + }; + if (input.group_name === (model.group_name ?? null) + && input.base_url === model.base_url + && input.api_key === model.api_key) continue; + await api.updateModel(model.model_hash, input); + } + await appStore.refresh(); + setSettingsGroup(null); + } catch (cause) { + message(errorText(cause)); + } finally { + setGroupSettingsBusy(false); + } + }; const reorderModels = useCallback(async (modelHashes: string[]) => { if (!await appStore.reorderCursorModels(modelHashes)) { message(appStore.getSnapshot().error || t("排序失败")); @@ -228,6 +267,7 @@ export function CursorSettingsPage() { onTestPluginModel={(model) => void testModel({ model_hash: model.id, display_name: model.displayName })} onPluginSettings={() => navigate("/plugins")} onReorder={reorderModels} + onGroupSettings={openGroupSettings} />; const refreshCa = async () => { @@ -274,6 +314,19 @@ export function CursorSettingsPage() { setCaCommand(null)} onConfirm={openCaTerminal}>
{t("需要授权安装证书")}{t("安装命令已自动复制。点击“打开终端”,将命令粘贴到终端中执行,并按提示输入密码。")}
{caCommand}
+ setSettingsGroup(null)} onSubmit={() => void saveGroupSettings()} submitLabel={t("保存")}> + {settingsGroup &&
+ + setGroupNameDraft(event.target.value)} /> + + + setGroupBaseUrlDraft(event.target.value)} /> + + + setGroupApiKeyDraft(event.target.value)} /> + +
} +
setDeleting(null)} onConfirm={() => { if (deleting) void appStore.deleteModel(deleting.model_hash); setDeleting(null); }}>

{t("确定删除这个模型吗?")}

; } @@ -283,6 +336,13 @@ function modelInput(model: Model): ModelInput { return input; } +/** 组内所有模型取值一致时返回该值,否则返回 null(表单留空表示保持不变)。 */ +function sharedValue(values: string[]): string | null { + const [first, ...rest] = values; + if (first === undefined) return null; + return rest.every((value) => value === first) ? first : null; +} + function draftInput(draft: CursorModelDraft): ModelInput { const model = { ...draft.model, diff --git a/apps/desktop/src/i18n/generated/catalog.json b/apps/desktop/src/i18n/generated/catalog.json index c3f45b4..5f06481 100644 --- a/apps/desktop/src/i18n/generated/catalog.json +++ b/apps/desktop/src/i18n/generated/catalog.json @@ -33,7 +33,7 @@ "refs": [ { "file": "features/models/CursorModelEditor.tsx", - "line": 151, + "line": 152, "column": 40 } ] @@ -106,12 +106,12 @@ "refs": [ { "file": "features/models/CursorModelCards.tsx", - "line": 151, + "line": 165, "column": 64 }, { "file": "features/models/CursorModelCards.tsx", - "line": 265, + "line": 279, "column": 70 }, { @@ -141,7 +141,7 @@ "refs": [ { "file": "features/models/CursorSettingsPage.tsx", - "line": 142, + "line": 148, "column": 27 } ] @@ -205,12 +205,12 @@ "refs": [ { "file": "features/models/CursorModelEditor.tsx", - "line": 181, + "line": 182, "column": 16 }, { "file": "features/models/CursorSettingsPage.tsx", - "line": 296, + "line": 356, "column": 73 } ] @@ -249,7 +249,7 @@ "refs": [ { "file": "features/models/CursorSettingsPage.tsx", - "line": 307, + "line": 367, "column": 89 } ] @@ -447,12 +447,12 @@ "refs": [ { "file": "features/models/CursorModelEditor.tsx", - "line": 168, + "line": 169, "column": 16 }, { "file": "features/models/CursorSettingsPage.tsx", - "line": 306, + "line": 366, "column": 36 } ] @@ -476,7 +476,7 @@ "refs": [ { "file": "features/models/CursorSettingsPage.tsx", - "line": 264, + "line": 304, "column": 211 } ] @@ -512,7 +512,7 @@ "refs": [ { "file": "features/models/CursorModelEditor.tsx", - "line": 153, + "line": 154, "column": 286 } ] @@ -524,12 +524,12 @@ "refs": [ { "file": "features/models/CursorModelEditor.tsx", - "line": 142, + "line": 143, "column": 59 }, { "file": "features/models/CursorModelEditor.tsx", - "line": 142, + "line": 143, "column": 124 } ] @@ -546,17 +546,17 @@ }, { "file": "features/models/CursorSettingsPage.tsx", - "line": 267, + "line": 307, "column": 51 }, { "file": "features/models/CursorSettingsPage.tsx", - "line": 267, + "line": 307, "column": 130 }, { "file": "features/models/CursorSettingsPage.tsx", - "line": 269, + "line": 309, "column": 74 } ] @@ -679,7 +679,7 @@ "refs": [ { "file": "features/models/CursorSettingsPage.tsx", - "line": 277, + "line": 330, "column": 52 } ] @@ -778,17 +778,17 @@ "refs": [ { "file": "features/models/CursorModelEditor.tsx", - "line": 157, + "line": 158, "column": 49 }, { "file": "features/models/CursorModelEditor.tsx", - "line": 159, + "line": 160, "column": 50 }, { "file": "features/models/CursorModelEditor.tsx", - "line": 162, + "line": 163, "column": 50 } ] @@ -802,7 +802,7 @@ "refs": [ { "file": "features/models/CursorSettingsPage.tsx", - "line": 300, + "line": 360, "column": 89 } ] @@ -838,7 +838,7 @@ "refs": [ { "file": "features/models/CursorSettingsPage.tsx", - "line": 262, + "line": 302, "column": 134 } ] @@ -902,7 +902,7 @@ "refs": [ { "file": "features/models/CursorModelEditor.tsx", - "line": 153, + "line": 154, "column": 42 } ] @@ -952,12 +952,12 @@ "refs": [ { "file": "features/models/CursorModelEditor.tsx", - "line": 164, + "line": 165, "column": 27 }, { "file": "features/models/CursorSettingsPage.tsx", - "line": 299, + "line": 359, "column": 186 } ] @@ -986,7 +986,7 @@ }, { "file": "features/models/CursorSettingsPage.tsx", - "line": 277, + "line": 330, "column": 76 }, { @@ -1079,7 +1079,7 @@ "refs": [ { "file": "features/models/CursorModelEditor.tsx", - "line": 154, + "line": 155, "column": 42 } ] @@ -1093,7 +1093,7 @@ "refs": [ { "file": "features/models/CursorSettingsPage.tsx", - "line": 313, + "line": 373, "column": 70 } ] @@ -1167,17 +1167,17 @@ "refs": [ { "file": "features/models/CursorModelCards.tsx", - "line": 153, + "line": 167, "column": 96 }, { "file": "features/models/CursorModelCards.tsx", - "line": 267, + "line": 281, "column": 102 }, { "file": "features/models/CursorSettingsPage.tsx", - "line": 277, + "line": 330, "column": 99 }, { @@ -1216,6 +1216,23 @@ } ] }, + "3260348163d03b8e": { + "source": "留空保持不变", + "kind": "text", + "placeholders": [], + "refs": [ + { + "file": "features/models/CursorSettingsPage.tsx", + "line": 323, + "column": 35 + }, + { + "file": "features/models/CursorSettingsPage.tsx", + "line": 326, + "column": 41 + } + ] + }, "32896fdaaaa4c106": { "source": "账号已保存,模型目录已同步。", "kind": "text", @@ -1269,7 +1286,7 @@ "refs": [ { "file": "features/models/CursorSettingsPage.tsx", - "line": 298, + "line": 358, "column": 123 } ] @@ -1505,7 +1522,7 @@ "refs": [ { "file": "features/models/CursorSettingsPage.tsx", - "line": 264, + "line": 304, "column": 197 } ] @@ -1517,7 +1534,7 @@ "refs": [ { "file": "features/models/CursorSettingsPage.tsx", - "line": 274, + "line": 314, "column": 80 }, { @@ -1630,7 +1647,7 @@ "refs": [ { "file": "features/models/CursorSettingsPage.tsx", - "line": 151, + "line": 157, "column": 27 } ] @@ -1754,12 +1771,12 @@ "refs": [ { "file": "features/models/CursorModelEditor.tsx", - "line": 157, + "line": 158, "column": 25 }, { "file": "features/models/CursorSettingsPage.tsx", - "line": 299, + "line": 359, "column": 34 } ] @@ -1771,22 +1788,22 @@ "refs": [ { "file": "features/models/CursorModelCards.tsx", - "line": 150, + "line": 164, "column": 86 }, { "file": "features/models/CursorModelCards.tsx", - "line": 173, + "line": 187, "column": 86 }, { "file": "features/models/CursorModelCards.tsx", - "line": 264, + "line": 278, "column": 92 }, { "file": "features/models/CursorSettingsPage.tsx", - "line": 269, + "line": 309, "column": 703 } ] @@ -1834,7 +1851,7 @@ "refs": [ { "file": "features/models/CursorModelEditor.tsx", - "line": 148, + "line": 149, "column": 145 } ] @@ -1875,7 +1892,7 @@ "refs": [ { "file": "features/models/CursorModelEditor.tsx", - "line": 206, + "line": 207, "column": 41 } ] @@ -1933,7 +1950,7 @@ "refs": [ { "file": "features/models/CursorSettingsPage.tsx", - "line": 188, + "line": 194, "column": 13 } ] @@ -1993,7 +2010,7 @@ "refs": [ { "file": "features/models/CursorSettingsPage.tsx", - "line": 213, + "line": 252, "column": 47 } ] @@ -2017,7 +2034,7 @@ "refs": [ { "file": "features/models/CursorSettingsPage.tsx", - "line": 267, + "line": 307, "column": 63 } ] @@ -2029,17 +2046,17 @@ "refs": [ { "file": "features/models/CursorModelEditor.tsx", - "line": 159, + "line": 160, "column": 27 }, { "file": "features/models/CursorModelEditor.tsx", - "line": 162, + "line": 163, "column": 27 }, { "file": "features/models/CursorSettingsPage.tsx", - "line": 299, + "line": 359, "column": 83 } ] @@ -2063,7 +2080,7 @@ "refs": [ { "file": "features/models/CursorModelEditor.tsx", - "line": 155, + "line": 156, "column": 69 } ] @@ -2140,7 +2157,7 @@ "refs": [ { "file": "features/models/CursorSettingsPage.tsx", - "line": 263, + "line": 303, "column": 122 } ] @@ -2314,12 +2331,12 @@ "refs": [ { "file": "features/models/CursorModelCards.tsx", - "line": 152, + "line": 166, "column": 64 }, { "file": "features/models/CursorModelCards.tsx", - "line": 266, + "line": 280, "column": 70 }, { @@ -2336,7 +2353,7 @@ "refs": [ { "file": "shared/api.ts", - "line": 486, + "line": 488, "column": 43 } ] @@ -2499,17 +2516,17 @@ "refs": [ { "file": "features/models/CursorModelCards.tsx", - "line": 150, + "line": 164, "column": 98 }, { "file": "features/models/CursorModelCards.tsx", - "line": 173, + "line": 187, "column": 98 }, { "file": "features/models/CursorModelCards.tsx", - "line": 264, + "line": 278, "column": 104 } ] @@ -2586,7 +2603,7 @@ "refs": [ { "file": "features/models/CursorSettingsPage.tsx", - "line": 261, + "line": 301, "column": 103 } ] @@ -2663,7 +2680,7 @@ "refs": [ { "file": "shared/api.ts", - "line": 481, + "line": 483, "column": 43 } ] @@ -2711,12 +2728,12 @@ "refs": [ { "file": "features/models/CursorModelEditor.tsx", - "line": 175, + "line": 176, "column": 16 }, { "file": "features/models/CursorSettingsPage.tsx", - "line": 294, + "line": 354, "column": 67 } ] @@ -2817,7 +2834,7 @@ "refs": [ { "file": "shared/api.ts", - "line": 418, + "line": 420, "column": 21 } ] @@ -2844,7 +2861,7 @@ "refs": [ { "file": "features/models/CursorSettingsPage.tsx", - "line": 186, + "line": 192, "column": 11 } ] @@ -2871,7 +2888,7 @@ "refs": [ { "file": "features/models/CursorModelEditor.tsx", - "line": 153, + "line": 154, "column": 298 } ] @@ -2987,7 +3004,7 @@ "refs": [ { "file": "features/models/CursorModelEditor.tsx", - "line": 148, + "line": 149, "column": 115 } ] @@ -3086,7 +3103,7 @@ "refs": [ { "file": "features/models/CursorModelEditor.tsx", - "line": 142, + "line": 143, "column": 76 } ] @@ -3124,7 +3141,7 @@ "refs": [ { "file": "features/models/CursorSettingsPage.tsx", - "line": 314, + "line": 374, "column": 87 } ] @@ -3208,7 +3225,7 @@ "refs": [ { "file": "features/models/CursorSettingsPage.tsx", - "line": 277, + "line": 330, "column": 250 } ] @@ -3314,12 +3331,12 @@ }, { "file": "features/models/CursorModelEditor.tsx", - "line": 163, + "line": 164, "column": 27 }, { "file": "features/models/CursorModelEditor.tsx", - "line": 163, + "line": 164, "column": 58 } ] @@ -3479,7 +3496,7 @@ "refs": [ { "file": "features/models/CursorSettingsPage.tsx", - "line": 197, + "line": 203, "column": 22 } ] @@ -3532,9 +3549,14 @@ "refs": [ { "file": "features/models/CursorSettingsPage.tsx", - "line": 269, + "line": 309, "column": 432 }, + { + "file": "features/models/CursorSettingsPage.tsx", + "line": 317, + "column": 193 + }, { "file": "features/settings/ProxySettingsCard.tsx", "line": 33, @@ -3616,7 +3638,7 @@ "refs": [ { "file": "features/models/CursorSettingsPage.tsx", - "line": 125, + "line": 131, "column": 15 } ] @@ -3671,7 +3693,7 @@ "refs": [ { "file": "features/models/CursorSettingsPage.tsx", - "line": 269, + "line": 309, "column": 715 } ] @@ -3705,7 +3727,7 @@ "refs": [ { "file": "features/models/CursorModelEditor.tsx", - "line": 149, + "line": 150, "column": 61 } ] @@ -3717,12 +3739,12 @@ "refs": [ { "file": "features/models/CursorModelCards.tsx", - "line": 246, + "line": 260, "column": 113 }, { "file": "features/models/CursorModelCards.tsx", - "line": 246, + "line": 260, "column": 131 } ] @@ -3744,7 +3766,7 @@ }, { "file": "features/models/CursorModelEditor.tsx", - "line": 154, + "line": 155, "column": 25 } ] @@ -3780,7 +3802,7 @@ "refs": [ { "file": "features/models/CursorModelEditor.tsx", - "line": 155, + "line": 156, "column": 118 } ] @@ -3915,6 +3937,23 @@ } ] }, + "b2617bf9ae663752": { + "source": "分组设置", + "kind": "text", + "placeholders": [], + "refs": [ + { + "file": "features/models/CursorModelCards.tsx", + "line": 132, + "column": 99 + }, + { + "file": "features/models/CursorSettingsPage.tsx", + "line": 317, + "column": 49 + } + ] + }, "b4411558b932266f": { "source": "上游类型", "kind": "text", @@ -4049,7 +4088,7 @@ "refs": [ { "file": "features/models/CursorSettingsPage.tsx", - "line": 274, + "line": 314, "column": 103 } ] @@ -4107,7 +4146,7 @@ "refs": [ { "file": "features/models/CursorModelEditor.tsx", - "line": 154, + "line": 155, "column": 99 } ] @@ -4134,7 +4173,7 @@ "refs": [ { "file": "features/models/CursorSettingsPage.tsx", - "line": 189, + "line": 195, "column": 13 } ] @@ -4192,6 +4231,23 @@ } ] }, + "be961dc60ab610da": { + "source": "修改后应用于该分组下的全部模型;留空保持各模型现有配置不变。", + "kind": "text", + "placeholders": [], + "refs": [ + { + "file": "features/models/CursorSettingsPage.tsx", + "line": 322, + "column": 45 + }, + { + "file": "features/models/CursorSettingsPage.tsx", + "line": 325, + "column": 42 + } + ] + }, "bf57afd709694b55": { "source": "概览时间范围", "kind": "text", @@ -4292,7 +4348,7 @@ "refs": [ { "file": "features/models/CursorSettingsPage.tsx", - "line": 269, + "line": 309, "column": 62 } ] @@ -4304,12 +4360,12 @@ "refs": [ { "file": "features/models/CursorModelEditor.tsx", - "line": 160, + "line": 161, "column": 27 }, { "file": "features/models/CursorModelEditor.tsx", - "line": 160, + "line": 161, "column": 58 } ] @@ -4386,7 +4442,7 @@ }, { "file": "features/models/CursorModelEditor.tsx", - "line": 153, + "line": 154, "column": 25 } ] @@ -4468,7 +4524,7 @@ "refs": [ { "file": "features/models/CursorModelEditor.tsx", - "line": 164, + "line": 165, "column": 50 } ] @@ -4533,8 +4589,13 @@ "refs": [ { "file": "features/models/CursorModelEditor.tsx", - "line": 148, + "line": 149, "column": 70 + }, + { + "file": "features/models/CursorSettingsPage.tsx", + "line": 322, + "column": 27 } ] }, @@ -4545,7 +4606,7 @@ "refs": [ { "file": "features/models/CursorSettingsPage.tsx", - "line": 269, + "line": 309, "column": 675 }, { @@ -4598,17 +4659,17 @@ "refs": [ { "file": "features/models/CursorModelEditor.tsx", - "line": 157, + "line": 158, "column": 121 }, { "file": "features/models/CursorModelEditor.tsx", - "line": 159, + "line": 160, "column": 122 }, { "file": "features/models/CursorModelEditor.tsx", - "line": 162, + "line": 163, "column": 122 } ] @@ -4620,7 +4681,7 @@ "refs": [ { "file": "features/models/CursorModelEditor.tsx", - "line": 164, + "line": 165, "column": 137 } ] @@ -4639,6 +4700,18 @@ } ] }, + "d7e266bdc8064193": { + "source": "分组名称", + "kind": "text", + "placeholders": [], + "refs": [ + { + "file": "features/models/CursorSettingsPage.tsx", + "line": 319, + "column": 27 + } + ] + }, "d86fa42c3848c680": { "source": "使用系统代理", "kind": "text", @@ -4675,7 +4748,7 @@ "refs": [ { "file": "features/models/CursorSettingsPage.tsx", - "line": 275, + "line": 315, "column": 47 } ] @@ -4699,12 +4772,12 @@ "refs": [ { "file": "features/models/CursorModelEditor.tsx", - "line": 110, + "line": 111, "column": 88 }, { "file": "features/models/CursorModelEditor.tsx", - "line": 155, + "line": 156, "column": 54 } ] @@ -4793,7 +4866,7 @@ "refs": [ { "file": "features/models/CursorModelCards.tsx", - "line": 174, + "line": 188, "column": 64 }, { @@ -4918,7 +4991,7 @@ "refs": [ { "file": "features/models/CursorSettingsPage.tsx", - "line": 209, + "line": 215, "column": 26 } ] @@ -5004,7 +5077,7 @@ "refs": [ { "file": "features/models/CursorModelEditor.tsx", - "line": 148, + "line": 149, "column": 54 } ] @@ -5086,12 +5159,12 @@ "refs": [ { "file": "features/models/CursorModelEditor.tsx", - "line": 141, + "line": 142, "column": 25 }, { "file": "features/models/CursorModelEditor.tsx", - "line": 141, + "line": 142, "column": 55 } ] @@ -5156,6 +5229,18 @@ } ] }, + "eb4a3db23661fb52": { + "source": "应用于该分组下的全部模型,并作为 Cursor 模型选择器中的徽章标签;清空则恢复显示服务器域名。", + "kind": "text", + "placeholders": [], + "refs": [ + { + "file": "features/models/CursorSettingsPage.tsx", + "line": 319, + "column": 44 + } + ] + }, "eb77492c9f76a7e1": { "source": "安装命令已自动复制。点击“打开终端”,将命令粘贴到终端中执行,并按提示输入密码。", "kind": "text", @@ -5163,7 +5248,7 @@ "refs": [ { "file": "features/models/CursorSettingsPage.tsx", - "line": 275, + "line": 315, "column": 77 } ] @@ -5192,7 +5277,7 @@ }, { "file": "features/models/CursorSettingsPage.tsx", - "line": 260, + "line": 300, "column": 69 } ] @@ -5204,7 +5289,7 @@ "refs": [ { "file": "features/models/CursorSettingsPage.tsx", - "line": 274, + "line": 314, "column": 53 } ] diff --git a/apps/desktop/src/i18n/locales/en-US.json b/apps/desktop/src/i18n/locales/en-US.json index aa941c2..d352612 100644 --- a/apps/desktop/src/i18n/locales/en-US.json +++ b/apps/desktop/src/i18n/locales/en-US.json @@ -79,6 +79,7 @@ "2f9daa828907b93f": "Delete", "2fe5a8d0eee9f14c": "Invalid", "303c30f301514250": "Search resources", + "3260348163d03b8e": "Leave blank to keep unchanged", "32896fdaaaa4c106": "Account saved and the model catalog is synced.", "346ff60e6c7c5181": "Reading…", "36f33adaf0942634": "Confirm", @@ -270,6 +271,7 @@ "b06325c5660f0c29": "Direct", "b16c3b2ecedd6fe1": "Cursor integration is active. Add a model configuration to use a BYOK model.", "b254ff315d861346": "Try initializing again", + "b2617bf9ae663752": "Group settings", "b4411558b932266f": "Provider type", "b4c9e08870d41aa2": "Initialize the plugin runtime first", "b502b1d414664337": "Prompt: {tokens}", @@ -290,6 +292,7 @@ "bb7efdcb6af6e805": "Default dark", "bda62ce1d5e4ace9": "Tell us why", "bda74b5674b6a57d": "Initialize plugins", + "be961dc60ab610da": "Applies to every model in this group when changed; leave blank to keep each model's current configuration.", "bf57afd709694b55": "Overview time range", "bfc01caf9fe0c841": "Cache hit rate {rate}", "c0b3fbff51ccc40b": "Done", @@ -322,6 +325,7 @@ "d60669bb26a22f5d": "Leave blank to use the default", "d6b1f203680f5496": "Leave blank to use adaptive thinking", "d766536c18e8e990": "Plugin runtime {version} is installed and ready to use.", + "d7e266bdc8064193": "Group name", "d86fa42c3848c680": "Use system proxy", "d8c47e9776cf1082": "Main menu", "da521d1c1cbd36af": "Authorization is required to install the certificate", @@ -361,6 +365,7 @@ "ea26b760e930a7ca": "Call observability", "eb11e2df1d8ae387": "Provider URL", "eb1be07f2ca6e506": "Estimated using Claude Opus 4.7 pricing.", + "eb4a3db23661fb52": "Applies to every model in this group and is used as the badge label in Cursor's model picker; clear it to fall back to the server domain.", "eb77492c9f76a7e1": "The install command has been copied. Click “Open terminal”, paste it into the terminal, and enter your password when prompted.", "eba54690937bc532": "Manage accounts", "ed31fbb483ee1b0a": "Actions", diff --git a/apps/desktop/src/i18n/locales/zh-CN.json b/apps/desktop/src/i18n/locales/zh-CN.json index 9aff95e..53ed05e 100644 --- a/apps/desktop/src/i18n/locales/zh-CN.json +++ b/apps/desktop/src/i18n/locales/zh-CN.json @@ -79,6 +79,7 @@ "2f9daa828907b93f": "删除", "2fe5a8d0eee9f14c": "已失效", "303c30f301514250": "搜索资源", + "3260348163d03b8e": "留空保持不变", "32896fdaaaa4c106": "账号已保存,模型目录已同步。", "346ff60e6c7c5181": "读取中…", "36f33adaf0942634": "确认", @@ -270,6 +271,7 @@ "b06325c5660f0c29": "直连", "b16c3b2ecedd6fe1": "Cursor 接管已生效;添加模型配置后即可使用 BYOK 模型。", "b254ff315d861346": "请重试初始化", + "b2617bf9ae663752": "分组设置", "b4411558b932266f": "上游类型", "b4c9e08870d41aa2": "需要先初始化插件运行时", "b502b1d414664337": "提示词:{tokens}", @@ -290,6 +292,7 @@ "bb7efdcb6af6e805": "默认暗色", "bda62ce1d5e4ace9": "可以告诉我们原因", "bda74b5674b6a57d": "初始化插件", + "be961dc60ab610da": "修改后应用于该分组下的全部模型;留空保持各模型现有配置不变。", "bf57afd709694b55": "概览时间范围", "bfc01caf9fe0c841": "缓存命中率 {rate}", "c0b3fbff51ccc40b": "完成", @@ -322,6 +325,7 @@ "d60669bb26a22f5d": "留空使用默认值", "d6b1f203680f5496": "留空使用 adaptive thinking", "d766536c18e8e990": "插件运行时 {version} 已安装,可以开始使用插件。", + "d7e266bdc8064193": "分组名称", "d86fa42c3848c680": "使用系统代理", "d8c47e9776cf1082": "主菜单", "da521d1c1cbd36af": "需要授权安装证书", @@ -361,6 +365,7 @@ "ea26b760e930a7ca": "调用观测", "eb11e2df1d8ae387": "上游地址", "eb1be07f2ca6e506": "按 Claude Opus 4.7 价格估算。", + "eb4a3db23661fb52": "应用于该分组下的全部模型,并作为 Cursor 模型选择器中的徽章标签;清空则恢复显示服务器域名。", "eb77492c9f76a7e1": "安装命令已自动复制。点击“打开终端”,将命令粘贴到终端中执行,并按提示输入密码。", "eba54690937bc532": "账号管理", "ed31fbb483ee1b0a": "操作", diff --git a/apps/desktop/src/shared/api.ts b/apps/desktop/src/shared/api.ts index 03d8e22..b5ecc91 100644 --- a/apps/desktop/src/shared/api.ts +++ b/apps/desktop/src/shared/api.ts @@ -7,6 +7,7 @@ export interface Model { model_hash: string; sort_order: number; display_name: string; + group_name: string | null; type: ModelType; base_url: string; use_full_url: boolean; @@ -33,6 +34,7 @@ export interface Model { export interface ModelInput { sort_order: number; display_name: string; + group_name: string | null; type: ModelType; base_url: string; use_full_url: boolean; diff --git a/server/migrations/0008_add_model_group_name.sql b/server/migrations/0008_add_model_group_name.sql new file mode 100644 index 0000000..a98c812 --- /dev/null +++ b/server/migrations/0008_add_model_group_name.sql @@ -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; diff --git a/server/src/api/cursor/handlers.rs b/server/src/api/cursor/handlers.rs index 888770e..6a608d2 100644 --- a/server/src/api/cursor/handlers.rs +++ b/server/src/api/cursor/handlers.rs @@ -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 { 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) } diff --git a/server/src/app.rs b/server/src/app.rs index 7b4a536..d4ec4b1 100644 --- a/server/src/app.rs +++ b/server/src/app.rs @@ -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)?; diff --git a/server/src/cursor/compile/context.rs b/server/src/cursor/compile/context.rs index f7b0a4b..e29e634 100644 --- a/server/src/cursor/compile/context.rs +++ b/server/src/cursor/compile/context.rs @@ -146,6 +146,37 @@ async fn decode_part( .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::>(); + 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::>(); + 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()); + } +} diff --git a/server/src/cursor/compile/run.rs b/server/src/cursor/compile/run.rs index 740ed57..f351ead 100644 --- a/server/src/cursor/compile/run.rs +++ b/server/src/cursor/compile/run.rs @@ -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, diff --git a/server/src/cursor/conversation/registry.rs b/server/src/cursor/conversation/registry.rs index d41479c..7a97a40 100644 --- a/server/src/cursor/conversation/registry.rs +++ b/server/src/cursor/conversation/registry.rs @@ -26,6 +26,8 @@ pub(crate) struct ConversationDependencies { pub provider: Arc, pub compiler: PromptCompiler, pub web_cache: WebCache, + /// 本地 rules 服务的 md 存储目录;编译请求上下文时合并其中的规则。 + pub local_rules_dir: Option, } struct RegistryInner { @@ -47,6 +49,7 @@ impl ConversationRegistry { provider: Arc, compiler: PromptCompiler, web_cache: WebCache, + local_rules_dir: Option, ) -> Self { Self { inner: Arc::new(RegistryInner { @@ -58,6 +61,7 @@ impl ConversationRegistry { provider, compiler, web_cache, + local_rules_dir, }, }), } diff --git a/server/src/cursor/conversation/runtime.rs b/server/src/cursor/conversation/runtime.rs index 0a37d91..6b9f939 100644 --- a/server/src/cursor/conversation/runtime.rs +++ b/server/src/cursor/conversation/runtime.rs @@ -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, }; diff --git a/server/src/cursor/services/knowledge/mod.rs b/server/src/cursor/services/knowledge/mod.rs new file mode 100644 index 0000000..2b62e1b --- /dev/null +++ b/server/src/cursor/services/knowledge/mod.rs @@ -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, +} + +#[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, + #[prost(string, optional, tag = "2")] + git_origin: Option, +} + +#[derive(Clone, PartialEq, Message)] +pub(crate) struct KnowledgeBaseListResponse { + #[prost(bool, tag = "1")] + success: bool, + #[prost(message, repeated, tag = "2")] + all_results: Vec, +} + +#[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, +} + +struct Inner { + store: RuleStore, + lock: tokio::sync::Mutex<()>, +} + +impl KnowledgeService { + pub fn managed() -> Result { + Self::with_root(config::managed_data_dir()?.join("rules")) + } + + /// 指定存储根目录构造;managed() 与集成测试共用。 + pub fn with_root(root: std::path::PathBuf) -> Result { + Ok(Self { + inner: Arc::new(Inner { + store: RuleStore::open(root)?, + lock: tokio::sync::Mutex::new(()), + }), + }) + } +} + +pub async fn add( + Extension(upstream): Extension, + Extension(service): Extension, + request: Request, +) -> Result> { + 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::(&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, + Extension(service): Extension, + request: Request, +) -> Result> { + 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::(&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, + Extension(service): Extension, + request: Request, +) -> Result> { + 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::(&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, + Extension(service): Extension, + request: Request, +) -> Result> { + 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::(&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) -> 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> { + 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) +} diff --git a/server/src/cursor/services/knowledge/store.rs b/server/src/cursor/services/knowledge/store.rs new file mode 100644 index 0000000..35e761f --- /dev/null +++ b/server/src/cursor/services/knowledge/store.rs @@ -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, + #[serde(default)] + journal: Vec, +} + +#[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 { + std::fs::create_dir_all(&root)?; + Ok(Self { root }) + } + + pub fn list(&self) -> Result> { + 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> { + 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> { + 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::::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 { + 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::>(), + ["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"); + } +} diff --git a/server/src/cursor/services/knowledge/sync.rs b/server/src/cursor/services/knowledge/sync.rs new file mode 100644 index 0000000..64e3eb9 --- /dev/null +++ b/server/src/cursor/services/knowledge/sync.rs @@ -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 { + 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) -> 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::>(); + store.replace_all(&records) +} + +async fn replay_add( + upstream: &proxy::CursorProxy, + headers: &HeaderMap, + store: &RuleStore, + id: &str, +) -> Result { + 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::(&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 { + 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::(&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 { + 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::(&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 { + 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 + } + } +} diff --git a/server/src/cursor/services/mod.rs b/server/src/cursor/services/mod.rs index 7c3c020..4ab2ec1 100644 --- a/server/src/cursor/services/mod.rs +++ b/server/src/cursor/services/mod.rs @@ -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; diff --git a/server/src/cursor/services/model_catalog.rs b/server/src/cursor/services/model_catalog.rs index 1d924d6..4c088dd 100644 --- a/server/src/cursor/services/model_catalog.rs +++ b/server/src/cursor/services/model_catalog.rs @@ -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, }], diff --git a/server/src/cursor/transport/registry.rs b/server/src/cursor/transport/registry.rs index 70b9513..db57032 100644 --- a/server/src/cursor/transport/registry.rs +++ b/server/src/cursor/transport/registry.rs @@ -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, + 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, + local_rules_dir: Option, ) -> Self { Self { inner: Arc::new(RegistryInner { @@ -80,6 +106,7 @@ impl TransportRegistry { provider, compiler, web_cache.clone(), + local_rules_dir, ), store, web_cache, diff --git a/server/src/local_app/proxy.rs b/server/src/local_app/proxy.rs index 0876de5..c67810c 100644 --- a/server/src/local_app/proxy.rs +++ b/server/src/local_app/proxy.rs @@ -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" ) diff --git a/server/src/model/configuration.rs b/server/src/model/configuration.rs index be45ce9..1fe7111 100644 --- a/server/src/model/configuration.rs +++ b/server/src/model/configuration.rs @@ -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, #[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, #[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 { 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 Result { Ok(ModelConfigInput { sort_order: model.sort, display_name: model.display_name.clone(), + group_name: None, model_type, base_url, use_full_url, diff --git a/server/src/store/migrations.rs b/server/src/store/migrations.rs index 405739f..1a06448 100644 --- a/server/src/store/migrations.rs +++ b/server/src/store/migrations.rs @@ -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); } } diff --git a/server/src/store/models.rs b/server/src/store/models.rs index 1c21bef..0e6eeef 100644 --- a/server/src/store/models.rs +++ b/server/src/store/models.rs @@ -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 { 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 { 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 { }) } +#[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> { row.try_get::, _>(column)? .map(|value| { diff --git a/server/tests/compaction.rs b/server/tests/compaction.rs index 910977c..46b8ea2 100644 --- a/server/tests/compaction.rs +++ b/server/tests/compaction.rs @@ -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, diff --git a/server/tests/interrupt.rs b/server/tests/interrupt.rs index d480a1f..d64dc11 100644 --- a/server/tests/interrupt.rs +++ b/server/tests/interrupt.rs @@ -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, diff --git a/server/tests/knowledge_rules.rs b/server/tests/knowledge_rules.rs new file mode 100644 index 0000000..1dc4bc5 --- /dev/null +++ b/server/tests/knowledge_rules.rs @@ -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, +} + +#[derive(Clone, PartialEq, Message)] +struct ListResponse { + #[prost(bool, tag = "1")] + success: bool, + #[prost(message, repeated, tag = "2")] + all_results: Vec, +} + +#[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 { + Request::post("/test") + .header(header::CONTENT_TYPE, "application/proto") + .body(Body::from(message.encode_to_vec())) + .unwrap() +} + +async fn decode(response: Response) -> 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); +} diff --git a/server/tests/local_rules_context.rs b/server/tests/local_rules_context.rs new file mode 100644 index 0000000..67d2bdc --- /dev/null +++ b/server/tests/local_rules_context.rs @@ -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::>(); + assert_eq!( + context_texts.len(), + 1, + "exactly one request-context message is projected" + ); + assert!( + context_texts[0].contains("\nAlways answer in haiku.\n"), + "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() + }, + )), + } +}