Compare commits

..
Author SHA1 Message Date
leookun c274a9db4c refactor: remove CursorAccountCard component and related localization entries
- Deleted the CursorAccountCard.vue component, which handled user account login and status.
- Removed associated localization entries from catalog.json and various language files.
- Updated Home.vue to eliminate references to the removed component.
- Refactored clientApi.js to remove unused account-related API functions.
2026-08-08 21:19:57 +08:00
104 changed files with 1950 additions and 6135 deletions
+1 -13
View File
@@ -2,11 +2,6 @@
# cursor-byok
cursor-byok 是 Cursor 后端的本地实现。
<br>
<br>
<a href="https://trendshift.io/repositories/39260?utm_source=repository-badge&amp;utm_medium=badge&amp;utm_campaign=badge-repository-39260" target="_blank" rel="noopener noreferrer"><img src="https://trendshift.io/api/badge/repositories/39260" alt="leookun/cursor-byok | Trendshift" width="250" height="55" /></a>
[使用教程](https://docs.leokun.cn) · [下载最新版](https://github.com/leookun/cursor-byok/releases/latest) · [问题反馈](https://github.com/leookun/cursor-byok/issues) · [English](./README.md)
[![Release](https://img.shields.io/github/v/release/leookun/cursor-byok?style=flat-square)](https://github.com/leookun/cursor-byok/releases/latest)
@@ -89,19 +84,12 @@ cursor-byok 在本机负责协议适配、模型请求转发、工具调用衔
- [Telegram 交流群](https://t.me/cursor_byok)
- QQ 交流群:`1095916242``1094411438``1095918002``1094419321`
<a href="https://trendshift.io/repositories/39260?utm_source=repository-badge&amp;utm_medium=badge&amp;utm_campaign=badge-repository-39260" target="_blank" rel="noopener noreferrer"><img src="https://trendshift.io/api/badge/repositories/39260" alt="leookun/cursor-byok | Trendshift" width="250" height="55" /></a>
## 开发与贡献
欢迎提交 Issue 和 Pull Request。开发环境、构建命令、项目结构及提交规范请阅读 [贡献指南](./CONTRIBUTING.md)。
## 贡献者名单
<a href="https://github.com/leookun/cursor-byok/graphs/contributors">
<img src="https://contrib.rocks/image?repo=leookun/cursor-byok" />
</a>
## 许可证
本项目基于 [MIT License](./LICENSE) 开源。
+2 -17
View File
@@ -1,20 +1,14 @@
<div align="center">
# cursor-byok
cursor-byok is a local implementation of Cursor's backend.
<br>
<br>
<a href="https://trendshift.io/repositories/39260?utm_source=repository-badge&amp;utm_medium=badge&amp;utm_campaign=badge-repository-39260" target="_blank" rel="noopener noreferrer"><img src="https://trendshift.io/api/badge/repositories/39260" alt="leookun/cursor-byok | Trendshift" width="250" height="55" /></a>
[User Guide](https://docs.leokun.cn) · [Download](https://github.com/leookun/cursor-byok/releases/latest) · [Report an Issue](https://github.com/leookun/cursor-byok/issues) · [中文版本说明](./README-CN.md)
[User Guide](https://docs.leokun.cn) · [Latest Release](https://github.com/leookun/cursor-byok/releases/latest) · [Report an Issue](https://github.com/leookun/cursor-byok/issues) · [简体中文](./README-CN.md)
[![Release](https://img.shields.io/github/v/release/leookun/cursor-byok?style=flat-square)](https://github.com/leookun/cursor-byok/releases/latest)
[![Downloads](https://img.shields.io/github/downloads/leookun/cursor-byok/total?style=flat-square)](https://github.com/leookun/cursor-byok/releases)
[![License](https://img.shields.io/github/license/leookun/cursor-byok?style=flat-square)](./LICENSE)
[![Platforms](https://img.shields.io/badge/platform-macOS%20%7C%20Windows%20%7C%20Linux-lightgrey?style=flat-square)](https://github.com/leookun/cursor-byok/releases/latest)
</div>
![Connect cursor-byok to a wide range of model APIs](./images/en-brand.png)
@@ -90,21 +84,12 @@ See the [release roadmap](https://github.com/leookun/cursor-byok/discussions/32)
- [Telegram community](https://t.me/cursor_byok)
- QQ groups: `1095916242`, `1094411438`, `1095918002`, `1094419321`
<a href="https://trendshift.io/repositories/39260?utm_source=repository-badge&amp;utm_medium=badge&amp;utm_campaign=badge-repository-39260" target="_blank" rel="noopener noreferrer"><img src="https://trendshift.io/api/badge/repositories/39260" alt="leookun/cursor-byok | Trendshift" width="250" height="55" /></a>
## Development and Contributing
Issues and pull requests are welcome. See the [Contributing Guide](./CONTRIBUTING_EN.md) for prerequisites, build commands, project structure, and contribution guidelines.
## Contributors
<a href="https://github.com/leookun/cursor-byok/graphs/contributors">
<img src="https://contrib.rocks/image?repo=leookun/cursor-byok" />
</a>
## License
This project is open source under the [MIT License](./LICENSE).
+1 -1
View File
@@ -8,7 +8,7 @@ info:
description: "Cursor助手"
copyright: "© 2026, Cursor助手"
comments: "Cursor助手"
version: "0.0.48"
version: "0.0.46"
dev_mode:
root_path: .
+2 -2
View File
@@ -17,9 +17,9 @@
<key>CFBundlePackageType</key>
<string>APPL</string>
<key>CFBundleShortVersionString</key>
<string>0.0.48</string>
<string>0.0.46</string>
<key>CFBundleVersion</key>
<string>0.0.48</string>
<string>0.0.46</string>
<key>LSMinimumSystemVersion</key>
<string>12.0.0</string>
<key>LSUIElement</key>
+2 -2
View File
@@ -17,9 +17,9 @@
<key>CFBundlePackageType</key>
<string>APPL</string>
<key>CFBundleShortVersionString</key>
<string>0.0.48</string>
<string>0.0.46</string>
<key>CFBundleVersion</key>
<string>0.0.48</string>
<string>0.0.46</string>
<key>LSMinimumSystemVersion</key>
<string>12.0.0</string>
<key>LSUIElement</key>
+15 -28
View File
@@ -6,7 +6,7 @@
name: "Cursor助手"
arch: ${GOARCH}
platform: "linux"
version: "0.0.48"
version: "0.0.46"
section: "default"
priority: "extra"
maintainer: ${GIT_COMMITTER_NAME} <${GIT_COMMITTER_EMAIL}>
@@ -24,24 +24,24 @@ contents:
- src: "./build/linux/Cursor助手.desktop"
dst: "/usr/share/applications/Cursor助手.desktop"
# Default dependencies for the GTK4 + WebKitGTK 6.0 stack (Ubuntu 24.04+ / Debian 13+)
# Default dependencies for Debian 12/Ubuntu 22.04+ with WebKit 4.1
depends:
- libgtk-4-1
- libwebkitgtk-6.0-4
- libgtk-3-0
- libwebkit2gtk-4.1-0
# Distribution-specific overrides for different package formats
# Distribution-specific overrides for different package formats and WebKit versions
overrides:
# RPM packages for Fedora / RHEL / AlmaLinux / Rocky Linux
# RPM packages for RHEL/CentOS/AlmaLinux/Rocky Linux (WebKit 4.0)
rpm:
depends:
- gtk4
- webkitgtk6.0
# Arch Linux packages
- gtk3
- webkit2gtk4.1
# Arch Linux packages (WebKit 4.1)
archlinux:
depends:
- gtk4
- webkitgtk-6.0
- gtk3
- webkit2gtk-4.1
# scripts section to ensure desktop database is updated after install
scripts:
@@ -50,26 +50,13 @@ scripts:
# preremove: "./build/linux/nfpm/scripts/preremove.sh"
# postremove: "./build/linux/nfpm/scripts/postremove.sh"
# If you build your app with -tags gtk3 (legacy WebKit2GTK 4.1 stack — supported through v3.0.x, removed in v3.1),
# replace the depends/overrides above with these:
#
# depends:
# - libgtk-3-0
# - libwebkit2gtk-4.1-0
# overrides:
# rpm:
# depends:
# - gtk3
# - webkit2gtk4.1
# archlinux:
# depends:
# - gtk3
# - webkit2gtk-4.1
#
# replaces:
# - foobar
# provides:
# - bar
# depends:
# - gtk3
# - libwebkit2gtk
# recommends:
# - whatever
# suggests:
+2 -2
View File
@@ -1,10 +1,10 @@
{
"fixed": {
"file_version": "0.0.48"
"file_version": "0.0.46"
},
"info": {
"0000": {
"ProductVersion": "0.0.48",
"ProductVersion": "0.0.46",
"CompanyName": "Cursor助手",
"FileDescription": "Cursor助手",
"LegalCopyright": "© 2026, Cursor助手",
+12 -37
View File
@@ -14,7 +14,7 @@
!define INFO_PRODUCTNAME "Cursor助手"
!endif
!ifndef INFO_PRODUCTVERSION
!define INFO_PRODUCTVERSION "0.0.48"
!define INFO_PRODUCTVERSION "0.0.46"
!endif
!ifndef INFO_COPYRIGHT
!define INFO_COPYRIGHT "© 2026, Cursor助手"
@@ -27,16 +27,8 @@
!endif
!define UNINST_KEY "Software\Microsoft\Windows\CurrentVersion\Uninstall\${UNINST_KEY_NAME}"
!ifndef WAILS_INSTALL_SCOPE
!define WAILS_INSTALL_SCOPE "machine"
!endif
!ifndef REQUEST_EXECUTION_LEVEL
!if "${WAILS_INSTALL_SCOPE}" == "user"
!define REQUEST_EXECUTION_LEVEL "user"
!else
!define REQUEST_EXECUTION_LEVEL "admin"
!endif
!define REQUEST_EXECUTION_LEVEL "admin"
!endif
RequestExecutionLevel "${REQUEST_EXECUTION_LEVEL}"
@@ -123,40 +115,23 @@ RequestExecutionLevel "${REQUEST_EXECUTION_LEVEL}"
WriteUninstaller "$INSTDIR\uninstall.exe"
SetRegView 64
!if "${WAILS_INSTALL_SCOPE}" == "user"
WriteRegStr HKCU "${UNINST_KEY}" "Publisher" "${INFO_COMPANYNAME}"
WriteRegStr HKCU "${UNINST_KEY}" "DisplayName" "${INFO_PRODUCTNAME}"
WriteRegStr HKCU "${UNINST_KEY}" "DisplayVersion" "${INFO_PRODUCTVERSION}"
WriteRegStr HKCU "${UNINST_KEY}" "DisplayIcon" "$INSTDIR\${PRODUCT_EXECUTABLE}"
WriteRegStr HKCU "${UNINST_KEY}" "UninstallString" "$\"$INSTDIR\uninstall.exe$\""
WriteRegStr HKCU "${UNINST_KEY}" "QuietUninstallString" "$\"$INSTDIR\uninstall.exe$\" /S"
WriteRegStr HKLM "${UNINST_KEY}" "Publisher" "${INFO_COMPANYNAME}"
WriteRegStr HKLM "${UNINST_KEY}" "DisplayName" "${INFO_PRODUCTNAME}"
WriteRegStr HKLM "${UNINST_KEY}" "DisplayVersion" "${INFO_PRODUCTVERSION}"
WriteRegStr HKLM "${UNINST_KEY}" "DisplayIcon" "$INSTDIR\${PRODUCT_EXECUTABLE}"
WriteRegStr HKLM "${UNINST_KEY}" "UninstallString" "$\"$INSTDIR\uninstall.exe$\""
WriteRegStr HKLM "${UNINST_KEY}" "QuietUninstallString" "$\"$INSTDIR\uninstall.exe$\" /S"
${GetSize} "$INSTDIR" "/S=0K" $0 $1 $2
IntFmt $0 "0x%08X" $0
WriteRegDWORD HKCU "${UNINST_KEY}" "EstimatedSize" "$0"
!else
WriteRegStr HKLM "${UNINST_KEY}" "Publisher" "${INFO_COMPANYNAME}"
WriteRegStr HKLM "${UNINST_KEY}" "DisplayName" "${INFO_PRODUCTNAME}"
WriteRegStr HKLM "${UNINST_KEY}" "DisplayVersion" "${INFO_PRODUCTVERSION}"
WriteRegStr HKLM "${UNINST_KEY}" "DisplayIcon" "$INSTDIR\${PRODUCT_EXECUTABLE}"
WriteRegStr HKLM "${UNINST_KEY}" "UninstallString" "$\"$INSTDIR\uninstall.exe$\""
WriteRegStr HKLM "${UNINST_KEY}" "QuietUninstallString" "$\"$INSTDIR\uninstall.exe$\" /S"
${GetSize} "$INSTDIR" "/S=0K" $0 $1 $2
IntFmt $0 "0x%08X" $0
WriteRegDWORD HKLM "${UNINST_KEY}" "EstimatedSize" "$0"
!endif
${GetSize} "$INSTDIR" "/S=0K" $0 $1 $2
IntFmt $0 "0x%08X" $0
WriteRegDWORD HKLM "${UNINST_KEY}" "EstimatedSize" "$0"
!macroend
!macro wails.deleteUninstaller
Delete "$INSTDIR\uninstall.exe"
SetRegView 64
!if "${WAILS_INSTALL_SCOPE}" == "user"
DeleteRegKey HKCU "${UNINST_KEY}"
!else
DeleteRegKey HKLM "${UNINST_KEY}"
!endif
DeleteRegKey HKLM "${UNINST_KEY}"
!macroend
!macro wails.setShellContext
+1 -1
View File
@@ -1,6 +1,6 @@
<?xml version="1.0" encoding="UTF-8" standalone="yes"?>
<assembly manifestVersion="1.0" xmlns="urn:schemas-microsoft-com:asm.v1" xmlns:asmv3="urn:schemas-microsoft-com:asm.v3">
<assemblyIdentity type="win32" name="com.cursor.wuxianxubei" version="0.0.48" processorArchitecture="*"/>
<assemblyIdentity type="win32" name="com.cursor.wuxianxubei" version="0.0.46" processorArchitecture="*"/>
<dependency>
<dependentAssembly>
<assemblyIdentity type="win32" name="Microsoft.Windows.Common-Controls" version="6.0.0.0" processorArchitecture="*" publicKeyToken="6595b64144ccf1df" language="*"/>
+1 -3
View File
@@ -28,7 +28,6 @@ type exchangeContext struct {
type Server struct {
config Config
certManager *certs.Manager
caCertPEM []byte
store *exchangeStore
counter atomic.Uint64
proxyServer *http.Server
@@ -44,14 +43,13 @@ func New(config Config) (*Server, error) {
if err := validateLoopbackAddress(config.UIAddr); err != nil {
return nil, err
}
manager, caCertPEM, err := certs.NewGeneratedManager()
manager, err := certs.NewEmbeddedManager()
if err != nil {
return nil, fmt.Errorf("加载 MITM CA 失败:%w", err)
}
server := &Server{
config: config,
certManager: manager,
caCertPEM: caCertPEM,
store: newExchangeStore(config.MaxExchanges),
}
proxyHandler, err := server.newProxyHandler()
+3 -1
View File
@@ -8,6 +8,8 @@ import (
"net/http"
"strings"
"time"
"cursor/internal/certs"
)
//go:embed web/*
@@ -91,7 +93,7 @@ func (server *Server) handleEvents(writer http.ResponseWriter, request *http.Req
func (server *Server) handleCACertificate(writer http.ResponseWriter, _ *http.Request) {
writer.Header().Set("Content-Type", "application/x-x509-ca-cert")
writer.Header().Set("Content-Disposition", `attachment; filename="cursor-local-proxy-ca.crt"`)
_, _ = writer.Write(server.caCertPEM)
_, _ = writer.Write(certs.EmbeddedCACertPEM())
}
func writeJSON(writer http.ResponseWriter, status int, payload any) {
@@ -1,202 +0,0 @@
<script setup>
import Button from "@/components/ui/Button.vue";
import Card from "@/components/ui/Card.vue";
import Tooltip from "@/components/ui/Tooltip.vue";
import { useMessage } from "@/composables/useMessage";
import { showModal } from "@/composables/useModal";
import {
disconnectCursorAccount,
getCursorAccountStatus,
startCursorAccountLogin,
} from "@/services/clientApi";
import { toUserError } from "@/state/appState";
import { Browser } from "@wailsio/runtime";
import { computed, onMounted, onUnmounted, ref } from "vue";
const CURSOR_ACCOUNT_CONTRIBUTOR_URL = "https://github.com/aike0210";
const message = useMessage();
const cursorAccountStatus = ref({
state: "signed_out",
authId: "",
email: "",
error: "",
});
const cursorAccountBusy = ref(false);
let cursorAccountTimer = null;
function maskCursorAccountIdentifier(value) {
const identifier = String(value || "").trim();
if (!identifier) return "";
const atIndex = identifier.indexOf("@");
if (atIndex > 0 && atIndex < identifier.length - 1) {
const localPart = identifier.slice(0, atIndex);
const domain = identifier.slice(atIndex + 1);
const maskedLocalPart = localPart.length <= 2
? `${localPart[0]}***`
: `${localPart[0]}***${localPart.at(-1)}`;
return `${maskedLocalPart}@${domain}`;
}
if (identifier.length <= 8) return "****";
return `${identifier.slice(0, 4)}****${identifier.slice(-4)}`;
}
const cursorAccountSignedIn = computed(
() => cursorAccountStatus.value.state === "signed_in",
);
const cursorAccountWaiting = computed(
() => cursorAccountStatus.value.state === "waiting",
);
const cursorAccountDisplayIdentifier = computed(() => {
if (!cursorAccountSignedIn.value) return "";
return maskCursorAccountIdentifier(
cursorAccountStatus.value.email || cursorAccountStatus.value.authId,
);
});
const cursorAccountStateText = computed(() => {
if (cursorAccountSignedIn.value) return "已经登录";
if (cursorAccountWaiting.value) return "等待浏览器登录";
return "未连接";
});
function showActionError(title, error) {
const detail = String(error || "服务错误").trim() || "服务错误";
message(`${title}${detail}`);
}
async function handleOpenContributor() {
try {
await Browser.OpenURL(CURSOR_ACCOUNT_CONTRIBUTOR_URL);
} catch (error) {
showActionError("打开贡献者主页失败", toUserError(error));
}
}
async function refreshCursorAccountStatus() {
cursorAccountStatus.value = await getCursorAccountStatus();
}
async function handleCursorAccountLogin() {
cursorAccountBusy.value = true;
try {
cursorAccountStatus.value = await startCursorAccountLogin();
} catch (error) {
showActionError("登录失败", toUserError(error));
await refreshCursorAccountStatus().catch(() => {});
} finally {
cursorAccountBusy.value = false;
}
}
async function handleCursorAccountDisconnect() {
const confirmed = await showModal({
title: "退出登录",
content: "只会退出 cursor-byok 中的 Cursor 账号,不会退出 Cursor 客户端。是否继续?",
confirmText: "退出登录",
cancelText: "取消",
showCancel: true,
});
if (!confirmed) return;
cursorAccountBusy.value = true;
try {
cursorAccountStatus.value = await disconnectCursorAccount();
} catch (error) {
showActionError("退出登录失败", toUserError(error));
} finally {
cursorAccountBusy.value = false;
}
}
onMounted(async () => {
await refreshCursorAccountStatus().catch(() => {});
cursorAccountTimer = window.setInterval(() => {
if (cursorAccountWaiting.value) {
void refreshCursorAccountStatus().catch(() => {});
}
}, 1500);
});
onUnmounted(() => {
if (cursorAccountTimer) {
window.clearInterval(cursorAccountTimer);
cursorAccountTimer = null;
}
});
</script>
<template>
<Card>
<div class="flex flex-col gap-3">
<div class="flex items-center justify-between gap-4">
<div class="flex min-w-0 flex-wrap items-center gap-2">
<h2 class="text-base font-medium text-white">Cursor 控制面账号</h2>
<span
class="rounded-full border border-[#3a3a3a] bg-[#202020] px-2 py-0.5 text-xs text-[#b8b8b8]"
>
{{ cursorAccountStateText }}
</span>
</div>
<div class="flex shrink-0 items-center gap-1 text-xs text-[#737373]">
<span>@aike0210</span>
<Tooltip>
<div class="flex min-w-[220px] flex-col gap-2">
<div>感谢 @aike0210 Cursor 控制面账号功能的贡献</div>
<button
type="button"
class="flex items-center gap-2 text-left text-[#8ab4f8] transition-colors duration-150 hover:text-[#b6d0fb]"
@click="handleOpenContributor"
>
<span class="icon-[mdi--github] text-[14px]"></span>
<span>github.com/aike0210</span>
<span class="icon-[mdi--open-in-new] text-[12px]"></span>
</button>
</div>
</Tooltip>
</div>
</div>
<div class="flex items-end justify-between gap-4">
<div class="min-w-0">
<div
v-if="cursorAccountDisplayIdentifier"
class="truncate text-sm text-[#d0d0d0]"
>
{{ cursorAccountDisplayIdentifier }}
</div>
<div class="mt-1 text-sm text-[#a3a3a3]">
独立用于插件Skills MCP不会改变 Cursor 客户端当前账号
</div>
<div v-if="cursorAccountWaiting" class="mt-1 text-sm text-[#d6a84b]">
请在浏览器完成登录完成后返回 Cursor 重新打开插件市场
</div>
<div
v-if="cursorAccountStatus.error"
class="mt-1 break-all text-sm text-[#e06c75]"
>
{{ cursorAccountStatus.error }}
</div>
</div>
<Button
v-if="cursorAccountSignedIn"
class="shrink-0"
:disabled="cursorAccountBusy"
@click="handleCursorAccountDisconnect"
>
退出登录
</Button>
<Button
v-else
class="shrink-0"
variant="primary"
:disabled="cursorAccountBusy || cursorAccountWaiting"
@click="handleCursorAccountLogin"
>
{{ cursorAccountWaiting ? "等待登录..." : "登录 Cursor" }}
</Button>
</div>
</div>
</Card>
</template>
@@ -19,7 +19,6 @@ const modelTypeOptions = [
];
const reasoningEffortOptions = [
{ label: "不设置", value: "", icon: "icon-[mdi--minus-circle-outline]" },
{ label: "低", value: "low", icon: "icon-[mdi--head-outline]" },
{ label: "中", value: "medium", icon: "icon-[mdi--head-lightbulb-outline]" },
{ label: "高", value: "high", icon: "icon-[mdi--brain]" },
+1 -2
View File
@@ -35,7 +35,6 @@ const modelTypeTabs = [
];
const reasoningEffortOptions = [
{ label: "不设置", value: "", icon: "icon-[mdi--minus-circle-outline]" },
{ label: "低", value: "low", icon: "icon-[mdi--head-outline]" },
{ label: "中", value: "medium", icon: "icon-[mdi--head-lightbulb-outline]" },
{ label: "高", value: "high", icon: "icon-[mdi--brain]" },
@@ -151,7 +150,7 @@ const fieldTips = {
baseURL: "模型服务的 API 根地址,通常为兼容 OpenAI 或 Anthropic 的接口入口。",
apiKey: "调用该模型服务需要使用的访问密钥。",
contextWindowTokens: "模型单次可接受的最大上下文 Token 数。留空时使用默认值。",
reasoningEffort: "仅当模型支持 reasoning_effort 时才选择推理强度;选择“不设置”后,请求不会携带该参数。越高通常越稳,但也可能更慢。",
reasoningEffort: "推理强度仅对部分支持 reasoning_effort 的模型生效,并不是所有模型都支持。越高通常越稳,但也可能更慢。",
maxCompletionTokens: "单次回复允许生成的最大 Token 数。留空时使用默认值。",
openAIEndpoint: "选择接口协议端点。选“自定义路径”时,请在接口地址栏填写完整请求地址(含 /chat/completions 或 /responses 路径后缀),系统会根据末段自动判断协议形态。",
openAIExtraParams: "开启后会把 JSON 对象覆盖到 OpenAI 请求体。同名字段以这里为准。OpenAI service_tier 支持 auto、default、flex、scale、priority。",
@@ -1,92 +0,0 @@
import { ref } from "vue";
import { Dialogs } from "@wailsio/runtime";
import { showModal } from "@/composables/useModal";
import {
appState,
exportUserConfigToFile,
importUserConfigFromFile,
toUserError,
} from "@/state/appState";
const YAML_FILE_FILTER = [{ DisplayName: "YAML 配置文件", Pattern: "*.yaml;*.yml" }];
export function useConfigTransfer({ message, showActionError }) {
const busy = ref(false);
async function exportConfig() {
const confirmed = await showModal({
title: "导出完整配置",
content: "导出文件会包含全部模型配置、API Key 和自定义请求头。请将文件保存在安全位置,避免泄露;Windows 上的访问权限取决于目标文件夹的安全设置。",
confirmText: "继续导出",
cancelText: "取消",
showCancel: true,
});
if (!confirmed) {
return;
}
try {
const path = await Dialogs.SaveFile({
Title: "导出 cursor-byok 配置",
Filename: "cursor-byok-config.yaml",
Filters: YAML_FILE_FILTER,
CanCreateDirectories: true,
AllowsOtherFiletypes: false,
});
if (!path) {
return;
}
busy.value = true;
const exportedPath = await exportUserConfigToFile(path);
message(`配置已导出到 ${exportedPath}`);
} catch (error) {
showActionError("导出失败", toUserError(error));
} finally {
busy.value = false;
}
}
async function importConfig() {
if (appState.serviceRunning || appState.backendRunning || appState.proxyRunning) {
showActionError("导入失败", "服务运行中不能导入完整配置,请先停止服务");
return;
}
try {
const path = await Dialogs.OpenFile({
Title: "导入 cursor-byok 配置",
Filters: YAML_FILE_FILTER,
CanChooseFiles: true,
CanChooseDirectories: false,
AllowsMultipleSelection: false,
AllowsOtherFiletypes: false,
});
if (!path) {
return;
}
const confirmed = await showModal({
title: "覆盖当前配置",
content: "导入会替换当前完整配置,包括模型、API Key、服务地址和其他设置。建议先导出备份,是否继续?",
confirmText: "确认导入",
cancelText: "取消",
showCancel: true,
});
if (!confirmed) {
return;
}
busy.value = true;
const imported = await importUserConfigFromFile(path);
message(`配置导入成功,共导入 ${imported.modelAdapters.length} 个模型`);
} catch (error) {
showActionError("导入失败", toUserError(error));
} finally {
busy.value = false;
}
}
return {
configTransferBusy: busy,
handleExportConfig: exportConfig,
handleImportConfig: importConfig,
};
}
File diff suppressed because it is too large Load Diff
+2 -34
View File
@@ -1,5 +1,4 @@
{
"01aaebc96af187a9": "Import failed",
"02216368edc68816": "No release notes",
"03b11112dc970014": "Base URL",
"04f632dd4f034d5e": "{0} context window must be a positive integer",
@@ -20,10 +19,8 @@
"1afed6a81a2512d2": "Select model",
"1baddde657dd2720": "Current outbound requests use system proxy",
"1bc77f5ab979f4c1": "Add Model Settings",
"1c631615c1d85c9e": "Log in to Cursor",
"1e238093b79b3165": "Uses 65536 by default when left blank",
"21296ab18ad9af25": "Extra Params JSON",
"23353f7d54f291af": "Another configuration operation is in progress. Please try again later.",
"24343a2096988d42": "Failed to open",
"253d4a3428c648fb": "Cache write tokens: {0}",
"26a3855aed1d8d17": "Service not running",
@@ -39,20 +36,15 @@
"3463c5585c246df9": "Total: {0}",
"3468b57e3edbc599": "Aggregated from turn summaries scanned from the history.",
"35076178fe79a210": "Configuration changed. Please test again.",
"361ea3efbffac8ba": "Continue Export",
"37d23612f78a2e63": "Restart Now to Update",
"392bd7c9d0ffcad8": "Overwrite Current Configuration",
"392d0dceb45998d3": "Extreme",
"393df9bb13ea4900": "Hit",
"3ab8cc15939f3b5c": "Log out",
"3af7e5489e61ea51": "Refreshing",
"3c2a9f9901109e75": "{0} type only supports OpenAI or Anthropic",
"3d13868593ae4eeb": "Interface Language",
"3d52574ce1500561": "Not connected",
"3ea83f9f55062582": "Release date: {0}",
"3edda85621fd03b2": "model adapters",
"3fd47edce45b3603": "Close",
"429ea3af443dd80b": "Export Configuration",
"42aa8e01e98c0d8c": "Total Duration",
"438bb77c1ecdfab0": "Cache Write: {0} × ${1}/1M = {2}",
"46056425a48ef05b": "Currently displayed using the reuse rate definition",
@@ -60,18 +52,14 @@
"472642d58d3d5a6d": "Formula: {0}",
"4923eeb7bd75cccd": "{0} model ID cannot be empty",
"497c85690c4cc0fc": "No data",
"4a8d6841b4023edf": "Confirm Import",
"4b5e0ae1288a9695": "No matching models",
"4c0a929bb86ce912": "Current: {0}",
"4d2b6e53be6002e5": "Cache Statistics Strategy: {0} ({1})",
"4d8c1c5b42830791": "Unknown",
"4e30d7c9ed2b0eee": "Not set",
"4f0982ba1d37e51b": "Current outbound requests use environment variable proxy",
"4ff5e6073183ced5": "YAML Configuration Files",
"5205125c0e91d346": "Maximum tokens an Anthropic model may generate in a single response. Leave blank to use the default.",
"54e6745ff43c9c74": "Sorting failed",
"56627c94a9decee6": "Max Output Tokens",
"58c6b0935a7216da": "Failed to open contributor profile",
"593a972852ba0004": "Cursor Assistant | Permanently Free | Custom API",
"59a2195a01a8b35b": "{0} must be a valid JSON object",
"5aa8f5590c940829": "Non-cache Input: {0}",
@@ -90,12 +78,10 @@
"66af574b8948fe83": "{0} API key cannot be empty",
"6744b4c6a9aa0038": "Disabled",
"675109292da4eb36": "Not tested yet",
"688102a402ba015a": "Waiting for login...",
"6a7b96f399e58138": "e.g. sk-xxxxxx",
"6aa8f49cc992dfd7": "Test",
"6ae23d6d7cb18592": "Service error",
"6d7872164df9b36e": "Normal Input: {0} × ${1}/1M = {2}",
"6e4445df69575d79": "Export failed",
"6e584e3d5ce64aa0": "Save Settings",
"6ec87609a8769425": "Default {0} / Count as creation {1}",
"72f6c3525c0192a7": "When enabled, cache creation is included in the denominator",
@@ -107,7 +93,6 @@
"77c9e582e85583af": "Test failed",
"7a26bf794e9fb6bf": "Used only for display in the UI, so you can distinguish different models.",
"7b6187c41e88b70c": "Testing...",
"7cd744da6bc8c0f1": "Import cursor-byok Configuration",
"7df7641e5e741346": "Cache Read / (Cache Read + Cache Creation + Non-cache Input)",
"7e9e334aeb0bdc07": "Service operation failed",
"7f68ebad19ba6bcd": "Check for Updates",
@@ -116,18 +101,14 @@
"8139cb3dd11f5a67": "When enabled, the JSON object will override the final request headers. Duplicate headers are determined by this field, and values must be strings.",
"8151e8704a7ca89e": "No matches",
"83913e71fcf7ff60": "Refresh successful",
"83be9cac28873059": "Cursor Control Plane Account",
"83fcfb4c1f2c1641": "Fetch Models",
"8672864e90417138": "Max",
"86df7ec743047234": "Service running",
"877b27aa9c9c2691": "Configuration exported to {0}",
"899add6275682210": "Uses 200000 by default when left blank",
"8a4ef3e48e4e8a5a": "Enabled",
"8b8428f714611458": "Only select a reasoning effort when the model supports reasoning_effort. When set to Not set, the request omits this parameter. Higher values are usually more stable, but may also be slower.",
"8c1935935600e336": "Model Test",
"8cbcf741e727dbf7": "Model Settings",
"8d1de152be6360ce": "Valid ratio: {0}",
"8d828cb97168d7cd": "Configuration imported successfully. Imported {0} models.",
"8f6f8d979c981ced": "Copied",
"8f8baf5d18dd0492": "When enabled, the JSON object will override the OpenAI request body. Duplicate fields are determined by this field. OpenAI service_tier supports auto, default, flex, scale, priority; priority can be used for high-priority/Fast scenarios.",
"8faa670b512b6b9b": "Open Model Settings",
@@ -166,7 +147,6 @@
"aed55419ce62f08e": "Switching...",
"b10041a13f5c55b1": "Model Output: {0} × ${1}/1M = {2}",
"b1c27820fec23edb": "High",
"b48b6518293a604e": "Import Configuration",
"b5409d4049286061": "Custom Path (Please enter the full request URL)",
"b571037dc396a00c": "Total request tokens include both prompt and model output.",
"b765005f69fa971f": "e.g. gpt-4.1",
@@ -179,16 +159,12 @@
"bb074b86a98f6911": "Context Window",
"bc87a4121a0873b3": "Refresh Stats",
"bd4464ea88d3f24a": "Total turns: {0}",
"bd4d7a3c6e5a1ac8": "{0} reasoning effort only supports Not set, low, medium, high, xhigh, and max",
"bddd504af0c92fd0": "System PAC/automatic proxy detected; current version is handled as a direct connection",
"bef280f9eb392495": "Conversation Turns",
"c228558cf257fc49": "Delete failed",
"c3d46b387eeadb23": "This only logs the Cursor account out of cursor-byok; it does not log out of the Cursor client. Continue?",
"c5af02060847d167": "Thinking effort for Anthropic adaptive thinking. Requests will consistently use the new thinking.type=adaptive.",
"c6868592796ac2b2": "No {0} models have been configured yet.",
"c69f5bce63b9f14c": "Settings Folder",
"c733ec1c8af53c9a": "Export Full Configuration",
"c8a52b66651d294c": "Failed to log out",
"c8c14507b2d37395": "Reasoning Effort",
"c98e118e0a43f078": "Model",
"c9dd59beefd7144f": "Cache Read / (Cache Read + Non-cache Input)",
@@ -196,8 +172,6 @@
"ca1d1059408b3837": "Invalid turns: {0}",
"cc5049729a2c10f1": "Test failed. Check the raw details.",
"cd7ca5fb221e1c53": "{0} cannot be empty",
"cf14b829486f866e": "Importing will replace the entire current configuration, including models, API keys, service addresses, and other settings. We recommend exporting a backup first. Continue?",
"cfa6c803eb3fc713": "Waiting for browser login",
"d0325067fed88e5a": "Cache hit rate {0}",
"d20ab96566d33f25": "{0} display name cannot be empty",
"d2243e1d44b2a94e": "Edit Model Settings",
@@ -205,27 +179,20 @@
"d373809ab86ba93b": "Copy",
"d3b1da3088ddd334": "Model test failed",
"d53d32f1a1211371": "Custom Headers JSON",
"d6ce4f0f88178144": "Used only for Plugins, Skills, and MCP; does not change the account in the Cursor client",
"d7889896c5b7732a": "Anthropic Extra Params JSON",
"d7da2aabd35772ec": "e.g. 200000 (leave blank to use the default)",
"d95e5cb6bdcee553": "Include Cache Creation",
"da590a8fe3ce4de0": "Please select",
"daede9881787abe7": "Notes",
"dbb4b5be9b5723dc": "{0} reasoning effort only supports low, medium, high, xhigh, and max",
"dbee6e7139243362": "{0} base URL cannot be empty",
"dc82c5e8fb2ab777": "Version: v{0}",
"de8184da1ef88d03": "Configured",
"df5e382e16994c53": "Export cursor-byok Configuration",
"e01c5dae36cf8c35": "When enabled, the JSON object will override the OpenAI request body. Duplicate fields are determined by this field. OpenAI service_tier supports auto, default, flex, scale, priority.",
"e09a2127af23b1d6": "The full configuration cannot be imported while the service is running. Stop the service first.",
"e14c41ef2b7253c9": "Total request tokens: {0}",
"e406825e0a72d2c2": "Local Settings",
"e4343921c928a856": "Login failed",
"e4c0daa3c4bea691": "Thanks to @aike0210 for contributing the Cursor control-plane account feature.",
"e53580f8031f13c0": "Complete login in the browser, then return to Cursor and reopen the plugin marketplace",
"e552c2accdbf5178": "Add Model",
"e6faccfddce722e8": "Cache read tokens: {0}",
"e899a318b5cf571e": "The exported file contains all model settings, API keys, and custom headers. Store it securely to prevent disclosure. On Windows, access depends on the security settings of the destination folder.",
"e8a0a6053998ebfa": "Logged in",
"eaffd48cd2ea9f1a": "e.g. https://api.anthropic.com",
"eb1be07f2ca6e506": "Estimated based on Claude Opus 4.7 pricing.",
"ec3b17a75db49e24": "{0} t/s | First token {1}",
@@ -234,6 +201,7 @@
"f0b6a23368dd47cc": "Enter a model ID directly, or select one from the list returned by the server.",
"f1aa7326f38b4c09": "Drag to reorder",
"f1e0fc261d42fe29": "Notes shown when hovering over the model list.",
"f363622480699c52": "Reasoning effort only applies to some models that support reasoning_effort. Not all models do. Higher values are usually more stable, but may also be slower.",
"f3a76d896853c1df": "Miss",
"f3fae6cccb9004b1": "Custom header name cannot be empty",
"f474a4108aba4c4c": "Stop Service",
+2 -34
View File
@@ -1,5 +1,4 @@
{
"01aaebc96af187a9": "インポートに失敗しました",
"02216368edc68816": "更新内容はありません",
"03b11112dc970014": "ベース URL",
"04f632dd4f034d5e": "{0} のコンテキストウィンドウは正の整数である必要があります",
@@ -20,10 +19,8 @@
"1afed6a81a2512d2": "モデルを選択",
"1baddde657dd2720": "現在のアウトバウンドリクエストはシステムプロキシを使用しています",
"1bc77f5ab979f4c1": "モデル設定を追加",
"1c631615c1d85c9e": "Cursor にログイン",
"1e238093b79b3165": "空欄で 65536",
"21296ab18ad9af25": "追加パラメータ JSON",
"23353f7d54f291af": "別の設定操作を実行中です。しばらくしてからもう一度お試しください。",
"24343a2096988d42": "開けませんでした",
"253d4a3428c648fb": "キャッシュ書き込み Token: {0}",
"26a3855aed1d8d17": "サービスは起動していません",
@@ -39,20 +36,15 @@
"3463c5585c246df9": "合計:{0}",
"3468b57e3edbc599": "履歴からスキャンした各ターンの summary を集計しています。",
"35076178fe79a210": "設定が変更されました。再テストしてください",
"361ea3efbffac8ba": "エクスポートを続行",
"37d23612f78a2e63": "今すぐ再起動して更新",
"392bd7c9d0ffcad8": "現在の設定を上書き",
"392d0dceb45998d3": "最高",
"393df9bb13ea4900": "ヒット",
"3ab8cc15939f3b5c": "ログアウト",
"3af7e5489e61ea51": "更新中",
"3c2a9f9901109e75": "{0} のタイプは OpenAI または Anthropic のみサポートします",
"3d13868593ae4eeb": "表示言語",
"3d52574ce1500561": "未接続",
"3ea83f9f55062582": "公開日時: {0}",
"3edda85621fd03b2": "件のモデルアダプター",
"3fd47edce45b3603": "閉じる",
"429ea3af443dd80b": "設定をエクスポート",
"42aa8e01e98c0d8c": "総所要時間",
"438bb77c1ecdfab0": "キャッシュ書き込み:{0} × ${1}/1M = {2}",
"46056425a48ef05b": "現在、再利用率の定義に従って表示されています",
@@ -60,18 +52,14 @@
"472642d58d3d5a6d": "数式:{0}",
"4923eeb7bd75cccd": "{0} のモデル ID は必須です",
"497c85690c4cc0fc": "データなし",
"4a8d6841b4023edf": "インポートを確認",
"4b5e0ae1288a9695": "一致するモデルがありません",
"4c0a929bb86ce912": "現在:{0}",
"4d2b6e53be6002e5": "キャッシュ統計ポリシー:{0}{1})",
"4d8c1c5b42830791": "不明",
"4e30d7c9ed2b0eee": "設定しない",
"4f0982ba1d37e51b": "現在のアウトバウンドリクエストは環境変数プロキシを使用しています",
"4ff5e6073183ced5": "YAML 設定ファイル",
"5205125c0e91d346": "Anthropic モデルが1回の応答で生成できる最大 Token 数。空欄の場合はデフォルト値を使用します。",
"54e6745ff43c9c74": "並べ替えに失敗しました",
"56627c94a9decee6": "最大出力 Token",
"58c6b0935a7216da": "コントリビューターのプロフィールを開けませんでした",
"593a972852ba0004": "Cursor アシスタント | 永久無料 | カスタム API",
"59a2195a01a8b35b": "{0}は有効なJSONオブジェクトである必要があります",
"5aa8f5590c940829": "非キャッシュ入力:{0}",
@@ -90,12 +78,10 @@
"66af574b8948fe83": "{0} の API キーは必須です",
"6744b4c6a9aa0038": "無効化",
"675109292da4eb36": "まだテストしていません",
"688102a402ba015a": "ログインを待っています...",
"6a7b96f399e58138": "例: sk-xxxxxx",
"6aa8f49cc992dfd7": "テスト",
"6ae23d6d7cb18592": "サービスエラー",
"6d7872164df9b36e": "通常入力:{0} × ${1}/1M = {2}",
"6e4445df69575d79": "エクスポートに失敗しました",
"6e584e3d5ce64aa0": "設定を保存",
"6ec87609a8769425": "デフォルト {0} / 作成カウント {1}",
"72f6c3525c0192a7": "有効にすると、キャッシュ作成が分母に含まれます",
@@ -107,7 +93,6 @@
"77c9e582e85583af": "テスト失敗",
"7a26bf794e9fb6bf": "UIでモデルを区別するための表示専用です。",
"7b6187c41e88b70c": "テスト中...",
"7cd744da6bc8c0f1": "cursor-byok 設定をインポート",
"7df7641e5e741346": "キャッシュ読み取り / (キャッシュ読み取り + キャッシュ作成 + 非キャッシュ入力)",
"7e9e334aeb0bdc07": "サービス操作に失敗しました",
"7f68ebad19ba6bcd": "アップデートを確認",
@@ -116,18 +101,14 @@
"8139cb3dd11f5a67": "有効にすると、JSONオブジェクトが最終的なリクエストヘッダーを上書きします。同名のヘッダーはこの設定が優先され、値は文字列である必要があります。",
"8151e8704a7ca89e": "一致する項目がありません",
"83913e71fcf7ff60": "更新しました",
"83be9cac28873059": "Cursor コントロールプレーンアカウント",
"83fcfb4c1f2c1641": "モデルを取得",
"8672864e90417138": "最大",
"86df7ec743047234": "サービス稼働中",
"877b27aa9c9c2691": "設定を {0} にエクスポートしました",
"899add6275682210": "空欄で 200000",
"8a4ef3e48e4e8a5a": "有効",
"8b8428f714611458": "モデルが reasoning_effort に対応している場合のみ推論強度を選択してください。「設定しない」を選ぶと、リクエストにこのパラメータは含まれません。値が高いほど安定しやすい反面、遅くなることがあります。",
"8c1935935600e336": "モデルテスト",
"8cbcf741e727dbf7": "モデル設定",
"8d1de152be6360ce": "有効率: {0}",
"8d828cb97168d7cd": "設定のインポートが完了しました。{0} 件のモデルをインポートしました",
"8f6f8d979c981ced": "コピーしました",
"8f8baf5d18dd0492": "有効にすると、JSONオブジェクトがOpenAIのリクエストボディを上書きします。同名のフィールドはこの設定が優先されます。OpenAIのservice_tierはauto、default、flex、scale、priorityをサポートしており、priorityは高優先度/Fastのシナリオで使用できます。",
"8faa670b512b6b9b": "モデル設定を開く",
@@ -166,7 +147,6 @@
"aed55419ce62f08e": "切替中...",
"b10041a13f5c55b1": "モデル出力:{0} × ${1}/1M = {2}",
"b1c27820fec23edb": "高",
"b48b6518293a604e": "設定をインポート",
"b5409d4049286061": "カスタムパス(完全なリクエストURLを入力してください)",
"b571037dc396a00c": "総リクエスト Token には Prompt とモデル出力の両方が含まれます。",
"b765005f69fa971f": "例: gpt-4.1",
@@ -179,16 +159,12 @@
"bb074b86a98f6911": "コンテキストウィンドウ",
"bc87a4121a0873b3": "統計を更新",
"bd4464ea88d3f24a": "総ターン: {0}",
"bd4d7a3c6e5a1ac8": "{0} の推論強度は「設定しない」、low、medium、high、xhigh、max のみサポートします",
"bddd504af0c92fd0": "システムのPAC/自動プロキシが検出されました。現在のバージョンは直接接続として処理されます",
"bef280f9eb392495": "会話ターン",
"c228558cf257fc49": "削除に失敗しました",
"c3d46b387eeadb23": "cursor-byok 内の Cursor アカウントからのみログアウトします。Cursor クライアントからはログアウトしません。続行しますか?",
"c5af02060847d167": "Anthropic adaptive thinkingの思考強度。リクエストは一貫して新しいthinking.type=adaptiveを使用します。",
"c6868592796ac2b2": "まだ {0} モデルが設定されていません。",
"c69f5bce63b9f14c": "設定フォルダー",
"c733ec1c8af53c9a": "完全な設定をエクスポート",
"c8a52b66651d294c": "ログアウトに失敗しました",
"c8c14507b2d37395": "推論強度",
"c98e118e0a43f078": "モデル",
"c9dd59beefd7144f": "キャッシュ読み取り / (キャッシュ読み取り + 非キャッシュ入力)",
@@ -196,8 +172,6 @@
"ca1d1059408b3837": "異常ターン: {0}",
"cc5049729a2c10f1": "テストに失敗しました。元の詳細情報を確認してください。",
"cd7ca5fb221e1c53": "{0}は空にできません",
"cf14b829486f866e": "インポートすると、モデル、API Key、サービスアドレス、その他の設定を含む現在の設定全体が置き換えられます。先にバックアップをエクスポートすることを推奨します。続行しますか?",
"cfa6c803eb3fc713": "ブラウザでのログインを待っています",
"d0325067fed88e5a": "キャッシュヒット率 {0}",
"d20ab96566d33f25": "{0} の表示名は必須です",
"d2243e1d44b2a94e": "モデル設定を編集",
@@ -205,27 +179,20 @@
"d373809ab86ba93b": "コピー",
"d3b1da3088ddd334": "モデルテストに失敗しました",
"d53d32f1a1211371": "カスタムヘッダー JSON",
"d6ce4f0f88178144": "プラグイン、Skills、MCP 専用です。Cursor クライアントの現在のアカウントは変更しません",
"d7889896c5b7732a": "Anthropic 追加パラメータ JSON",
"d7da2aabd35772ec": "例: 200000(空欄でデフォルト値)",
"d95e5cb6bdcee553": "キャッシュ作成を含める",
"da590a8fe3ce4de0": "選択してください",
"daede9881787abe7": "メモ",
"dbb4b5be9b5723dc": "{0} の推論強度は low、medium、high、xhigh、max のみサポートします",
"dbee6e7139243362": "{0} のベース URL は必須です",
"dc82c5e8fb2ab777": "バージョン: v{0}",
"de8184da1ef88d03": "設定済み",
"df5e382e16994c53": "cursor-byok 設定をエクスポート",
"e01c5dae36cf8c35": "有効にすると、JSONオブジェクトがOpenAIのリクエストボディを上書きします。同名のフィールドはこの設定が優先されます。OpenAIのservice_tierはauto、default、flex、scale、priorityをサポートしています。",
"e09a2127af23b1d6": "サービスの実行中は設定全体をインポートできません。先にサービスを停止してください。",
"e14c41ef2b7253c9": "総リクエスト Token: {0}",
"e406825e0a72d2c2": "ローカル設定",
"e4343921c928a856": "ログインに失敗しました",
"e4c0daa3c4bea691": "Cursor コントロールプレーンアカウント機能への @aike0210 の貢献に感謝します。",
"e53580f8031f13c0": "ブラウザでログインを完了し、Cursor に戻ってプラグインマーケットを開き直してください",
"e552c2accdbf5178": "モデルを追加",
"e6faccfddce722e8": "キャッシュ読込 Token: {0}",
"e899a318b5cf571e": "エクスポートファイルには、すべてのモデル設定、API キー、カスタムヘッダーが含まれます。漏えいを防ぐため、安全な場所に保存してください。Windows では、アクセス権限は保存先フォルダーのセキュリティ設定に依存します。",
"e8a0a6053998ebfa": "ログイン済み",
"eaffd48cd2ea9f1a": "例: https://api.anthropic.com",
"eb1be07f2ca6e506": "Claude Opus 4.7の価格に基づいて見積もられます。",
"ec3b17a75db49e24": "{0} t/s | 初回 Token {1}",
@@ -234,6 +201,7 @@
"f0b6a23368dd47cc": "モデルIDを直接入力するか、サーバーから返された一覧から選択します。",
"f1aa7326f38b4c09": "ドラッグして並べ替え",
"f1e0fc261d42fe29": "モデル一覧にホバーしたときに表示されるメモです。",
"f363622480699c52": "推論強度は reasoning_effort をサポートする一部のモデルでのみ有効です。すべてのモデルが対応しているわけではありません。値が高いほど安定しやすい反面、遅くなることがあります。",
"f3a76d896853c1df": "ミス",
"f3fae6cccb9004b1": "カスタムヘッダー名は空にできません",
"f474a4108aba4c4c": "サービスを停止",
+2 -34
View File
@@ -1,5 +1,4 @@
{
"01aaebc96af187a9": "Не удалось импортировать",
"02216368edc68816": "Нет примечаний к выпуску",
"03b11112dc970014": "Базовый URL",
"04f632dd4f034d5e": "Размер контекстного окна {0} должен быть положительным целым числом",
@@ -20,10 +19,8 @@
"1afed6a81a2512d2": "Выберите модель",
"1baddde657dd2720": "Исходящие запросы используют системный прокси",
"1bc77f5ab979f4c1": "Добавить настройки модели",
"1c631615c1d85c9e": "Войти в Cursor",
"1e238093b79b3165": "Если оставить пустым, используется 65536",
"21296ab18ad9af25": "Дополнительные параметры JSON",
"23353f7d54f291af": "Уже выполняется другая операция с конфигурацией. Повторите попытку позже.",
"24343a2096988d42": "Не удалось открыть",
"253d4a3428c648fb": "Токены записи в кеш: {0}",
"26a3855aed1d8d17": "Сервис не запущен",
@@ -39,20 +36,15 @@
"3463c5585c246df9": "Итого: {0}",
"3468b57e3edbc599": "Сводка составлена по данным ходов, найденным в истории.",
"35076178fe79a210": "Конфигурация изменена. Выполните проверку снова.",
"361ea3efbffac8ba": "Продолжить экспорт",
"37d23612f78a2e63": "Перезапустить и обновить",
"392bd7c9d0ffcad8": "Заменить текущую конфигурацию",
"392d0dceb45998d3": "Очень высокая",
"393df9bb13ea4900": "Попадание",
"3ab8cc15939f3b5c": "Выйти",
"3af7e5489e61ea51": "Обновление",
"3c2a9f9901109e75": "Тип {0} поддерживает только OpenAI или Anthropic",
"3d13868593ae4eeb": "Язык интерфейса",
"3d52574ce1500561": "Не подключено",
"3ea83f9f55062582": "Дата выпуска: {0}",
"3edda85621fd03b2": "адаптеров моделей",
"3fd47edce45b3603": "Закрыть",
"429ea3af443dd80b": "Экспортировать конфигурацию",
"42aa8e01e98c0d8c": "Общая длительность",
"438bb77c1ecdfab0": "Запись в кеш: {0} × ${1}/1M = {2}",
"46056425a48ef05b": "Сейчас используется расчет по коэффициенту повторного использования",
@@ -60,18 +52,14 @@
"472642d58d3d5a6d": "Формула: {0}",
"4923eeb7bd75cccd": "Идентификатор модели {0} не может быть пустым",
"497c85690c4cc0fc": "Нет данных",
"4a8d6841b4023edf": "Подтвердить импорт",
"4b5e0ae1288a9695": "Подходящих моделей нет",
"4c0a929bb86ce912": "Сейчас: {0}",
"4d2b6e53be6002e5": "Стратегия статистики кеша: {0} ({1})",
"4d8c1c5b42830791": "Неизвестно",
"4e30d7c9ed2b0eee": "Не задано",
"4f0982ba1d37e51b": "Исходящие запросы используют прокси из переменных окружения",
"4ff5e6073183ced5": "Файлы конфигурации YAML",
"5205125c0e91d346": "Максимальное число токенов, которое модель Anthropic может сгенерировать за один ответ. Оставьте поле пустым для значения по умолчанию.",
"54e6745ff43c9c74": "Не удалось изменить порядок",
"56627c94a9decee6": "Макс. выходных токенов",
"58c6b0935a7216da": "Не удалось открыть профиль участника",
"593a972852ba0004": "Cursor Assistant | Всегда бесплатно | Пользовательский API",
"59a2195a01a8b35b": "{0} должен быть допустимым объектом JSON",
"5aa8f5590c940829": "Ввод без кеша: {0}",
@@ -90,12 +78,10 @@
"66af574b8948fe83": "Ключ API {0} не может быть пустым",
"6744b4c6a9aa0038": "Выключено",
"675109292da4eb36": "Еще не проверено",
"688102a402ba015a": "Ожидание входа...",
"6a7b96f399e58138": "например, sk-xxxxxx",
"6aa8f49cc992dfd7": "Проверить",
"6ae23d6d7cb18592": "Ошибка сервиса",
"6d7872164df9b36e": "Обычный ввод: {0} × ${1}/1M = {2}",
"6e4445df69575d79": "Не удалось экспортировать",
"6e584e3d5ce64aa0": "Сохранить настройки",
"6ec87609a8769425": "По умолчанию {0} / С учетом создания {1}",
"72f6c3525c0192a7": "Если включено, создание кеша учитывается в знаменателе",
@@ -107,7 +93,6 @@
"77c9e582e85583af": "Проверка не пройдена",
"7a26bf794e9fb6bf": "Используется только для отображения в интерфейсе, чтобы различать модели.",
"7b6187c41e88b70c": "Проверка...",
"7cd744da6bc8c0f1": "Импорт конфигурации cursor-byok",
"7df7641e5e741346": "Чтение кеша / (Чтение кеша + Создание кеша + Ввод без кеша)",
"7e9e334aeb0bdc07": "Не удалось выполнить операцию с сервисом",
"7f68ebad19ba6bcd": "Проверить обновления",
@@ -116,18 +101,14 @@
"8139cb3dd11f5a67": "Если включено, объект JSON переопределит итоговые заголовки запроса. При совпадении имен используются значения отсюда; все значения должны быть строками.",
"8151e8704a7ca89e": "Совпадений нет",
"83913e71fcf7ff60": "Обновление выполнено",
"83be9cac28873059": "Аккаунт управляющего уровня Cursor",
"83fcfb4c1f2c1641": "Получить модели",
"8672864e90417138": "Максимальная",
"86df7ec743047234": "Сервис запущен",
"877b27aa9c9c2691": "Конфигурация экспортирована в {0}",
"899add6275682210": "Если оставить пустым, используется 200000",
"8a4ef3e48e4e8a5a": "Включено",
"8b8428f714611458": "Выбирайте интенсивность рассуждений только для моделей с поддержкой reasoning_effort. Если выбрать «Не задано», этот параметр не будет добавлен в запрос. Более высокие значения обычно дают более стабильный результат, но могут замедлить ответ.",
"8c1935935600e336": "Проверка модели",
"8cbcf741e727dbf7": "Настройки модели",
"8d1de152be6360ce": "Доля успешных: {0}",
"8d828cb97168d7cd": "Конфигурация успешно импортирована. Импортировано моделей: {0}",
"8f6f8d979c981ced": "Скопировано",
"8f8baf5d18dd0492": "Если включено, объект JSON переопределит тело запроса OpenAI. При совпадении полей используются значения отсюда. OpenAI service_tier поддерживает auto, default, flex, scale и priority; priority можно использовать для сценариев с высоким приоритетом/Fast.",
"8faa670b512b6b9b": "Открыть настройки модели",
@@ -166,7 +147,6 @@
"aed55419ce62f08e": "Переключение...",
"b10041a13f5c55b1": "Вывод модели: {0} × ${1}/1M = {2}",
"b1c27820fec23edb": "Высокая",
"b48b6518293a604e": "Импортировать конфигурацию",
"b5409d4049286061": "Пользовательский путь (введите полный URL запроса)",
"b571037dc396a00c": "Общее число токенов запроса включает Prompt и вывод модели.",
"b765005f69fa971f": "например, gpt-4.1",
@@ -179,16 +159,12 @@
"bb074b86a98f6911": "Контекстное окно",
"bc87a4121a0873b3": "Обновить статистику",
"bd4464ea88d3f24a": "Всего ходов: {0}",
"bd4d7a3c6e5a1ac8": "Интенсивность рассуждений {0} поддерживает только значения «Не задано», low, medium, high, xhigh и max",
"bddd504af0c92fd0": "Обнаружен системный PAC/автоматический прокси; в текущей версии используется прямое подключение",
"bef280f9eb392495": "Ходы диалога",
"c228558cf257fc49": "Не удалось удалить",
"c3d46b387eeadb23": "Будет выполнен выход только из аккаунта Cursor в cursor-byok. В клиенте Cursor вы останетесь в системе. Продолжить?",
"c5af02060847d167": "Интенсивность для адаптивных рассуждений Anthropic. В запросах всегда используется новый режим thinking.type=adaptive.",
"c6868592796ac2b2": "Модели {0} пока не настроены.",
"c69f5bce63b9f14c": "Папка настроек",
"c733ec1c8af53c9a": "Экспортировать всю конфигурацию",
"c8a52b66651d294c": "Не удалось выйти",
"c8c14507b2d37395": "Интенсивность рассуждений",
"c98e118e0a43f078": "Модель",
"c9dd59beefd7144f": "Чтение кеша / (Чтение кеша + Ввод без кеша)",
@@ -196,8 +172,6 @@
"ca1d1059408b3837": "Ошибочных ходов: {0}",
"cc5049729a2c10f1": "Тест не пройден. Проверьте исходные сведения.",
"cd7ca5fb221e1c53": "{0} не может быть пустым",
"cf14b829486f866e": "Импорт заменит всю текущую конфигурацию, включая модели, API-ключи, адреса сервисов и другие настройки. Рекомендуется сначала экспортировать резервную копию. Продолжить?",
"cfa6c803eb3fc713": "Ожидание входа в браузере",
"d0325067fed88e5a": "Доля попаданий в кеш: {0}",
"d20ab96566d33f25": "Отображаемое имя {0} не может быть пустым",
"d2243e1d44b2a94e": "Изменить настройки модели",
@@ -205,27 +179,20 @@
"d373809ab86ba93b": "Копировать",
"d3b1da3088ddd334": "Проверка модели не пройдена",
"d53d32f1a1211371": "Пользовательские заголовки JSON",
"d6ce4f0f88178144": "Используется только для Plugins, Skills и MCP; текущий аккаунт клиента Cursor не изменяется",
"d7889896c5b7732a": "Дополнительные параметры Anthropic JSON",
"d7da2aabd35772ec": "например, 200000 (оставьте пустым для значения по умолчанию)",
"d95e5cb6bdcee553": "Учитывать создание кеша",
"da590a8fe3ce4de0": "Выберите значение",
"daede9881787abe7": "Примечания",
"dbb4b5be9b5723dc": "Интенсивность рассуждений {0} поддерживает только low, medium, high, xhigh и max",
"dbee6e7139243362": "Базовый URL {0} не может быть пустым",
"dc82c5e8fb2ab777": "Версия: v{0}",
"de8184da1ef88d03": "Настроено",
"df5e382e16994c53": "Экспорт конфигурации cursor-byok",
"e01c5dae36cf8c35": "Если включено, объект JSON переопределит тело запроса OpenAI. При совпадении полей используются значения отсюда. OpenAI service_tier поддерживает auto, default, flex, scale и priority.",
"e09a2127af23b1d6": "Нельзя импортировать полную конфигурацию во время работы сервиса. Сначала остановите сервис.",
"e14c41ef2b7253c9": "Всего токенов запроса: {0}",
"e406825e0a72d2c2": "Локальные настройки",
"e4343921c928a856": "Не удалось войти",
"e4c0daa3c4bea691": "Спасибо @aike0210 за вклад в функцию аккаунта панели управления Cursor.",
"e53580f8031f13c0": "Завершите вход в браузере, затем вернитесь в Cursor и снова откройте магазин плагинов",
"e552c2accdbf5178": "Добавить модель",
"e6faccfddce722e8": "Токены чтения из кеша: {0}",
"e899a318b5cf571e": "Экспортированный файл содержит все настройки моделей, ключи API и пользовательские заголовки. Храните его в безопасном месте. В Windows доступ зависит от параметров безопасности папки назначения.",
"e8a0a6053998ebfa": "Выполнен вход",
"eaffd48cd2ea9f1a": "например, https://api.anthropic.com",
"eb1be07f2ca6e506": "Расчет основан на тарифах Claude Opus 4.7.",
"ec3b17a75db49e24": "{0} т/с | Первый токен {1}",
@@ -234,6 +201,7 @@
"f0b6a23368dd47cc": "Введите идентификатор модели вручную или выберите его из списка, полученного от сервера.",
"f1aa7326f38b4c09": "Перетащите, чтобы изменить порядок",
"f1e0fc261d42fe29": "Примечание, отображаемое при наведении на модель в списке.",
"f363622480699c52": "Интенсивность рассуждений применяется только к моделям с поддержкой reasoning_effort. Чем выше значение, тем обычно стабильнее результат, но ответ может формироваться медленнее.",
"f3a76d896853c1df": "Промах",
"f3fae6cccb9004b1": "Имя пользовательского заголовка не может быть пустым",
"f474a4108aba4c4c": "Остановить сервис",
+2 -34
View File
@@ -1,5 +1,4 @@
{
"01aaebc96af187a9": "导入失败",
"02216368edc68816": "无更新说明",
"03b11112dc970014": "接口地址",
"04f632dd4f034d5e": "{0} 的上下文窗口必须为正整数",
@@ -20,10 +19,8 @@
"1afed6a81a2512d2": "选择模型",
"1baddde657dd2720": "当前出站请求使用系统代理",
"1bc77f5ab979f4c1": "新增模型配置",
"1c631615c1d85c9e": "登录 Cursor",
"1e238093b79b3165": "留空时默认 65536",
"21296ab18ad9af25": "额外参数 JSON",
"23353f7d54f291af": "已有配置操作正在进行,请稍后再试",
"24343a2096988d42": "打开失败",
"253d4a3428c648fb": "缓存写入:{0}",
"26a3855aed1d8d17": "服务未启动",
@@ -39,20 +36,15 @@
"3463c5585c246df9": "合计:{0}",
"3468b57e3edbc599": "按历史记录里扫描到的回合 summary 汇总。",
"35076178fe79a210": "配置已变更,请重新测试",
"361ea3efbffac8ba": "继续导出",
"37d23612f78a2e63": "立即重启更新",
"392bd7c9d0ffcad8": "覆盖当前配置",
"392d0dceb45998d3": "极高",
"393df9bb13ea4900": "命中",
"3ab8cc15939f3b5c": "退出登录",
"3af7e5489e61ea51": "刷新中",
"3c2a9f9901109e75": "{0} 的类型仅支持 OpenAI 或 Anthropic",
"3d13868593ae4eeb": "界面语言",
"3d52574ce1500561": "未连接",
"3ea83f9f55062582": "发布时间:{0}",
"3edda85621fd03b2": "个模型适配器",
"3fd47edce45b3603": "关闭",
"429ea3af443dd80b": "导出配置",
"42aa8e01e98c0d8c": "总耗时",
"438bb77c1ecdfab0": "缓存写入:{0} × ${1}/1M = {2}",
"46056425a48ef05b": "当前按复用率口径显示",
@@ -60,18 +52,14 @@
"472642d58d3d5a6d": "公式:{0}",
"4923eeb7bd75cccd": "{0} 的模型标识不能为空",
"497c85690c4cc0fc": "暂无数据",
"4a8d6841b4023edf": "确认导入",
"4b5e0ae1288a9695": "没有匹配的模型",
"4c0a929bb86ce912": "当前:{0}",
"4d2b6e53be6002e5": "缓存统计策略:{0}{1}",
"4d8c1c5b42830791": "未知",
"4e30d7c9ed2b0eee": "不设置",
"4f0982ba1d37e51b": "当前出站请求使用环境变量代理",
"4ff5e6073183ced5": "YAML 配置文件",
"5205125c0e91d346": "Anthropic 模型单次回复允许生成的最大 Token 数。留空时使用默认值。",
"54e6745ff43c9c74": "排序失败",
"56627c94a9decee6": "最大输出 Token",
"58c6b0935a7216da": "打开贡献者主页失败",
"593a972852ba0004": "Cursor助手|永久免费|自定义API",
"59a2195a01a8b35b": "{0}必须是合法 JSON 对象",
"5aa8f5590c940829": "非缓存输入:{0}",
@@ -90,12 +78,10 @@
"66af574b8948fe83": "{0} 的访问密钥不能为空",
"6744b4c6a9aa0038": "已关闭",
"675109292da4eb36": "尚未测试",
"688102a402ba015a": "等待登录...",
"6a7b96f399e58138": "例如:sk-xxxxxx",
"6aa8f49cc992dfd7": "测试",
"6ae23d6d7cb18592": "服务错误",
"6d7872164df9b36e": "普通输入:{0} × ${1}/1M = {2}",
"6e4445df69575d79": "导出失败",
"6e584e3d5ce64aa0": "保存配置",
"6ec87609a8769425": "默认 {0} / 计入创建 {1}",
"72f6c3525c0192a7": "开启后把缓存创建纳入分母",
@@ -107,7 +93,6 @@
"77c9e582e85583af": "测试失败",
"7a26bf794e9fb6bf": "仅用于界面展示,便于你区分不同模型。",
"7b6187c41e88b70c": "测试中...",
"7cd744da6bc8c0f1": "导入 cursor-byok 配置",
"7df7641e5e741346": "缓存读取 /(缓存读取 + 缓存创建 + 非缓存输入)",
"7e9e334aeb0bdc07": "服务操作失败",
"7f68ebad19ba6bcd": "检查更新",
@@ -116,18 +101,14 @@
"8139cb3dd11f5a67": "开启后会把 JSON 对象覆盖到最终请求头。同名请求头以这里为准,值必须是字符串。",
"8151e8704a7ca89e": "没有匹配项",
"83913e71fcf7ff60": "刷新成功",
"83be9cac28873059": "Cursor 控制面账号",
"83fcfb4c1f2c1641": "获取模型",
"8672864e90417138": "最高",
"86df7ec743047234": "服务运行中",
"877b27aa9c9c2691": "配置已导出到 {0}",
"899add6275682210": "留空时默认 200000",
"8a4ef3e48e4e8a5a": "已开启",
"8b8428f714611458": "仅当模型支持 reasoning_effort 时才选择推理强度;选择“不设置”后,请求不会携带该参数。越高通常越稳,但也可能更慢。",
"8c1935935600e336": "模型测试",
"8cbcf741e727dbf7": "模型配置",
"8d1de152be6360ce": "有效占比:{0}",
"8d828cb97168d7cd": "配置导入成功,共导入 {0} 个模型",
"8f6f8d979c981ced": "已复制",
"8f8baf5d18dd0492": "开启后会把 JSON 对象覆盖到 OpenAI 请求体。同名字段以这里为准。OpenAI service_tier 支持 auto、default、flex、scale、prioritypriority 可用于高优先级/Fast 类场景。",
"8faa670b512b6b9b": "打开模型配置",
@@ -166,7 +147,6 @@
"aed55419ce62f08e": "切换中...",
"b10041a13f5c55b1": "模型输出:{0} × ${1}/1M = {2}",
"b1c27820fec23edb": "高",
"b48b6518293a604e": "导入配置",
"b5409d4049286061": "自定义路径(请输入完整请求地址)",
"b571037dc396a00c": "总请求 Token 包含 Prompt 和模型输出。",
"b765005f69fa971f": "例如:gpt-4.1",
@@ -179,16 +159,12 @@
"bb074b86a98f6911": "上下文窗口",
"bc87a4121a0873b3": "刷新统计",
"bd4464ea88d3f24a": "总轮次:{0}",
"bd4d7a3c6e5a1ac8": "{0} 的推理强度仅支持不设置、low、medium、high、xhigh、max",
"bddd504af0c92fd0": "检测到系统 PAC/自动代理,当前版本按直连处理",
"bef280f9eb392495": "对话轮次",
"c228558cf257fc49": "删除失败",
"c3d46b387eeadb23": "只会退出 cursor-byok 中的 Cursor 账号,不会退出 Cursor 客户端。是否继续?",
"c5af02060847d167": "Anthropic adaptive thinking 的思考强度。请求会固定使用新版 thinking.type=adaptive。",
"c6868592796ac2b2": "当前还没有配置任何 {0} 模型。",
"c69f5bce63b9f14c": "设置文件夹",
"c733ec1c8af53c9a": "导出完整配置",
"c8a52b66651d294c": "退出登录失败",
"c8c14507b2d37395": "推理强度",
"c98e118e0a43f078": "模型",
"c9dd59beefd7144f": "缓存读取 /(缓存读取 + 非缓存输入)",
@@ -196,8 +172,6 @@
"ca1d1059408b3837": "异常轮次:{0}",
"cc5049729a2c10f1": "测试失败,请查看原始信息",
"cd7ca5fb221e1c53": "{0}不能为空",
"cf14b829486f866e": "导入会替换当前完整配置,包括模型、API Key、服务地址和其他设置。建议先导出备份,是否继续?",
"cfa6c803eb3fc713": "等待浏览器登录",
"d0325067fed88e5a": "缓存命中率 {0}",
"d20ab96566d33f25": "{0} 的显示名称不能为空",
"d2243e1d44b2a94e": "编辑模型配置",
@@ -205,27 +179,20 @@
"d373809ab86ba93b": "拷贝",
"d3b1da3088ddd334": "模型测试失败",
"d53d32f1a1211371": "自定义请求头 JSON",
"d6ce4f0f88178144": "独立用于插件、Skills 和 MCP;不会改变 Cursor 客户端当前账号",
"d7889896c5b7732a": "Anthropic 额外参数 JSON",
"d7da2aabd35772ec": "例如:200000(留空用默认值)",
"d95e5cb6bdcee553": "计入缓存创建",
"da590a8fe3ce4de0": "请选择",
"daede9881787abe7": "备注",
"dbb4b5be9b5723dc": "{0} 的推理强度仅支持 low、medium、high、xhigh、max",
"dbee6e7139243362": "{0} 的接口地址不能为空",
"dc82c5e8fb2ab777": "版本:v{0}",
"de8184da1ef88d03": "已配置",
"df5e382e16994c53": "导出 cursor-byok 配置",
"e01c5dae36cf8c35": "开启后会把 JSON 对象覆盖到 OpenAI 请求体。同名字段以这里为准。OpenAI service_tier 支持 auto、default、flex、scale、priority。",
"e09a2127af23b1d6": "服务运行中不能导入完整配置,请先停止服务",
"e14c41ef2b7253c9": "总请求:{0}",
"e406825e0a72d2c2": "本地配置",
"e4343921c928a856": "登录失败",
"e4c0daa3c4bea691": "感谢 @aike0210 对 Cursor 控制面账号功能的贡献。",
"e53580f8031f13c0": "请在浏览器完成登录,完成后返回 Cursor 重新打开插件市场",
"e552c2accdbf5178": "新增模型",
"e6faccfddce722e8": "缓存读取:{0}",
"e899a318b5cf571e": "导出文件会包含全部模型配置、API Key 和自定义请求头。请将文件保存在安全位置,避免泄露;Windows 上的访问权限取决于目标文件夹的安全设置。",
"e8a0a6053998ebfa": "已经登录",
"eaffd48cd2ea9f1a": "例如:https://api.anthropic.com",
"eb1be07f2ca6e506": "按 Claude Opus 4.7 价格估算。",
"ec3b17a75db49e24": "{0} t/s | 首字 {1}",
@@ -234,6 +201,7 @@
"f0b6a23368dd47cc": "可以直接输入模型标识,或从服务端返回的列表中选择。",
"f1aa7326f38b4c09": "拖拽排序",
"f1e0fc261d42fe29": "模型列表 hover 时显示的备注说明。",
"f363622480699c52": "推理强度仅对部分支持 reasoning_effort 的模型生效,并不是所有模型都支持。越高通常越稳,但也可能更慢。",
"f3a76d896853c1df": "未命中",
"f3fae6cccb9004b1": "自定义请求头名称不能为空",
"f474a4108aba4c4c": "关闭服务",
-37
View File
@@ -1,10 +1,7 @@
import {
DisconnectCursorAccount,
GetCursorAccountStatus,
GetState,
LoadUserConfig,
SaveUserConfig,
StartCursorAccountLogin,
StartProxy,
StopProxy,
} from "@bindings/cursor/internal/bridge/proxyservice.js";
@@ -63,40 +60,6 @@ export function saveUserConfig(payload) {
return withApiLogging("SaveUserConfig", payload, () => SaveUserConfig(payload));
}
export function exportUserConfig(path) {
return withApiLogging("ExportUserConfig", { path }, () =>
Call.ByName(`${PROXY_SERVICE_NAME}.ExportUserConfig`, path),
);
}
export function importUserConfig(path) {
return Call.ByName(`${PROXY_SERVICE_NAME}.ImportUserConfig`, path).then(
(result) => {
console.log(`${API_LOG_PREFIX} ImportUserConfig response`, {
path,
modelCount: Array.isArray(result?.modelAdapters) ? result.modelAdapters.length : 0,
});
return result;
},
(error) => {
logError("ImportUserConfig", { path }, error);
throw error;
},
);
}
export function getCursorAccountStatus() {
return withApiLogging("GetCursorAccountStatus", undefined, () => GetCursorAccountStatus());
}
export function startCursorAccountLogin() {
return withApiLogging("StartCursorAccountLogin", undefined, () => StartCursorAccountLogin());
}
export function disconnectCursorAccount() {
return withApiLogging("DisconnectCursorAccount", undefined, () => DisconnectCursorAccount());
}
export function getProxyState() {
return withApiLogging("GetState", undefined, () => GetState());
}
+8 -39
View File
@@ -3,13 +3,11 @@ import { Events } from "@wailsio/runtime";
import dayjs from "dayjs";
import {
checkForUpdates,
exportUserConfig as exportUserConfigFile,
getAppVersion,
getHomeMetricsSummary,
getModelAdapterTestResults,
installReadyUpdate,
getProxyState,
importUserConfig as importUserConfigFile,
openConfigWindow as openConfig,
loadUserConfig,
openLogsDirectory,
@@ -20,14 +18,11 @@ import {
testModelAdapter,
fetchModelAdapterModels,
} from "@/services/clientApi";
import {
normalizeReasoningEffort,
SUPPORTED_REASONING_EFFORTS,
} from "@/state/modelAdapterReasoning";
const APP_STATE_STORAGE_KEY = "cursor-client:runtime-state:v2";
const GENERIC_SERVICE_ERROR = "服务错误";
const SUPPORTED_MODEL_ADAPTER_TYPES = new Set(["openai", "anthropic"]);
const SUPPORTED_REASONING_EFFORTS = new Set(["low", "medium", "high", "xhigh", "max"]);
const SUPPORTED_ANTHROPIC_THINKING_EFFORTS = new Set(["low", "medium", "high", "xhigh", "max"]);
export const ANTHROPIC_THINKING_EFFORT_DEFAULT = "xhigh";
export const OPENAI_ENDPOINT_RESPONSES = "/v1/responses";
@@ -170,7 +165,7 @@ export function buildModelAdapterTestRequestHash(source) {
normalizeBaseURL(adapter.baseURL),
asString(adapter.apiKey),
asString(adapter.modelID),
adapter.type === "openai" ? asString(adapter.reasoningEffort) : "",
adapter.type === "openai" ? asString(adapter.reasoningEffort || "medium") : "",
adapter.type === "openai" ? normalizeOpenAIEndpoint(adapter.openAIEndpoint) : "",
adapter.type === "openai" ? String(Boolean(adapter.openAIExtraParamsEnabled)) : "false",
adapter.type === "openai" && adapter.openAIExtraParamsEnabled ? asString(adapter.openAIExtraParamsJSON) : "",
@@ -259,7 +254,7 @@ export function createEmptyModelAdapter() {
apiKey: "",
tooltipData: "备注",
modelID: "",
reasoningEffort: "",
reasoningEffort: "medium",
openAIEndpoint: OPENAI_ENDPOINT_RESPONSES,
openAIExtraParamsEnabled: false,
openAIExtraParamsJSON: OPENAI_EXTRA_PARAMS_DEFAULT_JSON,
@@ -334,7 +329,7 @@ function validateAnthropicExtraParamsJSON(value) {
export function normalizeModelAdapter(source) {
const raw = source && typeof source === "object" ? source : {};
const normalizedType = asString(raw.type).toLowerCase();
const normalizedReasoningEffort = normalizeReasoningEffort(raw.reasoningEffort ?? raw.reasoning_effort);
const normalizedReasoningEffort = asString(raw.reasoningEffort || raw.reasoning_effort).toLowerCase();
const normalizedAnthropicThinkingEffort = asString(
raw.anthropicThinkingEffort
?? raw.anthropic_thinking_effort
@@ -367,7 +362,9 @@ export function normalizeModelAdapter(source) {
apiKey: asString(raw.apiKey || raw.key),
tooltipData: asString(raw.tooltipData),
modelID: asString(raw.modelID),
reasoningEffort: normalizedReasoningEffort,
reasoningEffort: SUPPORTED_REASONING_EFFORTS.has(normalizedReasoningEffort)
? normalizedReasoningEffort
: "medium",
openAIEndpoint: normalizedType === "openai" ? normalizedOpenAIEndpoint : "",
openAIExtraParamsEnabled,
openAIExtraParamsJSON,
@@ -445,7 +442,7 @@ export function validateModelAdapters(source) {
return `${prefix} 的上下文窗口必须为正整数`;
}
if (adapter.type === "openai" && !SUPPORTED_REASONING_EFFORTS.has(adapter.reasoningEffort)) {
return `${prefix} 的推理强度仅支持不设置、low、medium、high、xhigh、max`;
return `${prefix} 的推理强度仅支持 low、medium、high、xhigh、max`;
}
if (adapter.type === "anthropic" && adapter.anthropicMaxTokens && (!Number.isInteger(adapter.anthropicMaxTokens) || adapter.anthropicMaxTokens <= 0)) {
return `${prefix} 的最大输出 Token 必须为正整数`;
@@ -609,9 +606,6 @@ async function loadPersistedUserConfig() {
}
async function persistConfigPayload(config, { modelAdaptersOnly = false } = {}) {
if (appState.configSaving) {
return { ok: false, error: "已有配置操作正在进行,请稍后再试" };
}
const payload = buildConfigPayload(config);
const validationError = validateModelAdapters(payload.modelAdapters);
if (validationError) {
@@ -1092,10 +1086,6 @@ export async function refreshModelAdapterTestResults() {
export function startModelAdapterTest(adapter) {
const normalized = normalizeModelAdapter(adapter);
const validationError = validateModelAdapters([normalized]);
if (validationError) {
return Promise.reject(new Error(validationError));
}
return testModelAdapter(normalized).then((rawResult) => {
const result = normalizeModelAdapterTestResult(rawResult);
if (result.adapterID) {
@@ -1124,27 +1114,6 @@ export async function persistUserConfig() {
});
}
export async function exportUserConfigToFile(path) {
return exportUserConfigFile(path);
}
export async function importUserConfigFromFile(path) {
if (appState.serviceRunning || appState.backendRunning || appState.proxyRunning) {
throw new Error("服务运行中不能导入完整配置,请先停止服务");
}
if (appState.configSaving) {
throw new Error("已有配置操作正在进行,请稍后再试");
}
appState.configSaving = true;
try {
const imported = normalizeConfig(await importUserConfigFile(path));
applyConfigToState(imported);
return imported;
} finally {
appState.configSaving = false;
}
}
export async function saveIncludeCacheWriteInHitRate(value) {
const currentConfig = await loadPersistedUserConfig();
const previousValue = appState.includeCacheWriteInHitRate;
@@ -1,14 +0,0 @@
export const SUPPORTED_REASONING_EFFORTS = new Set(["", "low", "medium", "high", "xhigh", "max"]);
export function normalizeReasoningEffort(value) {
if (typeof value === "string") {
return value.trim().toLowerCase();
}
if (value instanceof String) {
return value.toString().trim().toLowerCase();
}
if (typeof value === "number" || typeof value === "boolean") {
return String(value).trim().toLowerCase();
}
return "";
}
@@ -1,20 +0,0 @@
import assert from "node:assert/strict";
import test from "node:test";
import {
normalizeReasoningEffort,
SUPPORTED_REASONING_EFFORTS,
} from "./modelAdapterReasoning.js";
test("normalizeReasoningEffort preserves blank and supported values", () => {
assert.equal(normalizeReasoningEffort(""), "");
assert.equal(normalizeReasoningEffort(" HIGH "), "high");
assert.equal(SUPPORTED_REASONING_EFFORTS.has(normalizeReasoningEffort("max")), true);
});
test("normalizeReasoningEffort preserves unknown values for validation", () => {
const normalized = normalizeReasoningEffort(" Unsupported ");
assert.equal(normalized, "unsupported");
assert.equal(SUPPORTED_REASONING_EFFORTS.has(normalized), false);
});
-3
View File
@@ -2,7 +2,6 @@
import Button from "@/components/ui/Button.vue";
import Card from "@/components/ui/Card.vue";
import HomeMetricsCard from "@/components/HomeMetricsCard.vue";
import CursorAccountCard from "@/components/CursorAccountCard.vue";
import { useMessage } from "@/composables/useMessage";
import { getAdRuntime } from "@/services/clientApi";
import {
@@ -174,8 +173,6 @@ onBeforeUnmount(() => {
</div>
</Card>
<CursorAccountCard />
<Card class="">
<div class="flex items-center justify-between gap-4">
<div>
+2 -23
View File
@@ -5,7 +5,6 @@ import ContentModal from "@/components/ui/ContentModal.vue";
import ModelAdapterTestCard from "@/components/ModelAdapterTestCard.vue";
import ModelEditor from "@/components/ModelEditor.vue";
import { useMessage } from "@/composables/useMessage";
import { useConfigTransfer } from "@/composables/useConfigTransfer";
import Sortable from "sortablejs";
import {
appState,
@@ -88,12 +87,6 @@ function showActionError(title, error) {
message(`${title}${detail}`);
}
const {
configTransferBusy,
handleExportConfig,
handleImportConfig,
} = useConfigTransfer({ message, showActionError });
function maskSecret(value) {
const text = String(value || "").trim();
if (!text) {
@@ -399,26 +392,12 @@ onBeforeUnmount(() => {
<div class="center-row gap-2">
<Button
variant="default"
:disabled="sortSaving || appState.configSaving || batchTesting || configTransferBusy || appState.serviceRunning || appState.backendRunning || appState.proxyRunning"
@click="handleImportConfig"
>
导入配置
</Button>
<Button
variant="default"
:disabled="sortSaving || appState.configSaving || batchTesting || configTransferBusy"
@click="handleExportConfig"
>
导出配置
</Button>
<Button
variant="default"
:disabled="sortSaving || appState.configSaving || configTransferBusy || (!batchTesting && filteredAdapters.length === 0)"
:disabled="sortSaving || appState.configSaving || (!batchTesting && filteredAdapters.length === 0)"
@click="handleTestAllModelAdapters"
>
{{ batchButtonText }}
</Button>
<Button variant="primary" :disabled="sortSaving || appState.configSaving || batchTesting || configTransferBusy" @click="openEditor()">新增模型</Button>
<Button variant="primary" :disabled="sortSaving || appState.configSaving || batchTesting" @click="openEditor()">新增模型</Button>
</div>
</div>
</div>
+19 -17
View File
@@ -70,21 +70,20 @@ func Run(resources EmbeddedResources) error {
logger.Init()
netproxy.InstallDefaultTransport()
if err := appdata.EnsureAssistantHome(); err != nil {
return err
}
certManager, caCertPEM, err := certs.LoadOrCreateManager(appdata.CACertFilePath(), appdata.CAKeyFilePath())
embeddedCACertPEM := certs.EmbeddedCACertPEM()
logEmbeddedCAInfo(embeddedCACertPEM)
certManager, err := certs.NewEmbeddedManager()
if err != nil {
return err
}
logCAInfo(caCertPEM)
defaultBackendBaseURL := "http://" + serverconfig.DefaultBackendListenAddr
defaultBackendBaseURL := browserReachableLoopbackBaseURL(serverconfig.DefaultBackendListenAddr)
proxyServer, err := mitm.NewProxyServer(serverconfig.DefaultProxyListenAddr, defaultBackendBaseURL, "", "", certManager)
if err != nil {
return err
}
proxyService := bridge.NewProxyService(proxyServer, certManager, caCertPEM)
proxyService := bridge.NewProxyService(proxyServer, certManager, embeddedCACertPEM)
adAssetBaseURL := defaultBackendBaseURL
if cfg, err := proxyService.LoadUserConfig(); err == nil {
adAssetBaseURL = browserReachableLoopbackBaseURL(cfg.BackendListenAddr)
@@ -434,29 +433,32 @@ func windowsAdditionalBrowserArgs() []string {
func browserReachableLoopbackBaseURL(listenAddr string) string {
host, port, err := net.SplitHostPort(strings.TrimSpace(listenAddr))
if err != nil || strings.TrimSpace(port) == "" {
return "http://" + serverconfig.DefaultBackendListenAddr
return "https://localhost:8000"
}
host = strings.TrimSpace(host)
if host == "" || host == "0.0.0.0" || host == "::" || host == "[::]" {
host = "127.0.0.1"
}
return "http://" + net.JoinHostPort(host, port)
if host == "127.0.0.1" || host == "::1" || host == "localhost" {
host = "localhost"
}
return "https://" + net.JoinHostPort(host, port)
}
// logCAInfo 记录当前安装专属 CA 的公开信息
func logCAInfo(certPEM []byte) {
// logEmbeddedCAInfo 用于处理与 logEmbeddedCAInfo 相关的逻辑
func logEmbeddedCAInfo(certPEM []byte) {
if len(certPEM) == 0 {
logger.Errorf("installation CA is empty")
logger.Errorf("embedded CA is empty")
return
}
cert, err := parseCert(certPEM)
cert, err := parseEmbeddedCert(certPEM)
if err != nil {
logger.Errorf("parse installation CA failed: %v", err)
logger.Errorf("parse embedded CA failed: %v", err)
return
}
sum := sha256.Sum256(cert.Raw)
logger.Infof(
"installation CA loaded: sha256=%s subject=%s valid=%s~%s",
"embedded CA loaded: sha256=%s subject=%s valid=%s~%s",
strings.ToUpper(hex.EncodeToString(sum[:])),
cert.Subject.String(),
cert.NotBefore.Format(time.RFC3339),
@@ -464,8 +466,8 @@ func logCAInfo(certPEM []byte) {
)
}
// parseCert 解析 DER 或 PEM 编码的证书
func parseCert(data []byte) (*x509.Certificate, error) {
// parseEmbeddedCert 用于处理与 parseEmbeddedCert 相关的逻辑
func parseEmbeddedCert(data []byte) (*x509.Certificate, error) {
if block, _ := pem.Decode(data); block != nil {
return x509.ParseCertificate(block.Bytes)
}
-5
View File
@@ -70,8 +70,3 @@ func LogsRootPath() string {
func CACertFilePath() string {
return filepath.Join(DataRootPath(), "ca.crt")
}
// CAKeyFilePath 返回仅供本安装使用的 CA 私钥路径。
func CAKeyFilePath() string {
return filepath.Join(DataRootPath(), "ca.key")
}
+1 -3
View File
@@ -97,7 +97,6 @@ internal/backend/
- `~/.cursor-local-assistant-v2/config.yaml`
- `~/.cursor-local-assistant-v2/data/ca.crt`
- `~/.cursor-local-assistant-v2/data/ca.key`
- `~/.cursor-local-assistant-v2/data/ads/`
- `~/.cursor-local-assistant-v2/history/`
- `~/.cursor-local-assistant-v2/logs/`
@@ -105,8 +104,7 @@ internal/backend/
约定:
- `config.yaml` 是用户配置
- `data/ca.crt`首次运行时为当前用户生成、注入给宿主的 CA 证书
- `data/ca.key` 是与该证书配套的本地私钥,权限固定为 `0600`,不得打包或提交到仓库
- `data/ca.crt` 是注入给宿主的 CA 证书
- `data/ads/` 是广告包与资源缓存目录
- `history/` 是会话事实与全局 usage JSON 目录,不属于日志
- `logs/` 只保留必要文本运行日志
+1 -79
View File
@@ -2,15 +2,8 @@
package execbridge
import (
"bytes"
"crypto/sha256"
"encoding/json"
"fmt"
"image"
_ "image/gif"
_ "image/jpeg"
_ "image/png"
"net/http"
"strings"
"sync/atomic"
"time"
@@ -38,18 +31,10 @@ type ExecApplyResult struct {
ToolResultPayload string
// ToolCall 保存可用于发 ToolCallCompletedUpdate 的工具调用对象;当前仅对支持 ToolCall 的执行型工具可用。
ToolCall *agentv1.ToolCall
// ContentBlobs 保存需要在提交 history 前写入内容寻址存储的二进制内容。
ContentBlobs []ContentBlob
// ExecuteHookResponse 保存 execute hook 的结构化响应。
ExecuteHookResponse *agentv1.ExecuteHookResponse
}
// ContentBlob 表示由内容哈希稳定寻址的执行结果二进制数据。
type ContentBlob struct {
ID []byte
Data []byte
}
// OpenExecContext 表示执行桥打开请求时需要的最小上下文。
type OpenExecContext struct {
ConversationID string
@@ -160,9 +145,6 @@ func (bridge *Bridge) ApplyExecClientMessage(msg *agentv1.ExecClientMessage, pen
readResult := normalizeReadResultForModel(msg.GetReadResult())
result.ToolResultPayload = summarizeReadResult(readResult)
result.ToolCall = buildReadCompletedToolCall(pending.ToolCallID, pending.ArgsJSON, readResult)
if contentBlob, ok := readImageContentBlob(readResult); ok {
result.ContentBlobs = []ContentBlob{contentBlob}
}
result.IsTerminal = true
return result, nil
case "write":
@@ -2303,64 +2285,6 @@ func buildReadMcpResourceCompletedToolCall(argsJSON []byte, result *agentv1.Read
}
}
func supportedReadImageMIMEType(data []byte) string {
if len(data) == 0 {
return ""
}
detected := strings.ToLower(strings.TrimSpace(http.DetectContentType(data)))
configuration, format, err := image.DecodeConfig(bytes.NewReader(data))
if err != nil || configuration.Width <= 0 || configuration.Height <= 0 {
return ""
}
switch strings.ToLower(strings.TrimSpace(format)) {
case "png":
if detected == "image/png" {
return detected
}
case "jpeg":
if detected == "image/jpeg" {
return detected
}
case "gif":
if detected == "image/gif" {
return detected
}
}
return ""
}
func readImageContentBlob(result *agentv1.ReadResult) (ContentBlob, bool) {
success := result.GetSuccess()
if success == nil {
return ContentBlob{}, false
}
data := success.GetData()
if supportedReadImageMIMEType(data) == "" {
return ContentBlob{}, false
}
digest := sha256.Sum256(data)
return ContentBlob{
ID: append([]byte(nil), digest[:]...),
Data: append([]byte(nil), data...),
}, true
}
func readImageBlobID(data []byte) ([]byte, bool) {
if supportedReadImageMIMEType(data) == "" {
return nil, false
}
digest := sha256.Sum256(data)
return append([]byte(nil), digest[:]...), true
}
func readImageDataBlobOutput(data []byte) *agentv1.ReadToolSuccess_DataBlobId {
blobID, ok := readImageBlobID(data)
if !ok {
return nil
}
return &agentv1.ReadToolSuccess_DataBlobId{DataBlobId: blobID}
}
// convertReadResultToReadToolResult 把 `ReadResult` 映射为 `ReadToolResult`。
func convertReadResultToReadToolResult(result *agentv1.ReadResult) *agentv1.ReadToolResult {
if result == nil {
@@ -2394,9 +2318,7 @@ func convertReadResultToReadToolResult(result *agentv1.ReadResult) *agentv1.Read
if content != "" {
toolSuccess.Output = &agentv1.ReadToolSuccess_Content{Content: content}
} else if len(data) > 0 {
if imageOutput := readImageDataBlobOutput(data); imageOutput != nil {
toolSuccess.Output = imageOutput
} else if len(data) > readReplayBinaryLimit {
if len(data) > readReplayBinaryLimit {
toolSuccess.ExceededLimit = true
toolSuccess.Output = &agentv1.ReadToolSuccess_Content{
Content: replayTruncationNotice("Read binary data", readReplayBinaryLimit, 0, len(data)),
@@ -1,125 +0,0 @@
package execbridge
import (
"bytes"
"crypto/sha256"
"image"
"image/color"
"image/png"
"strings"
"testing"
"cursor/gen/agentv1"
runtimecore "cursor/internal/backend/agent/core"
)
func TestApplyExecClientMessageReturnsContentAddressedReadImage(t *testing.T) {
imageData := validReadTestPNG(t)
wantBlobID := sha256.Sum256(imageData)
result, err := NewBridge().ApplyExecClientMessage(&agentv1.ExecClientMessage{
Message: &agentv1.ExecClientMessage_ReadResult{
ReadResult: &agentv1.ReadResult{
Result: &agentv1.ReadResult_Success{
Success: &agentv1.ReadSuccess{
Path: "diagram.png",
FileSize: int64(len(imageData)),
OutputBlobId: append([]byte(nil), wantBlobID[:]...),
Output: &agentv1.ReadSuccess_Data{Data: imageData},
},
},
},
},
}, runtimecore.PendingExec{
ExecKind: "read",
ToolCallID: "call-1",
ArgsJSON: []byte(`{"path":"diagram.png"}`),
})
if err != nil {
t.Fatalf("ApplyExecClientMessage() error = %v", err)
}
if len(result.ContentBlobs) != 1 {
t.Fatalf("content blob count = %d, want 1", len(result.ContentBlobs))
}
if !bytes.Equal(result.ContentBlobs[0].ID, wantBlobID[:]) || !bytes.Equal(result.ContentBlobs[0].Data, imageData) {
t.Fatalf("content blob = %#v", result.ContentBlobs[0])
}
readSuccess := result.ToolCall.GetReadToolCall().GetResult().GetSuccess()
if readSuccess == nil {
t.Fatal("read tool result is not successful")
}
if !bytes.Equal(readSuccess.GetDataBlobId(), wantBlobID[:]) {
t.Fatalf("data_blob_id = %x, want %x", readSuccess.GetDataBlobId(), wantBlobID)
}
if len(readSuccess.GetData()) != 0 {
t.Fatalf("read tool result retained %d image bytes", len(readSuccess.GetData()))
}
}
func TestApplyExecClientMessageUsesComputedImageBlobID(t *testing.T) {
imageData := validReadTestPNG(t)
wantBlobID := sha256.Sum256(imageData)
result, err := NewBridge().ApplyExecClientMessage(&agentv1.ExecClientMessage{
Message: &agentv1.ExecClientMessage_ReadResult{
ReadResult: &agentv1.ReadResult{
Result: &agentv1.ReadResult_Success{
Success: &agentv1.ReadSuccess{
Path: "diagram.png",
OutputBlobId: bytes.Repeat([]byte{0xff}, sha256.Size),
Output: &agentv1.ReadSuccess_Data{Data: imageData},
},
},
},
},
}, runtimecore.PendingExec{ExecKind: "read", ToolCallID: "call-1"})
if err != nil {
t.Fatalf("ApplyExecClientMessage() error = %v", err)
}
if !bytes.Equal(result.ContentBlobs[0].ID, wantBlobID[:]) {
t.Fatalf("content blob id = %x, want computed %x", result.ContentBlobs[0].ID, wantBlobID)
}
}
func TestConvertReadResultKeepsTextAndLimitsUnsupportedBinary(t *testing.T) {
textResult := convertReadResultToReadToolResult(&agentv1.ReadResult{
Result: &agentv1.ReadResult_Success{
Success: &agentv1.ReadSuccess{
Path: "notes.txt",
Output: &agentv1.ReadSuccess_Content{Content: "hello"},
},
},
})
if got := textResult.GetSuccess().GetContent(); got != "hello" {
t.Fatalf("text read content = %q, want hello", got)
}
largeBinary := bytes.Repeat([]byte{0xff}, readReplayBinaryLimit+1)
binaryResult := convertReadResultToReadToolResult(&agentv1.ReadResult{
Result: &agentv1.ReadResult_Success{
Success: &agentv1.ReadSuccess{
Path: "archive.bin",
Output: &agentv1.ReadSuccess_Data{Data: largeBinary},
},
},
})
binarySuccess := binaryResult.GetSuccess()
if binarySuccess == nil || !binarySuccess.GetExceededLimit() {
t.Fatal("large non-image binary was not limited")
}
if binarySuccess.GetData() != nil || binarySuccess.GetDataBlobId() != nil {
t.Fatal("large non-image binary was retained")
}
if !strings.Contains(binarySuccess.GetContent(), "Read binary data") {
t.Fatalf("large binary fallback = %q", binarySuccess.GetContent())
}
}
func validReadTestPNG(t *testing.T) []byte {
t.Helper()
value := image.NewRGBA(image.Rect(0, 0, 2, 2))
value.Set(0, 0, color.RGBA{R: 0x44, G: 0x88, B: 0xcc, A: 0xff})
var encoded bytes.Buffer
if err := png.Encode(&encoded, value); err != nil {
t.Fatalf("encode test png: %v", err)
}
return encoded.Bytes()
}
+2 -19
View File
@@ -1126,16 +1126,7 @@ func isAnthropicCacheableBlock(block map[string]any) bool {
case contentPartTypeText:
return strings.TrimSpace(anthropicStringField(block, "text")) != ""
case "tool_result":
switch content := block["content"].(type) {
case string:
return strings.TrimSpace(content) != ""
case []map[string]any:
return len(content) > 0
case []any:
return len(content) > 0
default:
return false
}
return strings.TrimSpace(anthropicStringField(block, "content")) != ""
case "tool_use":
return strings.TrimSpace(anthropicStringField(block, "id")) != "" && strings.TrimSpace(anthropicStringField(block, "name")) != ""
default:
@@ -1187,18 +1178,10 @@ func normalizeAnthropicProviderMessages(input []Message, thinkingEnabled bool, r
if toolUseID == "" {
return nil, nil, fmt.Errorf("anthropic tool message requires tool_call_id")
}
var content any = message.Content
if hasImageContentParts(message.ContentParts) {
contentBlocks, err := anthropicContentBlocks(message)
if err != nil {
return nil, nil, err
}
content = contentBlocks
}
pendingToolResults = append(pendingToolResults, map[string]any{
"type": "tool_result",
"tool_use_id": toolUseID,
"content": content,
"content": message.Content,
})
case "user", "assistant":
flushToolResults()
+3 -36
View File
@@ -49,8 +49,7 @@ type openAIResponsesRequestBody struct {
}
type openAIResponsesReasoning struct {
Effort string `json:"effort,omitempty"`
Summary string `json:"summary,omitempty"`
Effort string `json:"effort,omitempty"`
}
type openAIToolAccumulator struct {
@@ -945,7 +944,7 @@ func (adapter *OpenAIAdapter) streamResponses(ctx context.Context, req StreamReq
requestBody.Tools = tools
}
if effort := strings.TrimSpace(req.ReasoningEffort); effort != "" {
requestBody.Reasoning = &openAIResponsesReasoning{Effort: effort, Summary: "auto"}
requestBody.Reasoning = &openAIResponsesReasoning{Effort: effort}
requestBody.Include = []string{"reasoning.encrypted_content"}
}
body = requestBody
@@ -1956,7 +1955,6 @@ func normalizeOpenAIResponsesInput(messages []Message) (string, []map[string]any
instructionParts := make([]string, 0, 2)
items := make([]map[string]any, 0, len(messages))
responsesCallIDs := make(map[string]string)
emittedCallIDs := make(map[string]struct{})
activeAssistantReasoningKey := ""
for _, message := range messages {
role := strings.TrimSpace(message.Role)
@@ -1969,38 +1967,10 @@ func normalizeOpenAIResponsesInput(messages []Message) (string, []map[string]any
}
if role == "tool" && strings.TrimSpace(message.ToolCallID) != "" {
callID := openAIResponsesToolMessageCallID(message, responsesCallIDs)
if strings.TrimSpace(callID) != "" {
if _, ok := emittedCallIDs[callID]; !ok {
// 历史损坏时可能出现没有配对 function_call 的工具结果
// (例如旧版本回放逻辑剥离了 assistant 调用但保留了结果)。
// Responses API 会直接拒绝这种 input,这里补一个占位 function_call
// 让旧会话可以继续;无法补齐时丢弃该结果。
if name := strings.TrimSpace(message.Name); name != "" {
items = append(items, map[string]any{
"type": "function_call",
"call_id": callID,
"name": name,
"arguments": "{}",
"status": "completed",
})
emittedCallIDs[callID] = struct{}{}
} else {
continue
}
}
}
var output any = openAIResponsesMessageText(message)
if hasImageContentParts(message.ContentParts) {
content, err := openAIResponsesMessageContent(message, false)
if err != nil {
return "", nil, err
}
output = content
}
items = append(items, map[string]any{
"type": "function_call_output",
"call_id": callID,
"output": output,
"output": openAIResponsesMessageText(message),
})
activeAssistantReasoningKey = ""
continue
@@ -2055,9 +2025,6 @@ func normalizeOpenAIResponsesInput(messages []Message) (string, []map[string]any
toolItem["status"] = "completed"
}
items = append(items, toolItem)
if strings.TrimSpace(callID) != "" {
emittedCallIDs[strings.TrimSpace(callID)] = struct{}{}
}
}
}
}
@@ -1,96 +0,0 @@
package modeladapter
import "testing"
// 历史损坏时(旧版本回放逻辑剥离了 assistant 调用但保留了结果),
// function_call_output 会缺少配对的 function_callResponses API 会拒绝。
// 这里验证 adapter 会为孤儿结果补一个占位 function_call,让旧会话可以继续。
func TestNormalizeOpenAIResponsesInputSynthesizesCallForOrphanToolOutput(t *testing.T) {
messages := []Message{
{Role: "user", Content: "query"},
{Role: "assistant", Content: "我先快速定位上下文"},
{Role: "tool", Name: "Grep", ToolCallID: "call_xrN6", Content: "grep result"},
{Role: "user", Content: "next"},
}
_, items, err := normalizeOpenAIResponsesInput(messages)
if err != nil {
t.Fatalf("normalizeOpenAIResponsesInput failed: %v", err)
}
var callIndexes []int
var outputIndexes []int
for index, item := range items {
if item["type"] == "function_call" && item["call_id"] == "call_xrN6" {
callIndexes = append(callIndexes, index)
}
if item["type"] == "function_call_output" && item["call_id"] == "call_xrN6" {
outputIndexes = append(outputIndexes, index)
}
}
if len(callIndexes) != 1 {
t.Fatalf("expected 1 synthesized function_call, got %d: %+v", len(callIndexes), items)
}
if len(outputIndexes) != 1 || outputIndexes[0] != callIndexes[0]+1 {
t.Fatalf("expected function_call_output right after synthesized function_call, calls=%v outputs=%v", callIndexes, outputIndexes)
}
if got := items[callIndexes[0]]["name"]; got != "Grep" {
t.Fatalf("expected synthesized function_call name Grep, got %v", got)
}
}
// 连工具名都没有的孤儿结果只能丢弃,避免上游 400。
func TestNormalizeOpenAIResponsesInputDropsNamelessOrphanToolOutput(t *testing.T) {
messages := []Message{
{Role: "user", Content: "query"},
{Role: "tool", ToolCallID: "call_unknown", Content: "orphan result"},
{Role: "user", Content: "next"},
}
_, items, err := normalizeOpenAIResponsesInput(messages)
if err != nil {
t.Fatalf("normalizeOpenAIResponsesInput failed: %v", err)
}
for _, item := range items {
if item["type"] == "function_call_output" {
t.Fatalf("expected orphan function_call_output to be dropped, got %+v", item)
}
}
}
// 正常配对的调用与结果不应受防御逻辑影响。
func TestNormalizeOpenAIResponsesInputKeepsPairedCallAndOutput(t *testing.T) {
messages := []Message{
{Role: "user", Content: "query"},
{
Role: "assistant",
ToolCalls: []ToolCallDescriptor{{
ID: "call_xrN6",
Type: "function",
Function: ToolCallFunctionShape{
Name: "Grep",
Arguments: `{"pattern":"ref"}`,
},
}},
},
{Role: "tool", Name: "Grep", ToolCallID: "call_xrN6", Content: "grep result"},
{Role: "user", Content: "next"},
}
_, items, err := normalizeOpenAIResponsesInput(messages)
if err != nil {
t.Fatalf("normalizeOpenAIResponsesInput failed: %v", err)
}
var calls, outputs int
for _, item := range items {
if item["type"] == "function_call" && item["call_id"] == "call_xrN6" {
calls++
}
if item["type"] == "function_call_output" && item["call_id"] == "call_xrN6" {
outputs++
}
}
if calls != 1 || outputs != 1 {
t.Fatalf("expected exactly one paired function_call/output, got calls=%d outputs=%d", calls, outputs)
}
}
@@ -2,140 +2,12 @@ package modeladapter
import (
"context"
"encoding/json"
"fmt"
"net/http"
"net/http/httptest"
"testing"
)
func TestOpenAIResponsesRequestsReasoningSummary(t *testing.T) {
var requestBody map[string]any
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {
if err := json.NewDecoder(request.Body).Decode(&requestBody); err != nil {
http.Error(writer, err.Error(), http.StatusBadRequest)
return
}
writer.Header().Set("Content-Type", "text/event-stream")
_, _ = fmt.Fprint(writer, "data: {\"type\":\"response.reasoning_summary_text.delta\",\"delta\":\"check the request\"}\n\n")
_, _ = fmt.Fprint(writer, "data: {\"type\":\"response.completed\",\"response\":{\"id\":\"resp-1\",\"model\":\"gpt-5.6\",\"status\":\"completed\",\"output_text\":\"done\"}}\n\n")
_, _ = fmt.Fprint(writer, "data: [DONE]\n\n")
}))
defer server.Close()
adapter := &OpenAIAdapter{client: server.Client()}
events := make([]ModelEvent, 0, 4)
err := adapter.Stream(context.Background(), StreamRequest{
RequestID: "request-1",
RunID: "run-1",
ModelCallID: "model-call-1",
BaseURL: server.URL,
APIKey: "test-key",
ProviderModelID: "gpt-5.6",
OpenAIEndpoint: "/v1/responses",
ReasoningEffort: "high",
Messages: []Message{{Role: "user", Content: "hello"}},
MaxTokens: 128,
}, func(event ModelEvent) error {
events = append(events, event)
return nil
})
if err != nil {
t.Fatalf("stream failed: %v", err)
}
reasoning, ok := requestBody["reasoning"].(map[string]any)
if !ok {
t.Fatalf("reasoning request body missing: %#v", requestBody)
}
if got := reasoning["effort"]; got != "high" {
t.Fatalf("reasoning.effort = %#v, want high", got)
}
if got := reasoning["summary"]; got != "auto" {
t.Fatalf("reasoning.summary = %#v, want auto", got)
}
include, ok := requestBody["include"].([]any)
if !ok || len(include) != 1 || include[0] != "reasoning.encrypted_content" {
t.Fatalf("reasoning include = %#v, want encrypted content", requestBody["include"])
}
assertOpenAIEventKindCount(t, events, ModelEventKindThinkingDelta, 1)
assertOpenAIEventKindCount(t, events, ModelEventKindThinkingCompleted, 1)
assertOpenAIEventKindCount(t, events, ModelEventKindTextDelta, 1)
}
func TestOpenAIResponsesOmitsReasoningWhenEffortBlank(t *testing.T) {
var requestBody map[string]any
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {
if err := json.NewDecoder(request.Body).Decode(&requestBody); err != nil {
http.Error(writer, err.Error(), http.StatusBadRequest)
return
}
writer.Header().Set("Content-Type", "text/event-stream")
_, _ = fmt.Fprint(writer, "data: {\"type\":\"response.completed\",\"response\":{\"id\":\"resp-1\",\"model\":\"grok-composer-2.5-fast\",\"status\":\"completed\",\"output_text\":\"done\"}}\n\n")
_, _ = fmt.Fprint(writer, "data: [DONE]\n\n")
}))
defer server.Close()
adapter := &OpenAIAdapter{client: server.Client()}
err := adapter.Stream(context.Background(), StreamRequest{
RequestID: "request-1",
RunID: "run-1",
ModelCallID: "model-call-1",
BaseURL: server.URL,
APIKey: "test-key",
ProviderModelID: "grok-composer-2.5-fast",
OpenAIEndpoint: "/v1/responses",
Messages: []Message{{Role: "user", Content: "hello"}},
MaxTokens: 128,
}, func(ModelEvent) error { return nil })
if err != nil {
t.Fatalf("stream failed: %v", err)
}
if _, exists := requestBody["reasoning"]; exists {
t.Fatalf("reasoning should be omitted when effort is blank: %#v", requestBody["reasoning"])
}
if _, exists := requestBody["reasoning_effort"]; exists {
t.Fatalf("reasoning_effort should be omitted when effort is blank: %#v", requestBody["reasoning_effort"])
}
}
func TestOpenAIChatCompletionsOmitsReasoningWhenEffortBlank(t *testing.T) {
var requestBody map[string]any
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {
if err := json.NewDecoder(request.Body).Decode(&requestBody); err != nil {
http.Error(writer, err.Error(), http.StatusBadRequest)
return
}
writer.Header().Set("Content-Type", "text/event-stream")
_, _ = fmt.Fprint(writer, "data: {\"model\":\"grok-composer-2.5-fast\",\"choices\":[{\"delta\":{\"content\":\"done\"},\"finish_reason\":\"stop\"}]}\n\n")
_, _ = fmt.Fprint(writer, "data: [DONE]\n\n")
}))
defer server.Close()
adapter := &OpenAIAdapter{client: server.Client()}
err := adapter.Stream(context.Background(), StreamRequest{
RequestID: "request-1",
RunID: "run-1",
ModelCallID: "model-call-1",
BaseURL: server.URL,
APIKey: "test-key",
ProviderModelID: "grok-composer-2.5-fast",
OpenAIEndpoint: "/v1/chat/completions",
Messages: []Message{{Role: "user", Content: "hello"}},
MaxTokens: 128,
}, func(ModelEvent) error { return nil })
if err != nil {
t.Fatalf("stream failed: %v", err)
}
for _, field := range []string{"reasoning_effort", "reasoning", "include"} {
if value, exists := requestBody[field]; exists {
t.Fatalf("%s should be omitted when effort is blank: %#v", field, value)
}
}
}
func TestOpenAIChatCompletionsIgnoresBlankFinishReason(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {
writer.Header().Set("Content-Type", "text/event-stream")
+40 -65
View File
@@ -339,82 +339,57 @@ func mergeProviderReasoningMetadata(last *Message, current Message) {
}
}
// providerToolResponseWindowEnd 返回 assistant tool-call 消息的响应收集窗口右边界(不含)。
// 窗口内除 tool 结果消息外,还允许出现同轮穿插的纯文本 assistant 消息:
// 部分模型(如 gpt-5.3-codex-spark)会在同一条响应里先输出 function_call 再输出说明文本,
// 回放顺序为 assistant[tool_call] -> assistant[text] -> tool[result],若只收集紧邻的
// tool 消息,会把有结果回放的调用误判为悬空。
func providerToolResponseWindowEnd(messages []Message, index int) int {
end := index + 1
for end < len(messages) {
candidate := messages[end]
switch {
case strings.TrimSpace(candidate.Role) == "tool":
end++
case strings.TrimSpace(candidate.Role) == "assistant" && len(candidate.ToolCalls) == 0:
end++
default:
return end
}
}
return end
}
func trimDanglingAssistantToolCalls(input []Message) []Message {
if len(input) == 0 {
return nil
}
survivingToolCallIDs := make(map[string]struct{})
for index, message := range input {
if strings.TrimSpace(message.Role) != "assistant" || len(message.ToolCalls) == 0 {
continue
}
responded := make(map[string]struct{}, len(message.ToolCalls))
for scan := index + 1; scan < providerToolResponseWindowEnd(input, index); scan++ {
if strings.TrimSpace(input[scan].Role) != "tool" {
continue
}
if toolCallID := strings.TrimSpace(input[scan].ToolCallID); toolCallID != "" {
responded[toolCallID] = struct{}{}
}
}
for _, toolCall := range message.ToolCalls {
if toolCallID := strings.TrimSpace(toolCall.ID); toolCallID != "" {
if _, ok := responded[toolCallID]; ok {
survivingToolCallIDs[toolCallID] = struct{}{}
}
}
}
}
trimmed := make([]Message, 0, len(input))
for _, item := range input {
message := cloneProviderMessage(item)
if strings.TrimSpace(message.Role) == "assistant" && len(message.ToolCalls) > 0 {
nextToolCalls := make([]ToolCallDescriptor, 0, len(message.ToolCalls))
for _, toolCall := range message.ToolCalls {
if _, ok := survivingToolCallIDs[strings.TrimSpace(toolCall.ID)]; !ok {
continue
}
toolCall.Index = len(nextToolCalls)
nextToolCalls = append(nextToolCalls, toolCall)
}
if len(nextToolCalls) == 0 {
if strings.TrimSpace(message.Content) == "" && len(message.ContentParts) == 0 && strings.TrimSpace(message.ReasoningContent) == "" {
continue
}
message.ToolCalls = nil
} else {
message.ToolCalls = nextToolCalls
}
for index := 0; index < len(input); index++ {
message := cloneProviderMessage(input[index])
if strings.TrimSpace(message.Role) != "assistant" || len(message.ToolCalls) == 0 {
trimmed = append(trimmed, message)
continue
}
if strings.TrimSpace(message.Role) == "tool" && strings.TrimSpace(message.ToolCallID) != "" {
if _, ok := survivingToolCallIDs[strings.TrimSpace(message.ToolCallID)]; !ok {
end := index + 1
responded := make(map[string]struct{}, len(message.ToolCalls))
for end < len(input) && strings.TrimSpace(input[end].Role) == "tool" {
toolCallID := strings.TrimSpace(input[end].ToolCallID)
if toolCallID != "" {
responded[toolCallID] = struct{}{}
}
end++
}
nextToolCalls := make([]ToolCallDescriptor, 0, len(message.ToolCalls))
allowedToolCallIDs := make(map[string]struct{}, len(message.ToolCalls))
for _, toolCall := range message.ToolCalls {
toolCallID := strings.TrimSpace(toolCall.ID)
if _, ok := responded[toolCallID]; !ok {
continue
}
item := toolCall
item.Index = len(nextToolCalls)
nextToolCalls = append(nextToolCalls, item)
allowedToolCallIDs[toolCallID] = struct{}{}
}
trimmed = append(trimmed, message)
if len(nextToolCalls) > 0 {
message.ToolCalls = nextToolCalls
trimmed = append(trimmed, message)
for toolIndex := index + 1; toolIndex < end; toolIndex++ {
toolMessage := cloneProviderMessage(input[toolIndex])
if _, ok := allowedToolCallIDs[strings.TrimSpace(toolMessage.ToolCallID)]; !ok {
continue
}
trimmed = append(trimmed, toolMessage)
}
} else if strings.TrimSpace(message.Content) != "" || len(message.ContentParts) > 0 || strings.TrimSpace(message.ReasoningContent) != "" {
message.ToolCalls = nil
trimmed = append(trimmed, message)
}
index = end - 1
}
return trimmed
}
@@ -1,77 +0,0 @@
package modeladapter
import "testing"
func routerToolCall(id string, name string) ToolCallDescriptor {
return ToolCallDescriptor{
ID: id,
Type: "function",
Function: ToolCallFunctionShape{
Name: name,
Arguments: `{"pattern":"ref"}`,
},
}
}
// 模型在同一条响应里先输出 function_call 再输出说明文本时,
// sanitize 后的消息序列为 assistant[tool_call] -> assistant[text] -> tool[result]。
// trimDanglingAssistantToolCalls 需要越过中间的文本消息收集工具结果,
// 否则会产生孤儿 function_call_outputResponses API 400)。
func TestTrimDanglingAssistantToolCallsKeepsInterleavedTextResponses(t *testing.T) {
input := []Message{
{Role: "user", Content: "query"},
{
Role: "assistant",
ToolCalls: []ToolCallDescriptor{
routerToolCall("call_1", "Grep"),
routerToolCall("call_2", "Read"),
},
},
{Role: "tool", Name: "Grep", ToolCallID: "call_1", Content: "grep result"},
{Role: "assistant", Content: "我先快速定位上下文"},
{Role: "tool", Name: "Read", ToolCallID: "call_2", Content: "read result"},
{Role: "user", Content: "next"},
}
trimmed := sanitizeProviderMessages(input)
if len(trimmed) != 6 {
t.Fatalf("expected 6 messages, got %d: %+v", len(trimmed), trimmed)
}
if len(trimmed[1].ToolCalls) != 2 {
t.Fatalf("expected both tool calls to survive, got %+v", trimmed[1].ToolCalls)
}
if trimmed[3].Role != "assistant" || trimmed[3].Content == "" {
t.Fatalf("expected interleaved assistant text to survive, got %+v", trimmed[3])
}
if trimmed[4].ToolCallID != "call_2" {
t.Fatalf("expected call_2 result to survive, got %+v", trimmed[4])
}
}
// 完全没有结果回放的调用仍应被剥离,且孤儿 tool 结果不得保留。
func TestTrimDanglingAssistantToolCallsDropsUnrespondedCallsAndOrphanResults(t *testing.T) {
input := []Message{
{Role: "user", Content: "query"},
{
Role: "assistant",
ToolCalls: []ToolCallDescriptor{
routerToolCall("call_1", "Grep"),
routerToolCall("call_2", "Read"),
},
},
{Role: "tool", Name: "Grep", ToolCallID: "call_1", Content: "grep result"},
{Role: "tool", Name: "Read", ToolCallID: "call_3", Content: "orphan result"},
{Role: "user", Content: "next"},
}
trimmed := sanitizeProviderMessages(input)
if len(trimmed) != 4 {
t.Fatalf("expected 4 messages, got %d: %+v", len(trimmed), trimmed)
}
if len(trimmed[1].ToolCalls) != 1 || trimmed[1].ToolCalls[0].ID != "call_1" {
t.Fatalf("expected only call_1 to survive, got %+v", trimmed[1].ToolCalls)
}
if trimmed[2].ToolCallID != "call_1" {
t.Fatalf("expected only call_1 result to survive, got %+v", trimmed[2])
}
}
@@ -1,64 +1,10 @@
package modeladapter
import (
"context"
"reflect"
"testing"
"time"
legacyruntime "cursor/internal/runtime"
)
type recordingModelAdapter struct {
request StreamRequest
}
func (adapter *recordingModelAdapter) Stream(_ context.Context, req StreamRequest, _ func(ModelEvent) error) error {
adapter.request = req
return nil
}
type staticChannelResolver struct {
channel *legacyruntime.ResolvedChannel
}
func (resolver staticChannelResolver) SelectChannelForModel(context.Context, string) (*legacyruntime.ResolvedChannel, error) {
return resolver.channel, nil
}
func (staticChannelResolver) ProviderStreamIdleTimeout(context.Context) time.Duration {
return time.Second
}
func TestRouterRuntimeDisabledClearsReasoningEffort(t *testing.T) {
openAI := &recordingModelAdapter{}
router := &Router{
openai: openAI,
resolver: staticChannelResolver{channel: &legacyruntime.ResolvedChannel{
ID: "channel-a",
Provider: "openai",
Model: "grok-composer-2.5-fast",
ReasoningEffort: "medium",
}},
}
requestKnobs := map[string]any{"reasoning_effort": "medium"}
err := router.Stream(context.Background(), StreamRequest{
ModelID: "channel-a",
ThinkingEffort: "disabled",
RequestKnobs: requestKnobs,
}, func(ModelEvent) error { return nil })
if err != nil {
t.Fatalf("Stream returned error: %v", err)
}
if got := openAI.request.ReasoningEffort; got != "" {
t.Fatalf("ReasoningEffort = %q, want blank", got)
}
if _, exists := openAI.request.RequestKnobs["reasoning_effort"]; exists {
t.Fatalf("reasoning_effort knob should be removed: %#v", openAI.request.RequestKnobs)
}
}
func TestSanitizeProviderMessagesMergesLegacyAssistantTextAndToolCallTurnsIdempotently(t *testing.T) {
input := []Message{
{
@@ -1,88 +0,0 @@
package modeladapter
import (
"strings"
"testing"
)
func TestToolImageProviderEncodings(t *testing.T) {
message := toolImageMessageForTest()
t.Run("openai_chat", func(t *testing.T) {
items, err := normalizeOpenAIProviderMessages([]Message{message}, false)
if err != nil {
t.Fatalf("normalizeOpenAIProviderMessages() error = %v", err)
}
if len(items) != 1 || items[0]["role"] != "tool" || items[0]["tool_call_id"] != "call-1" {
t.Fatalf("openai chat tool message = %#v", items)
}
content, ok := items[0]["content"].([]map[string]any)
if !ok || len(content) != 2 {
t.Fatalf("openai chat content = %#v", items[0]["content"])
}
imageURL, ok := content[1]["image_url"].(map[string]any)
if content[1]["type"] != "image_url" || !ok || !strings.HasPrefix(imageURL["url"].(string), "data:image/png;base64,") {
t.Fatalf("openai chat image part = %#v", content[1])
}
})
t.Run("openai_responses", func(t *testing.T) {
_, items, err := normalizeOpenAIResponsesInput([]Message{message})
if err != nil {
t.Fatalf("normalizeOpenAIResponsesInput() error = %v", err)
}
// 孤儿 tool 结果(无前置 assistant 调用)会补一个占位 function_call
// 保证每个 function_call_output 都有配对调用。
if len(items) != 2 || items[0]["type"] != "function_call" || items[0]["call_id"] != "call-1" || items[1]["type"] != "function_call_output" {
t.Fatalf("openai responses items = %#v", items)
}
content, ok := items[1]["output"].([]map[string]any)
if !ok || len(content) != 2 {
t.Fatalf("openai responses output = %#v", items[1]["output"])
}
if content[0]["type"] != "input_text" || content[1]["type"] != "input_image" {
t.Fatalf("openai responses content = %#v", content)
}
})
t.Run("anthropic", func(t *testing.T) {
_, messages, err := normalizeAnthropicProviderMessages([]Message{message}, false, false)
if err != nil {
t.Fatalf("normalizeAnthropicProviderMessages() error = %v", err)
}
if len(messages) != 1 || messages[0].Role != "user" || len(messages[0].Content) != 1 {
t.Fatalf("anthropic messages = %#v", messages)
}
toolResult := messages[0].Content[0]
if toolResult["type"] != "tool_result" || toolResult["tool_use_id"] != "call-1" {
t.Fatalf("anthropic tool result = %#v", toolResult)
}
content, ok := toolResult["content"].([]map[string]any)
if !ok || len(content) != 2 {
t.Fatalf("anthropic tool content = %#v", toolResult["content"])
}
if content[0]["type"] != "text" || content[1]["type"] != "image" {
t.Fatalf("anthropic content blocks = %#v", content)
}
})
}
func toolImageMessageForTest() Message {
return Message{
Role: "tool",
Content: "read binary bytes=16",
ToolCallID: "call-1",
Name: "Read",
ContentParts: []ContentPart{
{Type: "text", Text: "read binary bytes=16"},
{
Type: "image",
Image: &ImageContent{
MIMEType: "image/png",
Path: "diagram.png",
Data: []byte("\x89PNG\r\n\x1a\nimage"),
},
},
},
}
}
-3
View File
@@ -730,9 +730,6 @@ func (service *Service) handleProviderDoneEvent(stream *ActiveStream, payload *s
service.setTurnPhase(stream, TurnPhaseFailed)
return service.closeStreamWithProviderError(stream, conversationID, turnSeq, requestID, accumulatedText, accumulatedReasoning, accumulatedReasoningSignature, accumulatedReasoningSignatureSource, accumulatedReasoningItemID, accumulatedReasoningStatus, accumulatedReasoningSummary, usage, providerErr, !hadToolInvocation)
}
if err := service.flushAssistantText(stream, conversationID, turnSeq, requestID, accumulatedText, accumulatedReasoning, accumulatedReasoningSignature, accumulatedReasoningSignatureSource, accumulatedReasoningItemID, accumulatedReasoningStatus, accumulatedReasoningSummary, !hadToolInvocation); err != nil {
return service.failStream(stream, "unknown", fmt.Errorf("flush failed provider output: %w", err))
}
service.setTurnPhase(stream, TurnPhaseFailed)
return service.failStream(stream, "unknown", payload.Err)
}
+54 -22
View File
@@ -19,87 +19,119 @@ type usageLookupRecord struct {
CreatedAt time.Time
}
type aiHandler struct {
mux *http.ServeMux
paths map[string]struct{}
}
func newAIHandlerMux() *aiHandler {
return &aiHandler{
mux: http.NewServeMux(),
paths: make(map[string]struct{}),
}
}
func (handler *aiHandler) Handle(pattern string, target http.Handler) {
handler.paths[pattern] = struct{}{}
handler.mux.Handle(pattern, target)
}
func (handler *aiHandler) HandlesPath(path string) bool {
if handler == nil {
return false
}
_, ok := handler.paths[path]
return ok
}
func (handler *aiHandler) ServeHTTP(writer http.ResponseWriter, request *http.Request) {
if handler == nil || handler.mux == nil {
http.NotFound(writer, request)
return
}
handler.mux.ServeHTTP(writer, request)
}
const (
dashboardServiceGetTokenUsageProcedure = "/aiserver.v1.DashboardService/GetTokenUsage"
dashboardServiceGetGlassEarlyPreviewEnrollmentProcedure = "/aiserver.v1.DashboardService/GetGlassEarlyPreviewEnrollment"
)
func newAIHandler(service *Service) http.Handler {
mux := http.NewServeMux()
mux.Handle(
func newAIHandler(service *Service) *aiHandler {
handler := newAIHandlerMux()
handler.Handle(
dashboardServiceGetTokenUsageProcedure,
connect.NewUnaryHandler(dashboardServiceGetTokenUsageProcedure, service.GetTokenUsage),
)
mux.Handle(
handler.Handle(
dashboardServiceGetGlassEarlyPreviewEnrollmentProcedure,
connect.NewUnaryHandler(dashboardServiceGetGlassEarlyPreviewEnrollmentProcedure, service.GetGlassEarlyPreviewEnrollment),
)
mux.Handle(
handler.Handle(
aiserverv1connect.AiServiceCountTokensProcedure,
connect.NewUnaryHandler(aiserverv1connect.AiServiceCountTokensProcedure, service.CountTokens),
)
mux.Handle(
handler.Handle(
aiserverv1connect.AiServiceGetThoughtAnnotationProcedure,
connect.NewUnaryHandler(aiserverv1connect.AiServiceGetThoughtAnnotationProcedure, service.GetThoughtAnnotation),
)
mux.Handle(
handler.Handle(
aiserverv1connect.AiServiceWriteGitCommitMessageProcedure,
connect.NewUnaryHandler(aiserverv1connect.AiServiceWriteGitCommitMessageProcedure, service.WriteGitCommitMessage),
)
mux.Handle(
handler.Handle(
aiserverv1connect.AiServiceCreateExperimentalIndexProcedure,
connect.NewUnaryHandler(aiserverv1connect.AiServiceCreateExperimentalIndexProcedure, service.CreateExperimentalIndex),
)
mux.Handle(
handler.Handle(
aiserverv1connect.AiServiceListExperimentalIndexFilesProcedure,
connect.NewUnaryHandler(aiserverv1connect.AiServiceListExperimentalIndexFilesProcedure, service.ListExperimentalIndexFiles),
)
mux.Handle(
handler.Handle(
aiserverv1connect.AiServiceListenExperimentalIndexProcedure,
connect.NewServerStreamHandler(aiserverv1connect.AiServiceListenExperimentalIndexProcedure, service.ListenExperimentalIndex),
)
mux.Handle(
handler.Handle(
aiserverv1connect.AiServiceRegisterFileToIndexProcedure,
connect.NewUnaryHandler(aiserverv1connect.AiServiceRegisterFileToIndexProcedure, service.RegisterFileToIndex),
)
mux.Handle(
handler.Handle(
aiserverv1connect.AiServiceSetupIndexDependenciesProcedure,
connect.NewUnaryHandler(aiserverv1connect.AiServiceSetupIndexDependenciesProcedure, service.SetupIndexDependencies),
)
mux.Handle(
handler.Handle(
aiserverv1connect.AiServiceComputeIndexTopoSortProcedure,
connect.NewUnaryHandler(aiserverv1connect.AiServiceComputeIndexTopoSortProcedure, service.ComputeIndexTopoSort),
)
mux.Handle(
handler.Handle(
aiserverv1connect.AiServiceDocumentationQueryProcedure,
connect.NewUnaryHandler(aiserverv1connect.AiServiceDocumentationQueryProcedure, service.DocumentationQuery),
)
mux.Handle(
handler.Handle(
aiserverv1connect.AiServiceAvailableDocsProcedure,
connect.NewUnaryHandler(aiserverv1connect.AiServiceAvailableDocsProcedure, service.AvailableDocs),
)
mux.Handle(
handler.Handle(
aiserverv1connect.AiServiceKnowledgeBaseAddProcedure,
connect.NewUnaryHandler(aiserverv1connect.AiServiceKnowledgeBaseAddProcedure, service.KnowledgeBaseAdd),
)
mux.Handle(
handler.Handle(
aiserverv1connect.AiServiceKnowledgeBaseListProcedure,
connect.NewUnaryHandler(aiserverv1connect.AiServiceKnowledgeBaseListProcedure, service.KnowledgeBaseList),
)
mux.Handle(
handler.Handle(
aiserverv1connect.AiServiceKnowledgeBaseRemoveProcedure,
connect.NewUnaryHandler(aiserverv1connect.AiServiceKnowledgeBaseRemoveProcedure, service.KnowledgeBaseRemove),
)
mux.Handle(
handler.Handle(
aiserverv1connect.AiServiceKnowledgeBaseUpdateProcedure,
connect.NewUnaryHandler(aiserverv1connect.AiServiceKnowledgeBaseUpdateProcedure, service.KnowledgeBaseUpdate),
)
mux.Handle(
handler.Handle(
aiserverv1connect.AiServiceFetchRelevantKnowledgeForConversationProcedure,
connect.NewUnaryHandler(aiserverv1connect.AiServiceFetchRelevantKnowledgeForConversationProcedure, service.FetchRelevantKnowledgeForConversation),
)
mux.Handle("/", http.NotFoundHandler())
return mux
return handler
}
func (service *Service) GetThoughtAnnotation(_ context.Context, req *connect.Request[aiserverv1.GetThoughtAnnotationRequest]) (*connect.Response[aiserverv1.GetThoughtAnnotationResponse], error) {
@@ -0,0 +1,20 @@
package forwarder
import (
"testing"
"cursor/gen/aiserverv1/aiserverv1connect"
)
func TestAIHandlerTracksLocallyImplementedPaths(t *testing.T) {
handler := newAIHandler(&Service{})
if !handler.HandlesPath(aiserverv1connect.AiServiceCountTokensProcedure) {
t.Fatalf("expected %q to be handled locally", aiserverv1connect.AiServiceCountTokensProcedure)
}
if !handler.HandlesPath(dashboardServiceGetTokenUsageProcedure) {
t.Fatalf("expected %q to be handled locally", dashboardServiceGetTokenUsageProcedure)
}
if handler.HandlesPath("/aiserver.v1.AiService/UnknownProcedure") {
t.Fatal("unknown AI procedure must fall through to upstream")
}
}
+57 -46
View File
@@ -19,29 +19,15 @@ type pendingCheckpointBlobWrite struct {
blob CheckpointBlob
}
func successfulCheckpointTerminalAction(completion *pendingTurnCompletion) checkpointTerminalAction {
func clonePendingTurnCompletion(completion *pendingTurnCompletion) *pendingTurnCompletion {
if completion == nil {
return checkpointTerminalAction{}
}
return checkpointTerminalAction{
Kind: checkpointTerminalActionComplete,
Completion: *completion,
}
}
func failedCheckpointTerminalAction(errorCode string, errorMessage string) checkpointTerminalAction {
return checkpointTerminalAction{
Kind: checkpointTerminalActionFail,
ErrorCode: strings.TrimSpace(errorCode),
ErrorMessage: strings.TrimSpace(errorMessage),
return nil
}
cloned := *completion
return &cloned
}
func (service *Service) queueCheckpointProjection(stream *ActiveStream, projection *CheckpointProjection, completion *pendingTurnCompletion) error {
return service.queueCheckpointProjectionWithTerminal(stream, projection, successfulCheckpointTerminalAction(completion))
}
func (service *Service) queueCheckpointProjectionWithTerminal(stream *ActiveStream, projection *CheckpointProjection, terminal checkpointTerminalAction) error {
if service == nil || stream == nil || projection == nil || projection.State == nil {
return nil
}
@@ -57,8 +43,8 @@ func (service *Service) queueCheckpointProjectionWithTerminal(stream *ActiveStre
if stream.ConfirmedCheckpointBlobs == nil {
stream.ConfirmedCheckpointBlobs = make(map[string]struct{})
}
if terminal.Kind == checkpointTerminalActionNone && stream.PendingCheckpoint != nil {
terminal = stream.PendingCheckpoint.Terminal
if completion == nil && stream.PendingCheckpoint != nil {
completion = stream.PendingCheckpoint.Completion
}
required := make(map[string]struct{}, len(projection.Blobs))
pendingKeys := make(map[string]struct{}, len(stream.PendingCheckpointBlobWrites))
@@ -88,11 +74,11 @@ func (service *Service) queueCheckpointProjectionWithTerminal(stream *ActiveStre
toWrite = append(toWrite, pendingCheckpointBlobWrite{requestID: requestID, blob: blob})
}
stream.PendingCheckpoint = &pendingCheckpointPublish{
State: state,
Required: required,
Terminal: terminal,
State: state,
Required: required,
Completion: clonePendingTurnCompletion(completion),
}
if terminal.Kind != checkpointTerminalActionNone {
if completion != nil {
stream.Phase = TurnPhaseCheckpointing
}
stream.UpdatedAt = time.Now().UTC()
@@ -108,8 +94,13 @@ func (service *Service) queueCheckpointProjectionWithTerminal(stream *ActiveStre
if service.checkpointProjectionReady(stream) {
return service.publishReadyCheckpoint(stream)
}
// Checkpoints reference these Blob IDs, so the client must confirm every
// required Blob before the checkpoint becomes visible.
// Keep the latest live UI state ahead of an immediate client abort. Blob writes are
// ordered before this snapshot; acknowledgements still gate terminal completion.
if completion == nil {
if err := service.publishPendingCheckpoint(stream); err != nil {
return service.finishAfterCheckpointSyncFailure(stream, fmt.Errorf("publish pending checkpoint: %w", err))
}
}
service.scheduleStreamTimer(
stream,
providerTimerKey(streamTimerCheckpointBlobs, ""),
@@ -122,6 +113,31 @@ func (service *Service) queueCheckpointProjectionWithTerminal(stream *ActiveStre
return nil
}
func (service *Service) publishPendingCheckpoint(stream *ActiveStream) error {
if service == nil || stream == nil {
return nil
}
stream.mu.Lock()
pending := stream.PendingCheckpoint
if pending == nil || pending.Published {
stream.mu.Unlock()
return nil
}
pending.Published = true
state := pending.State
stream.UpdatedAt = time.Now().UTC()
stream.mu.Unlock()
if err := service.broker.Publish(stream.RequestID, StreamEvent{Message: buildCheckpointMessage(state)}); err != nil {
stream.mu.Lock()
if stream.PendingCheckpoint == pending {
pending.Published = false
}
stream.mu.Unlock()
return err
}
return nil
}
func (service *Service) checkpointProjectionReady(stream *ActiveStream) bool {
if stream == nil {
return false
@@ -191,18 +207,24 @@ func (service *Service) publishReadyCheckpoint(stream *ActiveStream) error {
}
stream.PendingCheckpoint = nil
state := pending.State
terminal := pending.Terminal
completion := clonePendingTurnCompletion(pending.Completion)
published := pending.Published
stream.UpdatedAt = time.Now().UTC()
stream.mu.Unlock()
clearStreamTimer(stream, providerTimerKey(streamTimerCheckpointBlobs, ""))
if err := service.broker.Publish(stream.RequestID, StreamEvent{Message: buildCheckpointMessage(state)}); err != nil {
if terminal.Kind != checkpointTerminalActionNone {
log.Printf("forwarder checkpoint publish skipped before terminal request_id=%s err=%v", stream.RequestID, err)
return service.finishCheckpointTerminalAction(stream, terminal)
if !published {
if err := service.broker.Publish(stream.RequestID, StreamEvent{Message: buildCheckpointMessage(state)}); err != nil {
if completion != nil {
log.Printf("forwarder checkpoint publish skipped before successful terminal request_id=%s err=%v", stream.RequestID, err)
return service.finishSuccessfulTurnAfterCheckpoint(stream, *completion)
}
return err
}
return err
}
return service.finishCheckpointTerminalAction(stream, terminal)
if completion != nil {
return service.finishSuccessfulTurnAfterCheckpoint(stream, *completion)
}
return nil
}
func (service *Service) handleCheckpointBlobTimeout(stream *ActiveStream) error {
@@ -229,23 +251,12 @@ func (service *Service) finishAfterCheckpointSyncFailure(stream *ActiveStream, c
if cause != nil {
log.Printf("forwarder checkpoint blob sync skipped request_id=%s conversation_id=%s err=%v", stream.RequestID, stream.ConversationID, cause)
}
if pending != nil {
return service.finishCheckpointTerminalAction(stream, pending.Terminal)
if pending != nil && pending.Completion != nil {
return service.finishSuccessfulTurnAfterCheckpoint(stream, *pending.Completion)
}
return nil
}
func (service *Service) finishCheckpointTerminalAction(stream *ActiveStream, terminal checkpointTerminalAction) error {
switch terminal.Kind {
case checkpointTerminalActionComplete:
return service.finishSuccessfulTurnAfterCheckpoint(stream, terminal.Completion)
case checkpointTerminalActionFail:
return service.finishFailedTurnAfterCheckpoint(stream, terminal.ErrorCode, terminal.ErrorMessage)
default:
return nil
}
}
func (service *Service) discardPendingCheckpoint(stream *ActiveStream, reason string) {
if stream == nil {
return
@@ -8,43 +8,22 @@ import (
"cursor/gen/agentv1"
)
func TestCheckpointBlobSyncWaitsForAcknowledgementsBeforePublishingNonTerminalCheckpoint(t *testing.T) {
func TestCheckpointBlobSyncPublishesNonTerminalCheckpointBeforeAcknowledgements(t *testing.T) {
service, stream, projection := testCheckpointBlobProjection(t)
if err := service.queueCheckpointProjection(stream, projection, nil); err != nil {
t.Fatalf("queueCheckpointProjection() error = %v", err)
}
events := readCheckpointTestEvents(t, service, stream)
if len(events) != len(projection.Blobs) {
t.Fatalf("events before ACK = %d, want %d Blob writes", len(events), len(projection.Blobs))
if len(events) != len(projection.Blobs)+1 {
t.Fatalf("events before ACK = %d, want %d Blob writes and one checkpoint", len(events), len(projection.Blobs))
}
for _, event := range events {
for _, event := range events[:len(projection.Blobs)] {
if event.Message.GetKvServerMessage().GetSetBlobArgs() == nil {
t.Fatalf("event before ACK = %#v, want set_blob_args", event.Message)
}
}
stream.mu.Lock()
var firstRequestID uint32
for requestID := range stream.PendingCheckpointBlobWrites {
firstRequestID = requestID
break
}
stream.mu.Unlock()
if firstRequestID == 0 {
t.Fatal("checkpoint projection has no pending Blob writes")
}
if err := service.handleCheckpointBlobResult(stream, &agentv1.KvClientMessage{
Id: firstRequestID,
Message: &agentv1.KvClientMessage_SetBlobResult{
SetBlobResult: &agentv1.SetBlobResult{},
},
}); err != nil {
t.Fatalf("first Blob ACK error = %v", err)
}
for _, event := range readCheckpointTestEvents(t, service, stream) {
if event.Message.GetConversationCheckpointUpdate() != nil {
t.Fatal("checkpoint published after only a partial Blob acknowledgement")
}
if checkpoint := events[len(events)-1].Message.GetConversationCheckpointUpdate(); checkpoint == nil || len(checkpoint.GetTurns()) != 1 {
t.Fatalf("last event before ACK = %#v, want one Blob-backed turn", events[len(events)-1].Message)
}
acknowledgeCheckpointBlobs(t, service, stream)
@@ -119,123 +98,7 @@ func TestCheckpointBlobTimeoutDoesNotFailSuccessfulTurn(t *testing.T) {
}
}
func TestCheckpointBlobSyncPublishesCheckpointBeforeFailedTerminal(t *testing.T) {
service, stream, _ := testCheckpointBlobProjection(t)
if err := service.failActiveStream(
stream,
stream.ConversationID,
stream.RequestID,
"model-call-1",
"provider_error",
"provider failed",
); err != nil {
t.Fatalf("failActiveStream() error = %v", err)
}
for _, event := range readCheckpointTestEvents(t, service, stream) {
if event.Message.GetConversationCheckpointUpdate() != nil || event.End {
t.Fatalf("event before ACK = %#v, want only Blob writes", event)
}
}
stream.mu.Lock()
phaseBeforeACK := stream.Phase
statusBeforeACK := stream.Status
stream.mu.Unlock()
if phaseBeforeACK != TurnPhaseCheckpointing || isTerminalStreamStatus(statusBeforeACK) {
t.Fatalf("before ACK phase=%s status=%s, want checkpointing and non-terminal", phaseBeforeACK, statusBeforeACK)
}
acknowledgeCheckpointBlobs(t, service, stream)
events := readCheckpointTestEvents(t, service, stream)
checkpointIndex, endIndex := -1, -1
for index, event := range events {
switch {
case event.Message.GetConversationCheckpointUpdate() != nil:
checkpointIndex = index
case event.End:
endIndex = index
if event.TerminalErrorCode != "provider_error" || event.TerminalErrorMessage != "provider failed" {
t.Fatalf("terminal event = %#v, want provider error", event)
}
}
}
if checkpointIndex < 0 || endIndex <= checkpointIndex {
t.Fatalf("terminal order checkpoint=%d end=%d", checkpointIndex, endIndex)
}
stream.mu.Lock()
phaseAfterACK := stream.Phase
statusAfterACK := stream.Status
stream.mu.Unlock()
if phaseAfterACK != TurnPhaseFailed || statusAfterACK != StreamStatusFailed {
t.Fatalf("after ACK phase=%s status=%s, want failed", phaseAfterACK, statusAfterACK)
}
}
func TestCheckpointBlobTimeoutStillPublishesFailedTerminal(t *testing.T) {
service, stream, _ := testCheckpointBlobProjection(t)
if err := service.failActiveStream(
stream,
stream.ConversationID,
stream.RequestID,
"model-call-1",
"provider_error",
"provider failed",
); err != nil {
t.Fatalf("failActiveStream() error = %v", err)
}
if err := service.handleCheckpointBlobTimeout(stream); err != nil {
t.Fatalf("handleCheckpointBlobTimeout() error = %v", err)
}
events := readCheckpointTestEvents(t, service, stream)
var checkpoint, failedEnd bool
for _, event := range events {
checkpoint = checkpoint || event.Message.GetConversationCheckpointUpdate() != nil
failedEnd = failedEnd || event.End && event.TerminalErrorCode == "provider_error" && event.TerminalErrorMessage == "provider failed"
}
if checkpoint || !failedEnd {
t.Fatalf("timeout events checkpoint=%v failed_end=%v", checkpoint, failedEnd)
}
}
func TestManualCompactionNoopWaitsForCheckpointBeforeTerminal(t *testing.T) {
service, stream, _ := testCheckpointBlobProjection(t)
conversation, _, _, err := service.snapshotCheckpointConversation(stream)
if err != nil {
t.Fatalf("snapshotCheckpointConversation() error = %v", err)
}
if _, err := service.store.SaveConversationWithEntries(stream.ConversationID, conversation, conversation.Entries); err != nil {
t.Fatalf("SaveConversationWithEntries() error = %v", err)
}
if err := service.finishManualCompactionNoop(stream); err != nil {
t.Fatalf("finishManualCompactionNoop() error = %v", err)
}
for _, event := range readCheckpointTestEvents(t, service, stream) {
if event.Message.GetInteractionUpdate().GetTurnEnded() != nil || event.End {
t.Fatalf("terminal event before checkpoint Blob ACK = %#v", event)
}
}
acknowledgeCheckpointBlobs(t, service, stream)
events := readCheckpointTestEvents(t, service, stream)
checkpointIndex, turnEndedIndex, endIndex := -1, -1, -1
for index, event := range events {
switch {
case event.Message.GetConversationCheckpointUpdate() != nil:
checkpointIndex = index
case event.Message.GetInteractionUpdate().GetTurnEnded() != nil:
turnEndedIndex = index
case event.End:
endIndex = index
}
}
if checkpointIndex < 0 || turnEndedIndex <= checkpointIndex || endIndex <= turnEndedIndex {
t.Fatalf("terminal order checkpoint=%d turn_ended=%d end=%d", checkpointIndex, turnEndedIndex, endIndex)
}
}
func TestCancellationDiscardsUnpublishedCheckpointAndIgnoresLateAcknowledgements(t *testing.T) {
func TestCancellationKeepsPublishedCheckpointAndIgnoresLateAcknowledgements(t *testing.T) {
service, stream, projection := testCheckpointBlobProjection(t)
if err := service.queueCheckpointProjection(stream, projection, nil); err != nil {
t.Fatalf("queueCheckpointProjection() error = %v", err)
@@ -247,8 +110,8 @@ func TestCancellationDiscardsUnpublishedCheckpointAndIgnoresLateAcknowledgements
checkpointBeforeCancel++
}
}
if checkpointBeforeCancel != 0 {
t.Fatalf("checkpoints before cancel = %d, want 0", checkpointBeforeCancel)
if checkpointBeforeCancel != 1 {
t.Fatalf("checkpoints before cancel = %d, want 1", checkpointBeforeCancel)
}
stream.mu.Lock()
requestIDs := make([]uint32, 0, len(stream.PendingCheckpointBlobWrites))
@@ -286,7 +149,7 @@ func TestCancellationDiscardsUnpublishedCheckpointAndIgnoresLateAcknowledgements
stream.mu.Lock()
pending := stream.PendingCheckpoint
stream.mu.Unlock()
if checkpointCount != 0 || !canceledEnd || pending != nil {
if checkpointCount != 1 || !canceledEnd || pending != nil {
t.Fatalf("cancel events checkpoints=%d canceled_end=%v pending=%v", checkpointCount, canceledEnd, pending != nil)
}
}
+61 -26
View File
@@ -234,7 +234,7 @@ func (service *Service) buildLegacyCompactionPlan(base *compactionPlan, conversa
if conversation == nil || base == nil {
return nil, nil
}
candidates := buildContextCompactionCandidates(replayablePromptProjectionEntries(conversation.Entries), base.CurrentTurnSeq, base.CurrentRequestID)
candidates := buildContextCompactionCandidates(checkpointProjectionEntries(conversation.Entries), base.CurrentTurnSeq, base.CurrentRequestID)
if len(candidates) == 0 {
return nil, nil
}
@@ -260,7 +260,7 @@ func (service *Service) buildAutoCompactionPlanFromHistory(base *compactionPlan,
if err != nil {
return nil, err
}
currentCandidate, hasCurrentCandidate := buildCurrentTurnCompactionCandidate(replayablePromptProjectionEntries(conversation.Entries), base.CurrentTurnSeq, base.CurrentRequestID)
currentCandidate, hasCurrentCandidate := buildCurrentTurnCompactionCandidate(checkpointProjectionEntries(conversation.Entries), base.CurrentTurnSeq, base.CurrentRequestID)
if !hasCurrentCandidate {
return legacyPlan, nil
}
@@ -447,8 +447,16 @@ func (service *Service) handleCompactionEvent(stream *ActiveStream, payload *str
if err := service.completeManualCompactionTurn(stream); err != nil {
return service.failStream(stream, "unknown", err)
}
completion := manualCompactionTurnCompletion(stream)
return service.publishCheckpointWithCompletion(stream.RequestID, stream.ConversationID, &completion)
if err := service.broker.Publish(stream.RequestID, StreamEvent{
Message: buildTurnEndedMessage(0, 0, 0, 0),
}); err != nil {
return service.failStream(stream, "unknown", err)
}
if err := service.broker.Complete(stream.RequestID, "", ""); err != nil {
return service.failStream(stream, "unknown", err)
}
service.setTurnPhase(stream, TurnPhaseCompleted)
return nil
}
return service.requestProviderAction(stream, providerActionResume)
}
@@ -492,8 +500,12 @@ func (service *Service) finishManualCompactionNoop(stream *ActiveStream) error {
if err := service.completeManualCompactionTurn(stream); err != nil {
return err
}
completion := manualCompactionTurnCompletion(stream)
return service.publishCheckpointWithCompletion(stream.RequestID, stream.ConversationID, &completion)
if err := service.broker.Publish(stream.RequestID, StreamEvent{
Message: buildTurnEndedMessage(0, 0, 0, 0),
}); err != nil {
return err
}
return service.broker.Complete(stream.RequestID, "", "")
}
func (service *Service) completeManualCompactionTurn(stream *ActiveStream) error {
@@ -518,21 +530,10 @@ func (service *Service) completeManualCompactionTurn(stream *ActiveStream) error
if err := service.syncSummaryCarryForward(conversationID, requestID, modelCallID); err != nil {
return err
}
service.setTurnPhase(stream, TurnPhaseCompleted)
return nil
}
func manualCompactionTurnCompletion(stream *ActiveStream) pendingTurnCompletion {
if stream == nil {
return pendingTurnCompletion{}
}
return pendingTurnCompletion{
ConversationID: strings.TrimSpace(stream.ConversationID),
RequestID: strings.TrimSpace(stream.RequestID),
TurnSeq: stream.TurnSeq,
ModelCallID: "turn:" + strings.TrimSpace(stream.RequestID),
}
}
func (service *Service) publishSummaryCompleted(stream *ActiveStream, hookMessage string) error {
if service == nil || stream == nil {
return nil
@@ -567,7 +568,6 @@ func (service *Service) applyCompactionPlan(stream *ActiveStream, conversationID
if err != nil {
return err
}
originalEntryCount := len(candidateConversation.Entries)
if err := applyCompactionToConversation(candidateConversation, plan, summaryText); err != nil {
return err
}
@@ -582,9 +582,9 @@ func (service *Service) applyCompactionPlan(stream *ActiveStream, conversationID
if validationErr := validateCompactionCandidateBudget(recompiled, plan); validationErr != nil {
return validationErr
}
compactionEntries := append([]HistoryEntry(nil), candidateConversation.Entries[originalEntryCount:]...)
replacementEntries := append([]HistoryEntry(nil), candidateConversation.Entries...)
if service.store != nil {
persisted, _, err := service.store.AppendEntriesWithUpdate(conversationID, resetEntrySequences(compactionEntries), func(item *ConversationFile) error {
persisted, err := service.store.ReplaceEntries(conversationID, replacementEntries, func(item *ConversationFile) error {
if item == nil {
return nil
}
@@ -605,7 +605,10 @@ func (service *Service) applyCompactionPlan(stream *ActiveStream, conversationID
if item == nil {
return nil
}
appendEntriesInPlace(item, resetEntrySequences(compactionEntries))
item.Entries = nil
item.NextEntrySeq = 1
item.NextTurnSeq = 1
appendEntriesInPlace(item, resetEntrySequences(replacementEntries))
item.TokenDetailsUsedTokens = 0
clearConversationAutoCompactionState(item)
return nil
@@ -640,13 +643,14 @@ func applyCompactionToConversation(conversation *ConversationFile, plan *Pending
if conversation == nil || plan == nil {
return nil
}
compactionEntries, err := buildCompactedContextEntries(conversation, plan, summaryText)
replacementEntries, err := buildCompactedContextEntries(conversation, plan, summaryText)
if err != nil {
return err
}
// Canonical history stays append-only. The prompt projector applies the
// latest summary marker when constructing model-visible replay.
appendEntriesInPlace(conversation, resetEntrySequences(compactionEntries))
conversation.Entries = nil
conversation.NextEntrySeq = 1
conversation.NextTurnSeq = 1
appendEntriesInPlace(conversation, resetEntrySequences(replacementEntries))
conversation.TokenDetailsUsedTokens = 0
clearConversationAutoCompactionState(conversation)
if conversation.TokenDetailsMaxTokens == 0 {
@@ -667,9 +671,40 @@ func buildCompactedContextEntries(conversation *ConversationFile, plan *PendingC
if ok {
entries = append(entries, runtimeEntry)
}
if conversation == nil || !plan.PreserveCurrentTurnInputs {
return entries, nil
}
entries = append(entries, buildAutoCompactionPreservedCurrentTurnEntries(conversation.Entries, plan)...)
return entries, nil
}
func buildAutoCompactionPreservedCurrentTurnEntries(entries []HistoryEntry, plan *PendingCompaction) []HistoryEntry {
if len(entries) == 0 || plan == nil || !plan.PreserveCurrentTurnInputs {
return nil
}
latestToolCallID := latestCompletedToolCallIDForTurn(entries, plan.CurrentTurnSeq, plan.CurrentRequestID)
preservedIndexes := autoCompactionPreservedEntryIndexes(entries, plan.CurrentTurnSeq, plan.CurrentRequestID, latestToolCallID)
if len(preservedIndexes) == 0 {
return nil
}
preserved := make([]HistoryEntry, 0, len(preservedIndexes))
for index, entry := range entries {
if _, ok := preservedIndexes[index]; !ok {
continue
}
switch strings.TrimSpace(entry.Kind) {
case "compaction_summary", "compacted_summary", "compaction_request":
continue
case "tool_result":
if rewritten, ok := rewriteAutoCompactionToolResultEntry(entry, autoCompactionPreservedToolResultLimitBytes, false); ok {
entry = rewritten
}
}
preserved = append(preserved, entry)
}
return preserved
}
func newCompactionSummaryEntry(plan *PendingCompaction, summaryText string) HistoryEntry {
payload, _ := json.Marshal(compactionSummaryEntryPayload{
Summary: strings.TrimSpace(summaryText),
@@ -1,206 +0,0 @@
package forwarder
import (
"reflect"
"strings"
"testing"
"google.golang.org/protobuf/encoding/protojson"
"cursor/gen/agentv1"
)
func TestApplyCompactionToConversationPreservesCanonicalHistory(t *testing.T) {
conversation := compactionAppendOnlyConversation(t)
originalEntries := append([]HistoryEntry(nil), conversation.Entries...)
plan := &PendingCompaction{
Trigger: "manual",
CurrentTurnSeq: 2,
CurrentRequestID: "request-2",
}
if err := applyCompactionToConversation(conversation, plan, "earlier context summary"); err != nil {
t.Fatalf("applyCompactionToConversation() error = %v", err)
}
if len(conversation.Entries) <= len(originalEntries) {
t.Fatalf("entries after compaction = %d, want the %d original entries plus a summary marker", len(conversation.Entries), len(originalEntries))
}
if !reflect.DeepEqual(conversation.Entries[:len(originalEntries)], originalEntries) {
t.Fatal("compaction changed the canonical history prefix")
}
projector := NewHistoryProjector()
projection, err := projector.ProjectCheckpointProjection(conversation)
if err != nil {
t.Fatalf("ProjectCheckpointProjection() error = %v", err)
}
if len(projection.State.GetTurns()) != 2 {
t.Fatalf("checkpoint turns after compaction = %d, want 2 visible turns", len(projection.State.GetTurns()))
}
replay, err := projector.ProjectPromptReplay(conversation)
if err != nil {
t.Fatalf("ProjectPromptReplay() error = %v", err)
}
if len(replay) != 1 || replay[0].Role != "user" || !strings.Contains(replay[0].Content, "earlier context summary") {
t.Fatalf("prompt replay after compaction = %#v, want only the compacted summary", replay)
}
}
func TestCompactedPromptProjectionPlacesSummaryBeforePreservedCurrentTurn(t *testing.T) {
conversation := &ConversationFile{
ConversationID: "conversation-1",
RootConversationID: "conversation-1",
Mode: "agent",
NextTurnSeq: 1,
NextEntrySeq: 1,
}
appendEntriesInPlace(conversation, []HistoryEntry{
compactionTestUserEntry(t, 1, "request-1", "current question", "message-1"),
newToolCallEntry(1, "request-1", "call-1", "Read", "", "", checkpointTestReadToolCall(t, nil)),
newToolResultEntry(1, "request-1", "call-1", "Read", `{"path":"/tmp/example.txt"}`, "file contents", "", checkpointTestReadToolCall(t, nil)),
})
plan := &PendingCompaction{
Trigger: "auto",
CurrentTurnSeq: 1,
CurrentRequestID: "request-1",
PreserveCurrentTurnInputs: true,
}
if err := applyCompactionToConversation(conversation, plan, "current progress summary"); err != nil {
t.Fatalf("applyCompactionToConversation() error = %v", err)
}
projected := compactedPromptProjectionEntries(conversation.Entries)
promptKinds := make([]string, 0, len(projected))
for _, entry := range projected {
if isPromptReplayEntryKind(entry.Kind) {
promptKinds = append(promptKinds, entry.Kind)
}
}
want := []string{"compacted_summary", "user_message", "tool_call", "tool_result"}
if !reflect.DeepEqual(promptKinds, want) {
t.Fatalf("compacted prompt entry order = %#v, want %#v", promptKinds, want)
}
}
func TestCompactionPlanningDoesNotRecompactArchivedHistory(t *testing.T) {
conversation := compactionAppendOnlyConversation(t)
if err := applyCompactionToConversation(conversation, &PendingCompaction{
Trigger: "manual",
CurrentTurnSeq: 2,
CurrentRequestID: "request-2",
}, "archived history summary"); err != nil {
t.Fatalf("applyCompactionToConversation() error = %v", err)
}
appendEntriesInPlace(conversation, []HistoryEntry{
compactionTestUserEntry(t, 3, "request-3", "new question", "message-3"),
})
plan, err := (&Service{}).buildLegacyCompactionPlan(&compactionPlan{
CurrentTurnSeq: 3,
CurrentRequestID: "request-3",
}, conversation, false, 0)
if err != nil {
t.Fatalf("buildLegacyCompactionPlan() error = %v", err)
}
if plan != nil {
t.Fatalf("buildLegacyCompactionPlan() = %#v, want no already summarized candidates", plan)
}
}
func TestApplyCompactionPlanPersistsHistoryAppendOnly(t *testing.T) {
store := NewConversationFileStore(t.TempDir())
conversation := compactionAppendOnlyConversation(t)
if _, _, err := store.AppendEntries(conversation.ConversationID, resetEntrySequences(conversation.Entries)); err != nil {
t.Fatalf("AppendEntries() error = %v", err)
}
persisted, err := store.LoadConversation(conversation.ConversationID)
if err != nil {
t.Fatalf("initial LoadConversation() error = %v", err)
}
originalEntries := append([]HistoryEntry(nil), persisted.Entries...)
projector := NewHistoryProjector()
service := &Service{
store: store,
projector: projector,
compiler: compactionProjectionCompiler{projector: projector},
}
stream := &ActiveStream{
RequestID: "request-2",
ConversationID: conversation.ConversationID,
TurnSeq: 2,
Mode: agentv1.AgentMode_AGENT_MODE_AGENT,
CheckpointConversation: persisted,
}
plan := &PendingCompaction{
Trigger: "manual",
CurrentTurnSeq: 2,
CurrentRequestID: "request-2",
ContextWindowSize: 1_000_000,
}
if err := service.applyCompactionPlan(stream, conversation.ConversationID, plan, "persisted summary"); err != nil {
t.Fatalf("applyCompactionPlan() error = %v", err)
}
loaded, err := store.LoadConversation(conversation.ConversationID)
if err != nil {
t.Fatalf("LoadConversation() error = %v", err)
}
if len(loaded.Entries) <= len(originalEntries) {
t.Fatalf("persisted entries after compaction = %d, want more than %d", len(loaded.Entries), len(originalEntries))
}
for index := range originalEntries {
if !reflect.DeepEqual(loaded.Entries[index], originalEntries[index]) {
t.Fatalf("persisted history entry %d changed after compaction:\ngot %#v\nwant %#v", index, loaded.Entries[index], originalEntries[index])
}
}
}
type compactionProjectionCompiler struct {
projector *HistoryProjector
}
func (compiler compactionProjectionCompiler) Compile(conversation *ConversationFile, _ agentv1.AgentMode, _ string, _ string) (CompiledConversation, error) {
messages, err := compiler.projector.ProjectPromptReplay(conversation)
return CompiledConversation{Messages: messages}, err
}
func (compactionProjectionCompiler) DerivePromptContexts(*ConversationFile, agentv1.AgentMode, string) ([]PromptContextMessage, error) {
return nil, nil
}
func compactionAppendOnlyConversation(t *testing.T) *ConversationFile {
t.Helper()
conversation := &ConversationFile{
ConversationID: "conversation-1",
RootConversationID: "conversation-1",
Mode: "agent",
NextTurnSeq: 1,
NextEntrySeq: 1,
TokenDetailsUsedTokens: 42_000,
TokenDetailsMaxTokens: 50_000,
}
appendEntriesInPlace(conversation, []HistoryEntry{
compactionTestUserEntry(t, 1, "request-1", "first question", "message-1"),
newAssistantTextEntry(1, "request-1", "first answer", "", ""),
compactionTestUserEntry(t, 2, "request-2", "second question", "message-2"),
newAssistantTextEntry(2, "request-2", "second answer", "", ""),
})
return conversation
}
func compactionTestUserEntry(t *testing.T, turnSeq int64, requestID string, text string, messageID string) HistoryEntry {
t.Helper()
payload, err := protojson.Marshal(&agentv1.UserMessage{Text: text, MessageId: messageID})
if err != nil {
t.Fatalf("marshal user message: %v", err)
}
return HistoryEntry{
TurnSeq: turnSeq,
RequestID: requestID,
Role: "user",
Kind: "user_message",
Payload: payload,
}
}
var _ PromptCompiler = compactionProjectionCompiler{}
+2 -11
View File
@@ -20,21 +20,16 @@ type DefaultPromptCompiler struct {
catalog ToolCatalog
reminders ReminderInjector
rules *UserRuleStore
blobs contentBlobReader
}
// NewPromptCompiler 创建默认 prompt 编译器。
func NewPromptCompiler(projector *HistoryProjector, catalog ToolCatalog, reminders ReminderInjector, rules *UserRuleStore, blobReaders ...contentBlobReader) *DefaultPromptCompiler {
compiler := &DefaultPromptCompiler{
func NewPromptCompiler(projector *HistoryProjector, catalog ToolCatalog, reminders ReminderInjector, rules *UserRuleStore) *DefaultPromptCompiler {
return &DefaultPromptCompiler{
projector: projector,
catalog: catalog,
reminders: reminders,
rules: rules,
}
if len(blobReaders) > 0 {
compiler.blobs = blobReaders[0]
}
return compiler
}
// Compile 生成当前 turn 应发送给 provider 的消息和工具集合。
@@ -91,10 +86,6 @@ func (compiler *DefaultPromptCompiler) Compile(conversation *ConversationFile, m
if err != nil {
return CompiledConversation{}, err
}
replayMessages, err = enrichProviderReadImages(replayMessages, conversation, compiler.blobs)
if err != nil {
return CompiledConversation{}, err
}
messages = append(messages, replayMessages...)
return CompiledConversation{
Mode: normalizedMode,
@@ -1,111 +0,0 @@
// content_blob_store.go 负责持久化 history 引用的内容寻址二进制数据。
package forwarder
import (
"bytes"
"crypto/sha256"
"encoding/hex"
"fmt"
"os"
"path/filepath"
"strings"
)
const contentBlobDirectoryName = ".blobs"
// ContentBlobStore 使用 SHA-256 内容哈希保存不可变二进制数据。
type ContentBlobStore struct {
root string
}
// NewContentBlobStore 创建独立于 context.json 和 checkpoint 的内容寻址存储。
func NewContentBlobStore(historyRoot string) *ContentBlobStore {
historyRoot = strings.TrimSpace(historyRoot)
if historyRoot == "" {
return &ContentBlobStore{}
}
return &ContentBlobStore{root: filepath.Join(historyRoot, contentBlobDirectoryName, "sha256")}
}
// Put 校验内容哈希并幂等保存数据。
func (store *ContentBlobStore) Put(id []byte, data []byte) error {
if store == nil || strings.TrimSpace(store.root) == "" {
return fmt.Errorf("content blob store is not initialized")
}
normalizedID, err := normalizeContentBlobID(id)
if err != nil {
return err
}
digest := sha256.Sum256(data)
if !bytes.Equal(normalizedID, digest[:]) {
return fmt.Errorf("content blob id does not match payload sha256")
}
path := store.blobPath(normalizedID)
if existing, err := store.Get(normalizedID); err == nil {
if bytes.Equal(existing, data) {
return nil
}
return fmt.Errorf("content blob payload conflicts with existing id")
} else if !os.IsNotExist(err) {
return err
}
if err := os.MkdirAll(store.root, 0o700); err != nil {
return fmt.Errorf("create content blob directory: %w", err)
}
temporary, err := os.CreateTemp(store.root, ".blob-*")
if err != nil {
return fmt.Errorf("create content blob temporary file: %w", err)
}
temporaryPath := temporary.Name()
defer os.Remove(temporaryPath)
if err := temporary.Chmod(0o600); err != nil {
_ = temporary.Close()
return fmt.Errorf("set content blob permissions: %w", err)
}
if _, err := temporary.Write(data); err != nil {
_ = temporary.Close()
return fmt.Errorf("write content blob: %w", err)
}
if err := temporary.Sync(); err != nil {
_ = temporary.Close()
return fmt.Errorf("sync content blob: %w", err)
}
if err := temporary.Close(); err != nil {
return fmt.Errorf("close content blob: %w", err)
}
if err := os.Rename(temporaryPath, path); err != nil {
return fmt.Errorf("commit content blob: %w", err)
}
return nil
}
// Get 读取内容并再次校验哈希,避免损坏数据进入模型请求。
func (store *ContentBlobStore) Get(id []byte) ([]byte, error) {
if store == nil || strings.TrimSpace(store.root) == "" {
return nil, fmt.Errorf("content blob store is not initialized")
}
normalizedID, err := normalizeContentBlobID(id)
if err != nil {
return nil, err
}
data, err := os.ReadFile(store.blobPath(normalizedID))
if err != nil {
return nil, err
}
digest := sha256.Sum256(data)
if !bytes.Equal(normalizedID, digest[:]) {
return nil, fmt.Errorf("content blob sha256 verification failed")
}
return append([]byte(nil), data...), nil
}
func (store *ContentBlobStore) blobPath(id []byte) string {
return filepath.Join(store.root, hex.EncodeToString(id))
}
func normalizeContentBlobID(id []byte) ([]byte, error) {
if len(id) != sha256.Size {
return nil, fmt.Errorf("content blob id must be %d bytes", sha256.Size)
}
return append([]byte(nil), id...), nil
}
@@ -1,41 +0,0 @@
package forwarder
import (
"bytes"
"crypto/sha256"
"testing"
)
func TestContentBlobStorePutGetIsIdempotent(t *testing.T) {
store := NewContentBlobStore(t.TempDir())
data := []byte("stable blob bytes")
id := sha256.Sum256(data)
if err := store.Put(id[:], data); err != nil {
t.Fatalf("first Put() error = %v", err)
}
if err := store.Put(id[:], append([]byte(nil), data...)); err != nil {
t.Fatalf("second Put() error = %v", err)
}
got, err := store.Get(id[:])
if err != nil {
t.Fatalf("Get() error = %v", err)
}
if !bytes.Equal(got, data) {
t.Fatalf("Get() = %q, want %q", got, data)
}
got[0] ^= 0xff
again, err := store.Get(id[:])
if err != nil {
t.Fatalf("second Get() error = %v", err)
}
if !bytes.Equal(again, data) {
t.Fatalf("stored data was mutated: %q", again)
}
}
func TestContentBlobStoreRejectsMismatchedID(t *testing.T) {
store := NewContentBlobStore(t.TempDir())
if err := store.Put(bytes.Repeat([]byte{0xff}, sha256.Size), []byte("payload")); err == nil {
t.Fatal("Put() accepted mismatched content id")
}
}
+1 -13
View File
@@ -121,15 +121,10 @@ func (store *ConversationFileStore) LoadConversation(conversationID string) (*Co
// AppendEntries 把已经发生的语义事件追加到 context.json,并同步 state.json。
func (store *ConversationFileStore) AppendEntries(conversationID string, entries []HistoryEntry) (*ConversationFile, []HistoryEntry, error) {
return store.AppendEntriesWithUpdate(conversationID, entries, nil)
}
// AppendEntriesWithUpdate 原子追加 context entries,并在同一把会话锁内更新 state metadata。
func (store *ConversationFileStore) AppendEntriesWithUpdate(conversationID string, entries []HistoryEntry, update func(*ConversationFile) error) (*ConversationFile, []HistoryEntry, error) {
if store == nil {
return nil, nil, fmt.Errorf("conversation file store is nil")
}
if len(entries) == 0 && update == nil {
if len(entries) == 0 {
conversation, err := store.LoadConversation(conversationID)
return conversation, nil, err
}
@@ -167,11 +162,6 @@ func (store *ConversationFileStore) AppendEntriesWithUpdate(conversationID strin
conversation.Mode = alias
}
assigned := appendEntriesInPlace(conversation, entries)
if update != nil {
if err := update(conversation); err != nil {
return nil, nil, err
}
}
deriveConversationLoopState(conversation)
if err := store.writeConversationLocked(normalizedConversationID, conversation); err != nil {
return nil, nil, err
@@ -772,7 +762,6 @@ func mergeConversationMetadata(target *ConversationFile, source *ConversationFil
target.CurrentPlanText = source.CurrentPlanText
target.CurrentPlans = clonePlanRegistryEntries(source.CurrentPlans)
target.CurrentTodos = cloneTodoItems(source.CurrentTodos)
target.ImportedTurnIDs = cloneByteSlices(source.ImportedTurnIDs)
target.LatestRequestPrefix = cloneConversationRequestPrefix(source.LatestRequestPrefix)
target.LastProviderCall = cloneConversationProviderCall(source.LastProviderCall)
if !source.CreatedAt.IsZero() && (target.CreatedAt.IsZero() || source.CreatedAt.Before(target.CreatedAt)) {
@@ -905,7 +894,6 @@ func cloneConversationFile(conversation *ConversationFile) *ConversationFile {
cloned := *conversation
cloned.CurrentPlans = clonePlanRegistryEntries(conversation.CurrentPlans)
cloned.CurrentTodos = cloneTodoItems(conversation.CurrentTodos)
cloned.ImportedTurnIDs = cloneByteSlices(conversation.ImportedTurnIDs)
cloned.LatestRequestPrefix = cloneConversationRequestPrefix(conversation.LatestRequestPrefix)
cloned.LastProviderCall = cloneConversationProviderCall(conversation.LastProviderCall)
cloned.Entries = append([]HistoryEntry(nil), conversation.Entries...)
@@ -1,167 +0,0 @@
package forwarder
import (
"crypto/sha256"
"fmt"
"google.golang.org/protobuf/proto"
"cursor/gen/agentv1"
modeladapter "cursor/internal/backend/agent/model"
promptengine "cursor/internal/backend/agent/prompt"
)
type importedBlobStore map[string][]byte
func newImportedBlobStore(items []*agentv1.PreFetchedBlob) (importedBlobStore, error) {
if len(items) == 0 {
return nil, nil
}
store := make(importedBlobStore, len(items))
for _, item := range items {
if item == nil || len(item.GetId()) == 0 {
continue
}
if len(item.GetId()) != sha256.Size {
return nil, fmt.Errorf("prefetched blob id length %d, want %d", len(item.GetId()), sha256.Size)
}
digest := sha256.Sum256(item.GetValue())
if string(digest[:]) != string(item.GetId()) {
return nil, fmt.Errorf("prefetched blob %x failed SHA-256 validation", item.GetId())
}
store[string(item.GetId())] = append([]byte(nil), item.GetValue()...)
}
return store, nil
}
func (store importedBlobStore) resolve(id []byte) ([]byte, bool) {
if len(id) == 0 || len(store) == 0 {
return nil, false
}
value, ok := store[string(id)]
return append([]byte(nil), value...), ok
}
func decodeImportedTurn(raw []byte, blobs importedBlobStore) (*agentv1.ConversationTurnStructure, []byte, error) {
if data, ok := blobs.resolve(raw); ok {
turn := &agentv1.ConversationTurnStructure{}
if err := proto.Unmarshal(data, turn); err != nil || turn.GetTurn() == nil {
return nil, nil, fmt.Errorf("decode imported turn blob %x: %w", raw, firstNonNilError(err, fmt.Errorf("turn payload is empty")))
}
return turn, append([]byte(nil), raw...), nil
}
turn := &agentv1.ConversationTurnStructure{}
if err := proto.Unmarshal(raw, turn); err == nil && turn.GetTurn() != nil {
return turn, nil, nil
}
if len(raw) == sha256.Size {
return nil, append([]byte(nil), raw...), nil
}
return nil, nil, fmt.Errorf("decode imported inline turn")
}
func decodeImportedUserMessage(raw []byte, blobs importedBlobStore) (*agentv1.UserMessage, error) {
data := raw
if resolved, ok := blobs.resolve(raw); ok {
data = resolved
} else if len(raw) == sha256.Size {
candidate := &agentv1.UserMessage{}
if err := proto.Unmarshal(raw, candidate); err != nil || !hasKnownUserMessageContent(candidate) {
return nil, fmt.Errorf("missing prefetched user message blob %x", raw)
}
return candidate, nil
}
message := &agentv1.UserMessage{}
if err := proto.Unmarshal(data, message); err != nil {
return nil, fmt.Errorf("decode imported turn user_message: %w", err)
}
return message, nil
}
func decodeImportedStep(raw []byte, blobs importedBlobStore) (*agentv1.ConversationStep, error) {
data := raw
if resolved, ok := blobs.resolve(raw); ok {
data = resolved
} else if len(raw) == sha256.Size {
candidate := &agentv1.ConversationStep{}
if err := proto.Unmarshal(raw, candidate); err != nil || candidate.GetMessage() == nil {
return nil, fmt.Errorf("missing prefetched conversation step blob %x", raw)
}
return candidate, nil
}
step := &agentv1.ConversationStep{}
if err := proto.Unmarshal(data, step); err != nil {
return nil, fmt.Errorf("decode imported turn step: %w", err)
}
if step.GetMessage() == nil {
return nil, fmt.Errorf("decode imported turn step: payload is empty")
}
return step, nil
}
func importedBlobTurnMessages(turn *agentv1.ConversationTurnStructure, blobs importedBlobStore) ([]modeladapter.Message, error) {
if turn == nil || turn.GetAgentConversationTurn() == nil {
return nil, nil
}
agentTurn := turn.GetAgentConversationTurn()
messages := make([]modeladapter.Message, 0, 1+len(agentTurn.GetSteps()))
if len(agentTurn.GetUserMessage()) > 0 {
userMessage, err := decodeImportedUserMessage(agentTurn.GetUserMessage(), blobs)
if err != nil {
return nil, err
}
if replay, ok := promptengine.BuildUserMessageReplayMessage(userMessage); ok {
messages = append(messages, toModelMessage(replay))
}
}
for _, rawStep := range agentTurn.GetSteps() {
if len(rawStep) == 0 {
continue
}
step, err := decodeImportedStep(rawStep, blobs)
if err != nil {
return nil, err
}
for _, replay := range promptengine.BuildLegacyMessagesFromConversationStep(step) {
messages = append(messages, toModelMessage(replay))
}
}
return messages, nil
}
func importedTurnIDs(turns [][]byte, blobs importedBlobStore) ([][]byte, error) {
ids := make([][]byte, 0, len(turns))
for _, raw := range turns {
if len(raw) == 0 {
continue
}
_, id, err := decodeImportedTurn(raw, blobs)
if err != nil {
return nil, err
}
if len(id) > 0 {
ids = append(ids, id)
}
}
return ids, nil
}
func hasKnownUserMessageContent(message *agentv1.UserMessage) bool {
if message == nil {
return false
}
return message.GetText() != "" ||
message.GetMessageId() != "" ||
message.GetSelectedContext() != nil ||
message.GetRichText() != "" ||
len(message.GetConversationStateBlobId()) > 0 ||
len(message.GetTextBlobId()) > 0 ||
len(message.GetRichTextBlobId()) > 0
}
func firstNonNilError(err error, fallback error) error {
if err != nil {
return err
}
return fallback
}
@@ -1,109 +0,0 @@
package forwarder
import (
"crypto/sha256"
"testing"
"google.golang.org/protobuf/proto"
"cursor/gen/agentv1"
)
func TestImportedConversationStateRestoresBlobOnlyForkAndCheckpointPrefix(t *testing.T) {
parent := compactionAppendOnlyConversation(t)
parent.Entries = parent.Entries[:2]
parent.NextEntrySeq = 3
parent.NextTurnSeq = 2
projection, err := NewHistoryProjector().ProjectCheckpointProjection(parent)
if err != nil {
t.Fatalf("ProjectCheckpointProjection() error = %v", err)
}
prefetched := make([]*agentv1.PreFetchedBlob, 0, len(projection.Blobs))
for _, blob := range projection.Blobs {
prefetched = append(prefetched, &agentv1.PreFetchedBlob{Id: blob.ID, Value: blob.Data})
}
state := proto.Clone(projection.State).(*agentv1.ConversationStateStructure)
state.RootPromptMessagesJson = nil
conversation, err := newRuntimeConversation("fork-conversation", agentv1.AgentMode_AGENT_MODE_AGENT)
if err != nil {
t.Fatalf("newRuntimeConversation() error = %v", err)
}
entries, err := (&Service{}).importConversationState(conversation, state, prefetched)
if err != nil {
t.Fatalf("importConversationState() error = %v", err)
}
if len(conversation.ImportedTurnIDs) != 1 || conversation.NextTurnSeq != 2 {
t.Fatalf("imported prefix turns=%d next_turn_seq=%d, want 1 and 2", len(conversation.ImportedTurnIDs), conversation.NextTurnSeq)
}
if len(entries) != 2 {
t.Fatalf("imported model entries = %d, want parent user and assistant", len(entries))
}
appendEntriesInPlace(conversation, append(entries,
compactionTestUserEntry(t, 2, "request-2", "fork question", "message-2"),
))
forkProjection, err := NewHistoryProjector().ProjectCheckpointProjection(conversation)
if err != nil {
t.Fatalf("fork ProjectCheckpointProjection() error = %v", err)
}
if len(forkProjection.State.GetTurns()) != 2 {
t.Fatalf("fork checkpoint turns = %d, want imported parent plus local fork turn", len(forkProjection.State.GetTurns()))
}
if string(forkProjection.State.GetTurns()[0]) != string(projection.State.GetTurns()[0]) {
t.Fatal("fork checkpoint did not preserve the imported parent turn ID as its prefix")
}
}
func TestImportedConversationStateRejectsUnresolvedBlobTurn(t *testing.T) {
turnID := sha256.Sum256([]byte("missing imported turn"))
conversation, err := newRuntimeConversation("fork-conversation", agentv1.AgentMode_AGENT_MODE_AGENT)
if err != nil {
t.Fatalf("newRuntimeConversation() error = %v", err)
}
if _, err := (&Service{}).importConversationState(conversation, &agentv1.ConversationStateStructure{
Turns: [][]byte{turnID[:]},
}, nil); err == nil {
t.Fatal("importConversationState() accepted an unresolved Blob turn")
}
}
func TestImportedTurnIDsPersistThroughConversationStore(t *testing.T) {
store := NewConversationFileStore(t.TempDir())
turnID := sha256.Sum256([]byte("parent turn"))
conversation, err := newRuntimeConversation("fork-conversation", agentv1.AgentMode_AGENT_MODE_AGENT)
if err != nil {
t.Fatalf("newRuntimeConversation() error = %v", err)
}
conversation.ImportedTurnIDs = [][]byte{turnID[:]}
persisted, err := store.SaveConversationWithEntries(conversation.ConversationID, conversation, []HistoryEntry{
compactionTestUserEntry(t, 2, "request-2", "fork question", "message-2"),
})
if err != nil {
t.Fatalf("SaveConversationWithEntries() error = %v", err)
}
if len(persisted.ImportedTurnIDs) != 1 || string(persisted.ImportedTurnIDs[0]) != string(turnID[:]) {
t.Fatalf("persisted ImportedTurnIDs = %x, want %x", persisted.ImportedTurnIDs, turnID)
}
loaded, err := store.LoadConversation(conversation.ConversationID)
if err != nil {
t.Fatalf("LoadConversation() error = %v", err)
}
if len(loaded.ImportedTurnIDs) != 1 || string(loaded.ImportedTurnIDs[0]) != string(turnID[:]) {
t.Fatalf("loaded ImportedTurnIDs = %x, want %x", loaded.ImportedTurnIDs, turnID)
}
}
func TestRewindImportedTurnPrefixUsesClientForkPoint(t *testing.T) {
ids := make([][]byte, 3)
for index := range ids {
digest := sha256.Sum256([]byte{byte(index + 1)})
ids[index] = digest[:]
}
trimmed := rewindImportedTurnPrefix(ids, runRewindDecision{
TargetTurnSeq: 4,
HasClientTurnCount: true,
ClientTurnCount: 1,
})
if len(trimmed) != 1 || string(trimmed[0]) != string(ids[0]) {
t.Fatalf("rewindImportedTurnPrefix() = %x, want first imported turn only", trimmed)
}
}
@@ -2,7 +2,6 @@ package forwarder
import (
"encoding/json"
"errors"
"strings"
"testing"
)
@@ -126,53 +125,6 @@ func TestCancelPersistsInterruptedProviderOutputIdempotently(t *testing.T) {
}
}
func TestGenericProviderFailurePersistsAccumulatedOutput(t *testing.T) {
service, stream, _ := testCheckpointBlobProjection(t)
conversation, _, _, err := service.snapshotCheckpointConversation(stream)
if err != nil {
t.Fatalf("snapshotCheckpointConversation() error = %v", err)
}
if _, err := service.store.SaveConversationWithEntries(stream.ConversationID, conversation, conversation.Entries); err != nil {
t.Fatalf("SaveConversationWithEntries() error = %v", err)
}
stream.mu.Lock()
stream.CurrentModelCallID = "model-call-1"
stream.ProviderActive = true
stream.ProviderAccumulatedText = "partial answer before transport failure"
stream.Status = StreamStatusStreaming
stream.Phase = TurnPhaseProviderRunning
stream.mu.Unlock()
if err := service.handleProviderDoneEvent(stream, &streamProviderEvent{
Done: true,
Err: errors.New("transport failed"),
}); err != nil {
t.Fatalf("handleProviderDoneEvent() error = %v", err)
}
persisted, err := service.store.LoadConversation(stream.ConversationID)
if err != nil {
t.Fatalf("LoadConversation() error = %v", err)
}
foundPartialOutput := false
for _, entry := range persisted.Entries {
if entry.Kind != "assistant_text" {
continue
}
var payload assistantTextPayload
if err := json.Unmarshal(entry.Payload, &payload); err != nil {
t.Fatalf("decode assistant entry: %v", err)
}
if payload.Text == "partial answer before transport failure" {
foundPartialOutput = true
break
}
}
if !foundPartialOutput {
t.Fatal("generic provider failure discarded accumulated assistant output")
}
}
func TestCancelPreservesPersistedTurnActivityWithoutLiveAccumulator(t *testing.T) {
service, stream, _ := testCheckpointBlobProjection(t)
conversation, _, _, err := service.snapshotCheckpointConversation(stream)
+8
View File
@@ -32,3 +32,11 @@ func NewModule(historyRoot string, channelService modeladapter.ChannelResolver)
UploadServiceHandler: newUploadServiceHandler(service),
}
}
func (module *Module) HandlesAIPath(path string) bool {
if module == nil || module.AiHandler == nil {
return false
}
handler, ok := module.AiHandler.(interface{ HandlesPath(string) bool })
return ok && handler.HandlesPath(path)
}
+56 -85
View File
@@ -326,24 +326,20 @@ func compactedPromptProjectionEntries(entries []HistoryEntry) []HistoryEntry {
latestToolCallID := latestCompletedToolCallIDForTurn(entries, compactionPayload.CurrentTurnSeq, compactionPayload.CurrentRequestID)
preservedIndexes = autoCompactionPreservedEntryIndexes(entries, compactionPayload.CurrentTurnSeq, compactionPayload.CurrentRequestID, latestToolCallID)
}
filtered := make([]HistoryEntry, 0, len(entries)-compactionIndex+len(preservedIndexes))
for index := 0; index < compactionIndex; index++ {
if !isPromptReplayEntryKind(entries[index].Kind) {
filtered = append(filtered, entries[index])
filtered := make([]HistoryEntry, 0, len(entries)-compactionIndex)
for index, entry := range entries {
if index < compactionIndex && isPromptReplayEntryKind(entry.Kind) {
if _, ok := preservedIndexes[index]; !ok {
continue
}
}
}
filtered = append(filtered, entries[compactionIndex])
for index := 0; index < compactionIndex; index++ {
if _, ok := preservedIndexes[index]; !ok || isCompactionSummaryKind(entries[index].Kind) {
continue
}
entry := entries[index]
if rewritten, ok := compactedProjectionPreservedEntry(entry); ok {
entry = rewritten
if index < compactionIndex {
if rewritten, ok := compactedProjectionPreservedEntry(entry); ok {
entry = rewritten
}
}
filtered = append(filtered, entry)
}
filtered = append(filtered, entries[compactionIndex+1:]...)
return filtered
}
@@ -579,7 +575,7 @@ func (projector *HistoryProjector) ProjectCheckpointProjection(conversation *Con
if err != nil {
return nil, err
}
state.Turns = append(cloneByteSlices(conversation.ImportedTurnIDs), turnIDs...)
state.Turns = turnIDs
replayMessages, err := projector.ProjectPromptReplay(conversation)
if err != nil {
return nil, err
@@ -1196,82 +1192,57 @@ func mergeReplayReasoningMetadata(last *modeladapter.Message, current modeladapt
}
}
// replayToolResponseWindowEnd 返回 assistant tool-call 消息的响应收集窗口右边界(不含)。
// 窗口内除了 tool 结果消息,还允许出现同轮穿插的纯文本 assistant 消息:
// 部分模型(如 gpt-5.3-codex-spark)会在同一条响应里先输出 function_call 再输出说明文本,
// 落盘顺序为 tool_call → assistant_text → tool_result,若只收集紧邻的 tool 消息,
// 会把有结果回放的调用误判为悬空。
func replayToolResponseWindowEnd(messages []modeladapter.Message, index int) int {
end := index + 1
for end < len(messages) {
candidate := messages[end]
switch {
case strings.TrimSpace(candidate.Role) == "tool":
end++
case strings.TrimSpace(candidate.Role) == "assistant" && len(candidate.ToolCalls) == 0:
end++
default:
return end
}
}
return end
}
func trimReplayDanglingAssistantToolCalls(messages []modeladapter.Message) []modeladapter.Message {
if len(messages) == 0 {
return nil
}
survivingToolCallIDs := make(map[string]struct{})
for index, message := range messages {
if strings.TrimSpace(message.Role) != "assistant" || len(message.ToolCalls) == 0 {
continue
}
responded := make(map[string]struct{}, len(message.ToolCalls))
for scan := index + 1; scan < replayToolResponseWindowEnd(messages, index); scan++ {
if strings.TrimSpace(messages[scan].Role) != "tool" {
continue
}
if toolCallID := strings.TrimSpace(messages[scan].ToolCallID); toolCallID != "" {
responded[toolCallID] = struct{}{}
}
}
for _, toolCall := range message.ToolCalls {
if toolCallID := strings.TrimSpace(toolCall.ID); toolCallID != "" {
if _, ok := responded[toolCallID]; ok {
survivingToolCallIDs[toolCallID] = struct{}{}
}
}
}
}
trimmed := make([]modeladapter.Message, 0, len(messages))
for _, item := range messages {
message := cloneReplayModelMessage(item)
if strings.TrimSpace(message.Role) == "assistant" && len(message.ToolCalls) > 0 {
nextToolCalls := make([]modeladapter.ToolCallDescriptor, 0, len(message.ToolCalls))
for _, toolCall := range message.ToolCalls {
if _, ok := survivingToolCallIDs[strings.TrimSpace(toolCall.ID)]; !ok {
continue
}
toolCall.Index = len(nextToolCalls)
nextToolCalls = append(nextToolCalls, toolCall)
}
if len(nextToolCalls) == 0 {
if strings.TrimSpace(message.Content) == "" && len(message.ContentParts) == 0 && !hasReplayableReasoningPayload(message.ReasoningContent, message.ReasoningSignature, message.ReasoningSignatureSource) {
continue
}
message.ToolCalls = nil
} else {
message.ToolCalls = nextToolCalls
}
for index := 0; index < len(messages); index++ {
message := cloneReplayModelMessage(messages[index])
if strings.TrimSpace(message.Role) != "assistant" || len(message.ToolCalls) == 0 {
trimmed = append(trimmed, message)
continue
}
if strings.TrimSpace(message.Role) == "tool" && strings.TrimSpace(message.ToolCallID) != "" {
if _, ok := survivingToolCallIDs[strings.TrimSpace(message.ToolCallID)]; !ok {
end := index + 1
responded := make(map[string]struct{}, len(message.ToolCalls))
for end < len(messages) && strings.TrimSpace(messages[end].Role) == "tool" {
toolCallID := strings.TrimSpace(messages[end].ToolCallID)
if toolCallID != "" {
responded[toolCallID] = struct{}{}
}
end++
}
nextToolCalls := make([]modeladapter.ToolCallDescriptor, 0, len(message.ToolCalls))
allowedToolCallIDs := make(map[string]struct{}, len(message.ToolCalls))
for _, toolCall := range message.ToolCalls {
toolCallID := strings.TrimSpace(toolCall.ID)
if _, ok := responded[toolCallID]; !ok {
continue
}
item := toolCall
item.Index = len(nextToolCalls)
nextToolCalls = append(nextToolCalls, item)
allowedToolCallIDs[toolCallID] = struct{}{}
}
trimmed = append(trimmed, message)
if len(nextToolCalls) > 0 {
message.ToolCalls = nextToolCalls
trimmed = append(trimmed, message)
for toolIndex := index + 1; toolIndex < end; toolIndex++ {
toolMessage := cloneReplayModelMessage(messages[toolIndex])
if _, ok := allowedToolCallIDs[strings.TrimSpace(toolMessage.ToolCallID)]; !ok {
continue
}
trimmed = append(trimmed, toolMessage)
}
} else if strings.TrimSpace(message.Content) != "" || len(message.ContentParts) > 0 || hasReplayableReasoningPayload(message.ReasoningContent, message.ReasoningSignature, message.ReasoningSignatureSource) {
message.ToolCalls = nil
trimmed = append(trimmed, message)
}
index = end - 1
}
return trimmed
}
@@ -1321,7 +1292,7 @@ func filterCheckpointPersistentToolReplay(messages []promptengine.Message) []pro
return filtered
}
func restoreImportedReplayUserMessages(messages []promptengine.Message, importedTurns [][]byte, blobs importedBlobStore) []promptengine.Message {
func restoreImportedReplayUserMessages(messages []promptengine.Message, importedTurns [][]byte) []promptengine.Message {
if len(messages) == 0 || len(importedTurns) == 0 {
return messages
}
@@ -1330,16 +1301,16 @@ func restoreImportedReplayUserMessages(messages []promptengine.Message, imported
if len(rawTurn) == 0 {
continue
}
turn, _, err := decodeImportedTurn(rawTurn, blobs)
if err != nil || turn == nil {
turn := &agentv1.ConversationTurnStructure{}
if err := proto.Unmarshal(rawTurn, turn); err != nil {
continue
}
agentTurn := turn.GetAgentConversationTurn()
if agentTurn == nil || len(agentTurn.GetUserMessage()) == 0 {
continue
}
userMessage, err := decodeImportedUserMessage(agentTurn.GetUserMessage(), blobs)
if err != nil {
userMessage := &agentv1.UserMessage{}
if err := proto.Unmarshal(agentTurn.GetUserMessage(), userMessage); err != nil {
continue
}
replay, ok := promptengine.BuildUserMessageReplayMessage(userMessage)
@@ -1,180 +0,0 @@
package forwarder
import (
"bytes"
"crypto/sha256"
"encoding/base64"
"encoding/json"
"image"
"image/color"
"image/png"
"reflect"
"strings"
"testing"
"google.golang.org/protobuf/encoding/protojson"
"cursor/gen/agentv1"
modeladapter "cursor/internal/backend/agent/model"
)
func TestReadImageProjectionIsProviderOnlyAndIdempotent(t *testing.T) {
imageData := validForwarderTestPNG(t)
blobID := sha256.Sum256(imageData)
store := NewContentBlobStore(t.TempDir())
if err := store.Put(blobID[:], imageData); err != nil {
t.Fatalf("Put() error = %v", err)
}
conversation := readImageConversation(t, blobID[:], len(imageData))
projector := NewHistoryProjector()
canonical, err := projector.ProjectPromptReplay(conversation)
if err != nil {
t.Fatalf("ProjectPromptReplay() error = %v", err)
}
if len(canonical) != 2 {
t.Fatalf("canonical message count = %d, want 2", len(canonical))
}
if len(canonical[1].ContentParts) != 0 {
t.Fatalf("canonical replay contains image parts: %#v", canonical[1].ContentParts)
}
first, err := enrichProviderReadImages(canonical, conversation, store)
if err != nil {
t.Fatalf("first enrichment error = %v", err)
}
second, err := enrichProviderReadImages(canonical, conversation, store)
if err != nil {
t.Fatalf("second enrichment error = %v", err)
}
if !reflect.DeepEqual(first, second) {
t.Fatalf("provider enrichment is not idempotent\nfirst=%#v\nsecond=%#v", first, second)
}
reenriched, err := enrichProviderReadImages(first, conversation, store)
if err != nil {
t.Fatalf("re-enrichment error = %v", err)
}
if !reflect.DeepEqual(first, reenriched) {
t.Fatalf("provider enrichment changed an already enriched projection\nfirst=%#v\nreenriched=%#v", first, reenriched)
}
assertProviderReadImageMessage(t, first[1], imageData)
first[1].ContentParts[1].Image.Data[0] ^= 0xff
if bytes.Equal(first[1].ContentParts[1].Image.Data, second[1].ContentParts[1].Image.Data) {
t.Fatal("separate enrichments share mutable image bytes")
}
contextJSON, err := json.Marshal(conversation)
if err != nil {
t.Fatalf("marshal conversation: %v", err)
}
if bytes.Contains(contextJSON, imageData) || strings.Contains(string(contextJSON), base64.StdEncoding.EncodeToString(imageData)) {
t.Fatal("canonical conversation contains raw image bytes")
}
checkpoint, err := projector.ProjectCheckpointProjection(conversation)
if err != nil {
t.Fatalf("ProjectCheckpointProjection() error = %v", err)
}
checkpointJSON, err := json.Marshal(checkpoint)
if err != nil {
t.Fatalf("marshal checkpoint: %v", err)
}
if bytes.Contains(checkpointJSON, imageData) || strings.Contains(string(checkpointJSON), base64.StdEncoding.EncodeToString(imageData)) {
t.Fatal("checkpoint contains raw image bytes")
}
}
func TestProviderReadImageEnrichmentLeavesTextReadUnchanged(t *testing.T) {
toolCall := &agentv1.ToolCall{
Tool: &agentv1.ToolCall_ReadToolCall{
ReadToolCall: &agentv1.ReadToolCall{
Args: &agentv1.ReadToolArgs{Path: "notes.txt"},
Result: &agentv1.ReadToolResult{
Result: &agentv1.ReadToolResult_Success{
Success: &agentv1.ReadToolSuccess{
Path: "notes.txt",
Output: &agentv1.ReadToolSuccess_Content{Content: "hello"},
},
},
},
},
},
}
encoded, err := protojson.Marshal(toolCall)
if err != nil {
t.Fatalf("marshal tool call: %v", err)
}
conversation := &ConversationFile{Entries: []HistoryEntry{
newToolResultEntry(1, "request-1", "call-1", "Read", `{"path":"notes.txt"}`, "hello", "", encoded),
}}
messages := []modeladapter.Message{{Role: "tool", ToolCallID: "call-1", Name: "Read", Content: "hello"}}
got, err := enrichProviderReadImages(messages, conversation, NewContentBlobStore(t.TempDir()))
if err != nil {
t.Fatalf("enrichProviderReadImages() error = %v", err)
}
if !reflect.DeepEqual(got, messages) {
t.Fatalf("text read changed: got=%#v want=%#v", got, messages)
}
}
func readImageConversation(t *testing.T, blobID []byte, fileSize int) *ConversationFile {
t.Helper()
toolCall := &agentv1.ToolCall{
Tool: &agentv1.ToolCall_ReadToolCall{
ReadToolCall: &agentv1.ReadToolCall{
Args: &agentv1.ReadToolArgs{Path: "diagram.png"},
Result: &agentv1.ReadToolResult{
Result: &agentv1.ReadToolResult_Success{
Success: &agentv1.ReadToolSuccess{
FileSize: uint32(fileSize),
Path: "diagram.png",
Output: &agentv1.ReadToolSuccess_DataBlobId{DataBlobId: append([]byte(nil), blobID...)},
},
},
},
},
},
}
encoded, err := protojson.Marshal(toolCall)
if err != nil {
t.Fatalf("marshal tool call: %v", err)
}
return &ConversationFile{
ConversationID: "conversation-1",
Mode: "agent",
NextTurnSeq: 2,
Entries: []HistoryEntry{
newToolResultEntry(1, "request-1", "call-1", "Read", `{"path":"diagram.png"}`, "read binary bytes", "", encoded),
},
}
}
func assertProviderReadImageMessage(t *testing.T, message modeladapter.Message, imageData []byte) {
t.Helper()
if message.Role != "tool" || message.ToolCallID != "call-1" || message.Name != "Read" {
t.Fatalf("tool message metadata = %#v", message)
}
if len(message.ContentParts) != 2 {
t.Fatalf("content part count = %d, want text and image", len(message.ContentParts))
}
if message.ContentParts[0].Type != "text" || message.ContentParts[0].Text != message.Content {
t.Fatalf("text content part = %#v", message.ContentParts[0])
}
imagePart := message.ContentParts[1]
if imagePart.Type != "image" || imagePart.Image == nil {
t.Fatalf("image content part = %#v", imagePart)
}
if imagePart.Image.MIMEType != "image/png" || imagePart.Image.Path != "diagram.png" || !bytes.Equal(imagePart.Image.Data, imageData) {
t.Fatalf("image content = %#v", imagePart.Image)
}
}
func validForwarderTestPNG(t *testing.T) []byte {
t.Helper()
value := image.NewRGBA(image.Rect(0, 0, 2, 2))
value.Set(0, 0, color.RGBA{R: 0x44, G: 0x88, B: 0xcc, A: 0xff})
var encoded bytes.Buffer
if err := png.Encode(&encoded, value); err != nil {
t.Fatalf("encode test png: %v", err)
}
return encoded.Bytes()
}
@@ -1,89 +0,0 @@
package forwarder
import (
"testing"
modeladapter "cursor/internal/backend/agent/model"
)
func replayToolCall(id string, name string) modeladapter.ToolCallDescriptor {
return modeladapter.ToolCallDescriptor{
ID: id,
Type: "function",
Function: modeladapter.ToolCallFunctionShape{
Name: name,
Arguments: "{}",
},
}
}
// 模型在同一条响应里先输出 function_call 再输出说明文本时,
// 历史回放顺序为 assistant[tool_call] -> assistant[text] -> tool[result]。
// trimReplayDanglingAssistantToolCalls 需要越过中间的文本消息收集工具结果。
func TestTrimReplayDanglingAssistantToolCallsKeepsInterleavedTextResponses(t *testing.T) {
messages := []modeladapter.Message{
{Role: "user", Content: "query"},
{Role: "assistant", ToolCalls: []modeladapter.ToolCallDescriptor{replayToolCall("call_1", "Grep")}},
{Role: "assistant", Content: "我先快速定位上下文"},
{Role: "tool", Name: "Grep", ToolCallID: "call_1", Content: "grep result"},
{Role: "user", Content: "next"},
}
trimmed := trimReplayDanglingAssistantToolCalls(messages)
if len(trimmed) != 5 {
t.Fatalf("expected 5 messages, got %d: %+v", len(trimmed), trimmed)
}
if len(trimmed[1].ToolCalls) != 1 || trimmed[1].ToolCalls[0].ID != "call_1" {
t.Fatalf("expected assistant tool call call_1 to survive, got %+v", trimmed[1].ToolCalls)
}
if trimmed[3].Role != "tool" || trimmed[3].ToolCallID != "call_1" {
t.Fatalf("expected tool result call_1 to survive, got %+v", trimmed[3])
}
}
// 没有任何结果回放的调用仍应被剥离;被剥离调用对应的 tool 结果(若存在)也不得保留。
func TestTrimReplayDanglingAssistantToolCallsDropsUnrespondedCallsAndOrphanResults(t *testing.T) {
messages := []modeladapter.Message{
{Role: "user", Content: "query"},
{Role: "assistant", ToolCalls: []modeladapter.ToolCallDescriptor{
replayToolCall("call_1", "Grep"),
replayToolCall("call_2", "Read"),
}},
{Role: "tool", Name: "Grep", ToolCallID: "call_1", Content: "grep result"},
{Role: "tool", Name: "Read", ToolCallID: "call_3", Content: "orphan result"},
{Role: "user", Content: "next"},
}
trimmed := trimReplayDanglingAssistantToolCalls(messages)
if len(trimmed) != 4 {
t.Fatalf("expected 4 messages, got %d: %+v", len(trimmed), trimmed)
}
if len(trimmed[1].ToolCalls) != 1 || trimmed[1].ToolCalls[0].ID != "call_1" {
t.Fatalf("expected only call_1 to survive, got %+v", trimmed[1].ToolCalls)
}
if trimmed[2].ToolCallID != "call_1" {
t.Fatalf("expected only call_1 result to survive, got %+v", trimmed[2])
}
}
// 调用全部悬空但消息携带可回放 reasoning 时,保留为无调用的 assistant 消息。
func TestTrimReplayDanglingAssistantToolCallsKeepsReasoningOnlyShell(t *testing.T) {
messages := []modeladapter.Message{
{Role: "user", Content: "query"},
{
Role: "assistant",
ToolCalls: []modeladapter.ToolCallDescriptor{replayToolCall("call_1", "Grep")},
ReasoningSignature: "sig",
ReasoningSignatureSource: modeladapter.ReasoningSignatureSourceOpenAIResponses,
},
{Role: "user", Content: "next"},
}
trimmed := trimReplayDanglingAssistantToolCalls(messages)
if len(trimmed) != 3 {
t.Fatalf("expected 3 messages, got %d: %+v", len(trimmed), trimmed)
}
if len(trimmed[1].ToolCalls) != 0 {
t.Fatalf("expected tool calls to be trimmed, got %+v", trimmed[1].ToolCalls)
}
}
@@ -1,174 +0,0 @@
// provider_read_images.go 负责在 provider 请求边界按 blob 引用补全 Read 图片。
package forwarder
import (
"bytes"
"encoding/json"
"fmt"
"image"
_ "image/gif"
_ "image/jpeg"
_ "image/png"
"net/http"
"strings"
"google.golang.org/protobuf/encoding/protojson"
"cursor/gen/agentv1"
modeladapter "cursor/internal/backend/agent/model"
)
type contentBlobReader interface {
Get(id []byte) ([]byte, error)
}
type providerReadImageReference struct {
blobID []byte
path string
fileSize uint32
}
// enrichProviderReadImages 只为本次 provider 请求加载图片,不修改 canonical history 投影。
func enrichProviderReadImages(messages []modeladapter.Message, conversation *ConversationFile, blobs contentBlobReader) ([]modeladapter.Message, error) {
cloned := cloneProviderEnrichmentMessages(messages)
references, err := collectProviderReadImageReferences(conversation)
if err != nil {
return nil, err
}
if len(references) == 0 {
return cloned, nil
}
for index := range cloned {
message := &cloned[index]
if strings.TrimSpace(message.Role) != "tool" {
continue
}
reference, ok := references[strings.TrimSpace(message.ToolCallID)]
if !ok {
continue
}
if blobs == nil {
return nil, fmt.Errorf("provider read image blob store is not initialized")
}
data, err := blobs.Get(reference.blobID)
if err != nil {
return nil, fmt.Errorf("load read image blob for tool call %s: %w", message.ToolCallID, err)
}
mimeType := validatedProviderReadImageMIMEType(data)
if mimeType == "" {
return nil, fmt.Errorf("read image blob for tool call %s is not a supported image", message.ToolCallID)
}
summary := "Read image file: " + reference.path
message.Content = summary
message.ContentParts = []modeladapter.ContentPart{
{Type: "text", Text: summary},
{
Type: "image",
Image: &modeladapter.ImageContent{
MIMEType: mimeType,
Path: reference.path,
Data: append([]byte(nil), data...),
},
},
}
}
return cloned, nil
}
func collectProviderReadImageReferences(conversation *ConversationFile) (map[string]providerReadImageReference, error) {
references := make(map[string]providerReadImageReference)
if conversation == nil {
return references, nil
}
for _, entry := range conversation.Entries {
if strings.TrimSpace(entry.Kind) != "tool_result" {
continue
}
var payload toolResultEntryPayload
if err := json.Unmarshal(entry.Payload, &payload); err != nil {
return nil, fmt.Errorf("decode read image tool result entry: %w", err)
}
if len(payload.ToolCall) == 0 {
continue
}
toolCall := &agentv1.ToolCall{}
if err := protojson.Unmarshal(payload.ToolCall, toolCall); err != nil {
return nil, fmt.Errorf("decode read image tool call: %w", err)
}
readToolCall := toolCall.GetReadToolCall()
if readToolCall == nil || readToolCall.GetResult().GetSuccess() == nil {
continue
}
success := readToolCall.GetResult().GetSuccess()
blobID := success.GetDataBlobId()
if len(blobID) == 0 {
continue
}
toolCallID := strings.TrimSpace(firstNonEmpty(payload.ToolCallID, entry.ToolCallID))
if toolCallID == "" {
continue
}
reference := providerReadImageReference{
blobID: append([]byte(nil), blobID...),
path: firstNonEmpty(strings.TrimSpace(success.GetPath()), strings.TrimSpace(readToolCall.GetArgs().GetPath())),
fileSize: success.GetFileSize(),
}
if existing, ok := references[toolCallID]; ok {
if !bytes.Equal(existing.blobID, reference.blobID) || existing.path != reference.path || existing.fileSize != reference.fileSize {
return nil, fmt.Errorf("conflicting read image references for tool call %s", toolCallID)
}
continue
}
references[toolCallID] = reference
}
return references, nil
}
func cloneProviderEnrichmentMessages(messages []modeladapter.Message) []modeladapter.Message {
if len(messages) == 0 {
return nil
}
cloned := make([]modeladapter.Message, 0, len(messages))
for _, message := range messages {
item := cloneReplayModelMessage(message)
if len(message.ContentParts) > 0 {
item.ContentParts = make([]modeladapter.ContentPart, len(message.ContentParts))
for index, part := range message.ContentParts {
item.ContentParts[index] = part
if part.Image != nil {
imageCopy := *part.Image
imageCopy.Data = append([]byte(nil), part.Image.Data...)
item.ContentParts[index].Image = &imageCopy
}
}
}
cloned = append(cloned, item)
}
return cloned
}
func validatedProviderReadImageMIMEType(data []byte) string {
if len(data) == 0 {
return ""
}
detected := strings.ToLower(strings.TrimSpace(http.DetectContentType(data)))
configuration, format, err := image.DecodeConfig(bytes.NewReader(data))
if err != nil || configuration.Width <= 0 || configuration.Height <= 0 {
return ""
}
switch strings.ToLower(strings.TrimSpace(format)) {
case "png":
if detected == "image/png" {
return detected
}
case "jpeg":
if detected == "image/jpeg" {
return detected
}
case "gif":
if detected == "image/gif" {
return detected
}
}
return ""
}
-21
View File
@@ -224,7 +224,6 @@ func (service *Service) applyRunRewindToConversation(conversation *ConversationF
conversation.Entries = nil
conversation.NextEntrySeq = 1
conversation.NextTurnSeq = 1
conversation.ImportedTurnIDs = rewindImportedTurnPrefix(conversation.ImportedTurnIDs, decision)
appendEntriesInPlace(conversation, appendReplacementRunEntries(decision.PrefixEntries, entries))
applyRunRewindConversationState(conversation, intent, turnSeq)
deriveConversationLoopState(conversation)
@@ -270,30 +269,10 @@ func applyRunRewindMetadata(conversation *ConversationFile, source *Conversation
if source.TokenDetailsMaxTokens > 0 {
conversation.TokenDetailsMaxTokens = source.TokenDetailsMaxTokens
}
decision := runRewindDecision{TargetTurnSeq: turnSeq}
if intent.ConversationState != nil {
decision.HasClientTurnCount = true
decision.ClientTurnCount = len(intent.ConversationState.GetTurns())
}
conversation.ImportedTurnIDs = rewindImportedTurnPrefix(source.ImportedTurnIDs, decision)
}
applyRunRewindConversationState(conversation, intent, turnSeq)
}
func rewindImportedTurnPrefix(importedTurnIDs [][]byte, decision runRewindDecision) [][]byte {
keep := decision.TargetTurnSeq - 1
if decision.HasClientTurnCount {
keep = int64(decision.ClientTurnCount)
}
if keep <= 0 || len(importedTurnIDs) == 0 {
return nil
}
if keep > int64(len(importedTurnIDs)) {
keep = int64(len(importedTurnIDs))
}
return cloneByteSlices(importedTurnIDs[:keep])
}
func (service *Service) logRunRewindDecision(requestID string, conversationID string, eventName string, decision runRewindDecision) {
if service == nil || !decision.Evaluated {
return
@@ -50,7 +50,7 @@ func (service *Service) bootstrapRuntimeConversation(intent InboundIntent) (*Con
}
importedEntries := []HistoryEntry(nil)
if len(conversation.Entries) == 0 && intent.ConversationState != nil {
importedEntries, err = service.importConversationState(conversation, intent.ConversationState, intent.PreFetchedBlobs)
importedEntries, err = service.importConversationState(conversation, intent.ConversationState)
if err != nil {
return nil, agentv1.AgentMode_AGENT_MODE_AGENT, 0, nil, err
}
@@ -138,7 +138,6 @@ func (service *Service) syncConversationRecord(conversationID string, conversati
item.AutoCompactionReserveTokens = conversation.AutoCompactionReserveTokens
item.AutoCompactionTriggeredAt = conversation.AutoCompactionTriggeredAt
item.AutoCompactionSourceModelCallID = conversation.AutoCompactionSourceModelCallID
item.ImportedTurnIDs = cloneByteSlices(conversation.ImportedTurnIDs)
item.LatestRequestPrefix = cloneConversationRequestPrefix(conversation.LatestRequestPrefix)
item.LastProviderCall = cloneConversationProviderCall(conversation.LastProviderCall)
item.CreatedAt = conversation.CreatedAt
+12 -51
View File
@@ -248,7 +248,6 @@ func subagentModelOverrideSummaries(overrides map[string]runtimecore.SubagentMod
type Service struct {
store *ConversationFileStore
contentBlobs *ContentBlobStore
usageStore *UsageFileStore
codebaseIndexStore *CodebaseIndexStore
docsIndexStore *DocsIndexStore
@@ -275,7 +274,6 @@ type agentModelMemory interface {
func NewService(historyRoot string, resolver modeladapter.ChannelResolver) *Service {
projector := NewHistoryProjector()
store := NewConversationFileStore(historyRoot)
contentBlobs := NewContentBlobStore(historyRoot)
broker := NewStreamBroker()
rules := NewUserRuleStore(appdata.RulesRootPath())
var modelMemory agentModelMemory
@@ -289,13 +287,12 @@ func NewService(historyRoot string, resolver modeladapter.ChannelResolver) *Serv
debug := newDebugRecorder(historyRoot, broker, debugConfig)
service := &Service{
store: store,
contentBlobs: contentBlobs,
usageStore: NewUsageFileStore(historyRoot),
codebaseIndexStore: NewCodebaseIndexStore(appdata.CodebaseIndexRootPath()),
docsIndexStore: NewDocsIndexStore(appdata.DocsIndexRootPath()),
rules: rules,
projector: projector,
compiler: NewPromptCompiler(projector, NewToolCatalog(), NewReminderInjector(), rules, contentBlobs),
compiler: NewPromptCompiler(projector, NewToolCatalog(), NewReminderInjector(), rules),
provider: NewProviderGateway(resolver),
resolver: resolver,
modelMemory: modelMemory,
@@ -320,7 +317,6 @@ func newServiceWithDependencies(store *ConversationFileStore, projector *History
debug := newDebugRecorder(historyRoot, broker, nil)
return &Service{
store: store,
contentBlobs: NewContentBlobStore(historyRoot),
rules: NewUserRuleStore(appdata.RulesRootPath()),
projector: projector,
compiler: compiler,
@@ -563,7 +559,6 @@ func (service *Service) decodeInboundIntent(requestID string, message *agentv1.A
}
intent.ConversationID = conversationID
intent.ConversationState = runRequest.GetConversationState()
intent.PreFetchedBlobs = runRequest.GetPreFetchedBlobs()
intent.UserMessage = extractUserMessage(message)
intent.RequestContext = extractRequestContext(message)
if service.shouldIgnoreEmptyResumeRunRequest(requestID, runRequest, intent.UserMessage, intent.RequestContext) {
@@ -611,7 +606,6 @@ func (service *Service) decodeInboundIntent(requestID string, message *agentv1.A
intent.ConversationID = conversationID
intent.SubagentTypeName = strings.TrimSpace(prewarmRequest.GetSubagentTypeName())
intent.ConversationState = prewarmRequest.GetConversationState()
intent.PreFetchedBlobs = prewarmRequest.GetPreFetchedBlobs()
intent.Mode, intent.ModeSource, intent.HasExplicitMode, err = extractPrewarmMode(prewarmRequest)
if err != nil {
return InboundIntent{}, err
@@ -1018,9 +1012,6 @@ func (service *Service) handleExecResult(intent InboundIntent) error {
if !result.IsTerminal {
return nil
}
if err := service.persistExecContentBlobs(result.ContentBlobs); err != nil {
return err
}
markExecCompleted(stream, pending)
backgroundShellToolCallID := ""
if strings.TrimSpace(pending.ExecKind) == "shell" && shellToolCallIsBackgrounded(result.ToolCall) {
@@ -1059,21 +1050,6 @@ func (service *Service) handleExecResult(intent InboundIntent) error {
return service.reconcileStream(stream)
}
func (service *Service) persistExecContentBlobs(blobs []execbridge.ContentBlob) error {
if len(blobs) == 0 {
return nil
}
if service == nil || service.contentBlobs == nil {
return fmt.Errorf("content blob store is not initialized")
}
for _, blob := range blobs {
if err := service.contentBlobs.Put(blob.ID, blob.Data); err != nil {
return fmt.Errorf("persist exec content blob: %w", err)
}
}
return nil
}
// handleExecControl 处理执行桥控制面结果,例如 stream_close 或 throw。
func (service *Service) handleExecControl(intent InboundIntent) error {
stream, ok := service.broker.Get(intent.RequestID)
@@ -2270,15 +2246,6 @@ func (service *Service) finishSuccessfulTurnAfterCheckpoint(stream *ActiveStream
return nil
}
func (service *Service) finishFailedTurnAfterCheckpoint(stream *ActiveStream, terminalCode string, terminalMessage string) error {
if stream == nil {
return nil
}
err := service.broker.Fail(stream.RequestID, terminalCode, terminalMessage)
service.setTurnPhase(stream, TurnPhaseFailed)
return err
}
func (service *Service) failStreamIfNonTerminal(stream *ActiveStream, terminalCode string, cause error) error {
if stream == nil || cause == nil {
return nil
@@ -2298,10 +2265,6 @@ func (service *Service) publishCheckpoint(requestID string, conversationID strin
}
func (service *Service) publishCheckpointWithCompletion(requestID string, _ string, completion *pendingTurnCompletion) error {
return service.publishCheckpointWithTerminalAction(requestID, successfulCheckpointTerminalAction(completion))
}
func (service *Service) publishCheckpointWithTerminalAction(requestID string, terminal checkpointTerminalAction) error {
stream, ok := service.broker.Get(requestID)
if !ok || stream == nil {
return fmt.Errorf("request is not active: %s", requestID)
@@ -2319,7 +2282,7 @@ func (service *Service) publishCheckpointWithTerminalAction(requestID string, te
}
projection.State.PendingToolCalls = buildPendingToolCalls(pendingExecs, pendingInteractions)
service.rewriteCheckpointTokenDetailsForClient(stream, conversation, projection.State)
return service.queueCheckpointProjectionWithTerminal(stream, projection, terminal)
return service.queueCheckpointProjection(stream, projection, completion)
}
func (service *Service) rewriteCheckpointTokenDetailsForClient(stream *ActiveStream, conversation *ConversationFile, state *agentv1.ConversationStateStructure) {
@@ -2459,20 +2422,18 @@ func (service *Service) failActiveStream(stream *ActiveStream, conversationID st
if cancel != nil {
cancel()
}
if err := service.syncSummaryCarryForward(conversationID, requestID, modelCallID); err != nil {
log.Printf(
"forwarder summary sync before failed terminal skipped request_id=%s model_call_id=%s err=%v",
strings.TrimSpace(requestID),
strings.TrimSpace(modelCallID),
err,
)
service.setTurnPhase(stream, TurnPhaseFailed)
var firstErr error
if err := service.syncSummaryCarryForward(conversationID, requestID, modelCallID); err != nil && firstErr == nil {
firstErr = err
}
terminal := failedCheckpointTerminalAction(terminalCode, terminalMessage)
if err := service.publishCheckpointWithTerminalAction(requestID, terminal); err != nil {
log.Printf("forwarder checkpoint queue before failed terminal skipped request_id=%s err=%v", strings.TrimSpace(requestID), err)
return service.finishFailedTurnAfterCheckpoint(stream, terminalCode, terminalMessage)
if err := service.publishCheckpoint(requestID, conversationID); err != nil && firstErr == nil {
firstErr = err
}
return nil
if err := service.broker.Fail(requestID, terminalCode, terminalMessage); err != nil && firstErr == nil {
firstErr = err
}
return firstErr
}
// buildRunEntries 构造一次 run intent 需要写入 history 的首批 entry。
+29 -28
View File
@@ -45,25 +45,13 @@ func (snapshot turnUsageSnapshot) requestTokensTotal() int64 {
return snapshot.promptTokensTotal() + nonNegativeInt64(snapshot.OutputTokens)
}
func (service *Service) importConversationState(item *ConversationFile, state *agentv1.ConversationStateStructure, prefetchedBlobs []*agentv1.PreFetchedBlob) ([]HistoryEntry, error) {
func (service *Service) importConversationState(item *ConversationFile, state *agentv1.ConversationStateStructure) ([]HistoryEntry, error) {
if item == nil || state == nil {
return nil, nil
}
blobs, err := newImportedBlobStore(prefetchedBlobs)
if err != nil {
return nil, err
}
importedIDs, err := importedTurnIDs(state.GetTurns(), blobs)
if err != nil {
return nil, err
}
item.TokenDetailsUsedTokens = state.GetTokenDetails().GetUsedTokens()
item.ImportedTurnIDs = importedIDs
if minimumNextTurnSeq := int64(len(item.ImportedTurnIDs)) + 1; item.NextTurnSeq < minimumNextTurnSeq {
item.NextTurnSeq = minimumNextTurnSeq
}
entries := make([]HistoryEntry, 0, 2)
if messages, err := importedConversationStateModelMessagesWithBlobs(state, blobs); err != nil {
if messages, err := importedConversationStateModelMessages(state); err != nil {
return nil, err
} else {
for _, message := range messages {
@@ -117,10 +105,6 @@ func (service *Service) importConversationState(item *ConversationFile, state *a
}
func importedConversationStateModelMessages(state *agentv1.ConversationStateStructure) ([]modeladapter.Message, error) {
return importedConversationStateModelMessagesWithBlobs(state, nil)
}
func importedConversationStateModelMessagesWithBlobs(state *agentv1.ConversationStateStructure, blobs importedBlobStore) ([]modeladapter.Message, error) {
if state == nil {
return nil, nil
}
@@ -129,7 +113,7 @@ func importedConversationStateModelMessagesWithBlobs(state *agentv1.Conversation
if err != nil {
return nil, fmt.Errorf("decode imported replay messages: %w", err)
}
decoded = restoreImportedReplayUserMessages(decoded, state.GetTurns(), blobs)
decoded = restoreImportedReplayUserMessages(decoded, state.GetTurns())
decoded = filterLegacyPlainWriteReplay(decoded)
decoded = filterInternalPromptContextReplay(decoded)
messages := make([]modeladapter.Message, 0, len(decoded))
@@ -149,18 +133,35 @@ func importedConversationStateModelMessagesWithBlobs(state *agentv1.Conversation
if len(rawTurn) == 0 {
continue
}
turn, turnID, err := decodeImportedTurn(rawTurn, blobs)
if err != nil {
return nil, err
turn := &agentv1.ConversationTurnStructure{}
if err := proto.Unmarshal(rawTurn, turn); err != nil {
return nil, fmt.Errorf("decode imported turn: %w", err)
}
if turn == nil && len(turnID) > 0 {
return nil, fmt.Errorf("missing prefetched turn blob %x", turnID)
agentTurn := turn.GetAgentConversationTurn()
if agentTurn == nil {
continue
}
turnMessages, err := importedBlobTurnMessages(turn, blobs)
if err != nil {
return nil, err
if rawUser := agentTurn.GetUserMessage(); len(rawUser) > 0 {
userMessage := &agentv1.UserMessage{}
if err := proto.Unmarshal(rawUser, userMessage); err != nil {
return nil, fmt.Errorf("decode imported turn user_message: %w", err)
}
if replay, ok := promptengine.BuildUserMessageReplayMessage(userMessage); ok {
messages = append(messages, toModelMessage(replay))
}
}
for _, rawStep := range agentTurn.GetSteps() {
if len(rawStep) == 0 {
continue
}
step := &agentv1.ConversationStep{}
if err := proto.Unmarshal(rawStep, step); err != nil {
return nil, fmt.Errorf("decode imported turn step: %w", err)
}
for _, replay := range promptengine.BuildLegacyMessagesFromConversationStep(step) {
messages = append(messages, toModelMessage(replay))
}
}
messages = append(messages, turnMessages...)
}
return normalizeReplayMessageSequence(messages), nil
}
+4 -20
View File
@@ -39,7 +39,6 @@ type ConversationFile struct {
CurrentPlanText string `json:"current_plan_text,omitempty"`
CurrentPlans map[string]*agentv1.PlanRegistryEntry `json:"current_plans,omitempty"`
CurrentTodos []*agentv1.TodoItem `json:"current_todos,omitempty"`
ImportedTurnIDs [][]byte `json:"imported_turn_ids,omitempty"`
LatestRequestPrefix *ConversationRequestPrefix `json:"latest_request_prefix,omitempty"`
LastProviderCall *ConversationProviderCall `json:"last_provider_call,omitempty"`
CreatedAt time.Time `json:"created_at"`
@@ -225,25 +224,11 @@ type pendingTurnCompletion struct {
Disposition pendingCompletionDisposition
}
type checkpointTerminalActionKind uint8
const (
checkpointTerminalActionNone checkpointTerminalActionKind = iota
checkpointTerminalActionComplete
checkpointTerminalActionFail
)
type checkpointTerminalAction struct {
Kind checkpointTerminalActionKind
Completion pendingTurnCompletion
ErrorCode string
ErrorMessage string
}
type pendingCheckpointPublish struct {
State *agentv1.ConversationStateStructure
Required map[string]struct{}
Terminal checkpointTerminalAction
State *agentv1.ConversationStateStructure
Required map[string]struct{}
Completion *pendingTurnCompletion
Published bool
}
type PendingCompaction struct {
@@ -445,7 +430,6 @@ type InboundIntent struct {
SubagentTypeName string
SubagentModelOverrides map[string]runtimecore.SubagentModelOverrideSelection
ConversationState *agentv1.ConversationStateStructure
PreFetchedBlobs []*agentv1.PreFetchedBlob
UserMessage *agentv1.UserMessage
RequestContext *agentv1.RequestContext
ClientMessage *agentv1.AgentClientMessage
+119 -278
View File
@@ -2,6 +2,8 @@ package backend
import (
"context"
"crypto/tls"
"crypto/x509"
"fmt"
"net"
"net/http"
@@ -27,11 +29,11 @@ const healthPath = "/healthz"
const tabServerBaseURL = "https://tab.leokun.cn"
type Host struct {
store *serverconfig.Store
listenAddr string
configs *serverconfig.Manager
healthHTTP *http.Client
controlPlaneAuth upstream.AuthorizationProvider
store *serverconfig.Store
listenAddr string
configs *serverconfig.Manager
healthHTTP *http.Client
tlsCertificate *tls.Certificate
runMu sync.RWMutex
httpServer *http.Server
@@ -41,7 +43,20 @@ type Host struct {
mux http.Handler
}
func NewHost(store *serverconfig.Store, controlPlaneAuth upstream.AuthorizationProvider) (*Host, error) {
type HostOption func(*Host) error
func WithTLSCertificate(certificate *tls.Certificate) HostOption {
return func(host *Host) error {
if certificate == nil || len(certificate.Certificate) == 0 || certificate.PrivateKey == nil {
return fmt.Errorf("backend TLS certificate is invalid")
}
copied := *certificate
host.tlsCertificate = &copied
return nil
}
}
func NewHost(store *serverconfig.Store, options ...HostOption) (*Host, error) {
if store == nil {
return nil, fmt.Errorf("backend config store is required")
}
@@ -51,12 +66,19 @@ func NewHost(store *serverconfig.Store, controlPlaneAuth upstream.AuthorizationP
}
cfg := configs.Current()
host := &Host{
store: store,
listenAddr: cfg.BackendListenAddr,
configs: configs,
healthHTTP: newLoopbackHTTPClient(),
controlPlaneAuth: controlPlaneAuth,
store: store,
listenAddr: cfg.BackendListenAddr,
configs: configs,
}
for _, option := range options {
if option == nil {
continue
}
if err := option(host); err != nil {
return nil, err
}
}
host.healthHTTP = newLoopbackHTTPClient(host.tlsCertificate)
if err := host.rebuild(cfg); err != nil {
return nil, err
}
@@ -107,7 +129,14 @@ func (host *Host) BaseURL() string {
if listenAddr == "" {
return ""
}
return "http://" + listenAddr
if host.tlsCertificate == nil {
return "http://" + listenAddr
}
serverName := "localhost"
if _, port, err := net.SplitHostPort(listenAddr); err == nil {
return "https://" + net.JoinHostPort(serverName, port)
}
return "https://" + listenAddr
}
func (host *Host) IsRunning() bool {
@@ -153,6 +182,12 @@ func (host *Host) Start() error {
host.lastRunErr = fmt.Errorf("监听内置后端 %s 失败: %w", host.listenAddr, err)
return host.lastRunErr
}
if host.tlsCertificate != nil {
listener = tls.NewListener(listener, &tls.Config{
Certificates: []tls.Certificate{*host.tlsCertificate},
MinVersion: tls.VersionTLS12,
})
}
host.listenAddr = listener.Addr().String()
host.httpServer = httpServer
host.lastRunErr = nil
@@ -202,7 +237,7 @@ func (host *Host) HealthCheck(ctx context.Context) error {
}
client := host.healthHTTP
if client == nil {
client = newLoopbackHTTPClient()
client = newLoopbackHTTPClient(host.tlsCertificate)
}
response, err := client.Do(request)
if err != nil {
@@ -245,19 +280,34 @@ func (host *Host) InProcessHealthCheck() error {
return nil
}
func newLoopbackHTTPClient() *http.Client {
func newLoopbackHTTPClient(certificate *tls.Certificate) *http.Client {
transport := &http.Transport{
Proxy: nil,
DialContext: (&net.Dialer{
Timeout: 1 * time.Second,
KeepAlive: 30 * time.Second,
}).DialContext,
ForceAttemptHTTP2: false,
MaxIdleConns: 1,
MaxIdleConnsPerHost: 1,
IdleConnTimeout: 30 * time.Second,
}
if certificate != nil {
roots := x509.NewCertPool()
for _, rawCertificate := range certificate.Certificate[1:] {
parsed, err := x509.ParseCertificate(rawCertificate)
if err == nil {
roots.AddCert(parsed)
}
}
transport.TLSClientConfig = &tls.Config{
MinVersion: tls.VersionTLS12,
RootCAs: roots,
ServerName: "localhost",
}
}
return &http.Client{
Transport: &http.Transport{
Proxy: nil,
DialContext: (&net.Dialer{
Timeout: 1 * time.Second,
KeepAlive: 30 * time.Second,
}).DialContext,
ForceAttemptHTTP2: false,
MaxIdleConns: 1,
MaxIdleConnsPerHost: 1,
IdleConnTimeout: 30 * time.Second,
},
Transport: transport,
}
}
@@ -276,8 +326,20 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
SystemSettingService: &serverSystemSettings{configs: host.configs},
HTTPClient: netproxy.NewHTTPClient(30000 * time.Second),
}
fallbackForward := upstream.FallbackForwardAction(
routeDeps,
upstream.CompatRouteConfig{Name: "upstream_fallback"},
upstream.DefaultCursorUpstreamBaseURL,
)
localAIAction := server.HTTPHandlerAction(agentModule.AiHandler)
aiServiceAction := func(ctx *server.Context) error {
if ctx != nil && ctx.Request != nil && ctx.Request.URL != nil && agentModule.HandlesAIPath(ctx.Request.URL.Path) {
return localAIAction(ctx)
}
return fallbackForward(ctx)
}
host.mux = server.New(
host.mux = withLocalBackendCORS(server.New(
server.Use(
server.Recover(),
server.ServerContext(),
@@ -427,19 +489,11 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
StatusCode: http.StatusOK,
})),
),
server.POST("/oauth/token",
server.Name("oauth_token"),
server.GET("/auth/cursor_dev_session_token",
server.Name("auth_cursor_dev_session_token"),
server.HTTP(),
server.Local(upstream.MockOAuthAction(routeDeps, upstream.CompatRouteConfig{
Name: "oauth_token",
StatusCode: http.StatusOK,
})),
),
server.POST("/aiserver.v1.AuthService/GetEmail",
server.Name("auth_service_get_email"),
server.ConnectUnary(),
server.Local(upstream.MockAuthEmailAction(routeDeps, upstream.CompatRouteConfig{
Name: "auth_service_get_email",
server.Local(upstream.MockDevSessionTokenAction(routeDeps, upstream.CompatRouteConfig{
Name: "auth_cursor_dev_session_token",
StatusCode: http.StatusOK,
})),
),
@@ -476,17 +530,14 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
server.Any("/aiserver.v1.AiService/*",
server.Name("ai_service"),
server.HTTP(),
server.Local(server.HTTPHandlerAction(agentModule.AiHandler)),
server.Local(aiServiceAction),
),
tabServerProcedure("/aiserver.v1.CppService/AvailableModels", "cpp_available_models", server.ConnectUnary(), routeDeps),
tabServerProcedure("/aiserver.v1.CppService/RecordCppFate", "cpp_record_cpp_fate", server.ConnectUnary(), routeDeps),
server.Any("/aiserver.v1.CppService/*",
server.Name("cpp_service"),
server.HTTP(),
server.Local(func(ctx *server.Context) error {
http.NotFound(ctx.Writer, ctx.Request)
return nil
}),
server.Local(fallbackForward),
),
tabServerProcedure("/aiserver.v1.FileSyncService/FSSyncFile", "file_sync_sync_file", server.ConnectUnary(), routeDeps),
tabServerProcedure("/aiserver.v1.FileSyncService/FSIsEnabledForUser", "file_sync_is_enabled_for_user", server.ConnectUnary(), routeDeps),
@@ -495,10 +546,7 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
server.Any("/aiserver.v1.FileSyncService/*",
server.Name("file_sync"),
server.HTTP(),
server.Local(func(ctx *server.Context) error {
http.NotFound(ctx.Writer, ctx.Request)
return nil
}),
server.Local(fallbackForward),
),
server.POST("/aiserver.v1.DashboardService/GetTokenUsage",
server.Name("dashboard_token_usage"),
@@ -530,21 +578,6 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
MockBuilder: upstream.DashboardTeamsMockBuilder,
})),
),
server.POST("/aiserver.v1.DashboardService/GetManagedSkills",
server.Name("dashboard_get_managed_skills"),
server.ConnectUnary(),
server.Local(cursorControlPlaneAction(
host.controlPlaneAuth,
routeDeps,
"dashboard_get_managed_skills",
upstream.MockProtoAction(routeDeps, upstream.CompatRouteConfig{
Name: "dashboard_get_managed_skills",
StatusCode: http.StatusOK,
MockProtoType: "aiserver.v1.GetManagedSkillsResponse",
MockBuilder: upstream.DashboardManagedSkillsMockBuilder,
}),
)),
),
server.POST("/aiserver.v1.DashboardService/GetTeamAdminSettingsOrEmptyIfNotInTeam",
server.Name("dashboard_get_team_admin_settings_or_empty"),
server.ConnectUnary(),
@@ -565,76 +598,6 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
MockBuilder: upstream.EmptyMockBuilder,
})),
),
server.POST("/aiserver.v1.DashboardService/ListMarketplaces",
server.Name("dashboard_list_marketplaces"),
server.ConnectUnary(),
server.Local(upstream.MockProtoAction(routeDeps, upstream.CompatRouteConfig{
Name: "dashboard_list_marketplaces",
StatusCode: http.StatusOK,
MockProtoType: "aiserver.v1.ListMarketplacesResponse",
MockBuilder: upstream.EmptyMockBuilder,
})),
),
server.POST("/aiserver.v1.DashboardService/GetGlobalCommands",
server.Name("dashboard_get_global_commands"),
server.ConnectUnary(),
server.Local(upstream.MockProtoAction(routeDeps, upstream.CompatRouteConfig{
Name: "dashboard_get_global_commands",
StatusCode: http.StatusOK,
MockProtoType: "aiserver.v1.GetGlobalCommandsResponse",
MockBuilder: upstream.EmptyMockBuilder,
})),
),
server.POST("/aiserver.v1.DashboardService/GetEffectiveUserPlugins",
server.Name("dashboard_get_effective_user_plugins"),
server.ConnectUnary(),
server.Local(upstream.MockProtoAction(routeDeps, upstream.CompatRouteConfig{
Name: "dashboard_get_effective_user_plugins",
StatusCode: http.StatusOK,
MockProtoType: "aiserver.v1.GetEffectiveUserPluginsResponse",
MockBuilder: upstream.EmptyMockBuilder,
})),
),
server.POST("/aiserver.v1.DashboardService/RegisterMarketplaceAndPlugins",
server.Name("dashboard_register_marketplace_and_plugins"),
server.ConnectUnary(),
server.Local(upstream.MockProtoAction(routeDeps, upstream.CompatRouteConfig{
Name: "dashboard_register_marketplace_and_plugins",
StatusCode: http.StatusOK,
MockProtoType: "aiserver.v1.RegisterMarketplaceAndPluginsResponse",
MockBuilder: upstream.EmptyMockBuilder,
})),
),
server.POST("/aiserver.v1.DashboardService/GetCliDownloadUrl",
server.Name("dashboard_get_cli_download_url"),
server.ConnectUnary(),
server.Local(upstream.MockProtoAction(routeDeps, upstream.CompatRouteConfig{
Name: "dashboard_get_cli_download_url",
StatusCode: http.StatusOK,
MockProtoType: "aiserver.v1.GetCliDownloadUrlResponse",
MockBuilder: upstream.EmptyMockBuilder,
})),
),
server.POST("/aiserver.v1.DashboardService/GetMe",
server.Name("dashboard_get_me"),
server.ConnectUnary(),
server.Local(upstream.MockProtoAction(routeDeps, upstream.CompatRouteConfig{
Name: "dashboard_get_me",
StatusCode: http.StatusOK,
MockProtoType: "aiserver.v1.GetMeResponse",
MockBuilder: upstream.DashboardGetMeMockBuilder,
})),
),
server.POST("/aiserver.v1.DashboardService/GetUserPrivacyMode",
server.Name("dashboard_user_privacy_mode"),
server.ConnectUnary(),
server.Local(upstream.MockProtoAction(routeDeps, upstream.CompatRouteConfig{
Name: "dashboard_user_privacy_mode",
StatusCode: http.StatusOK,
MockProtoType: "aiserver.v1.GetUserPrivacyModeResponse",
MockBuilder: upstream.DashboardUserPrivacyModeMockBuilder,
})),
),
server.POST("/aiserver.v1.DashboardService/GetPlanInfo",
server.Name("dashboard_plan_info"),
server.ConnectUnary(),
@@ -665,118 +628,36 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
MockBuilder: upstream.DashboardIsOnNewPricingMockBuilder,
})),
),
cursorControlPlaneProcedure("/aiserver.v1.DashboardService/AddMarketplace", "dashboard_add_marketplace", server.ConnectUnary(), host.controlPlaneAuth, routeDeps),
cursorControlPlaneProcedure("/aiserver.v1.DashboardService/AddMcpServersFromPlugin", "dashboard_add_mcp_servers_from_plugin", server.ConnectUnary(), host.controlPlaneAuth, routeDeps),
cursorControlPlaneProcedure("/aiserver.v1.DashboardService/BatchGetPluginMcpConfig", "dashboard_batch_get_plugin_mcp_config", server.ConnectUnary(), host.controlPlaneAuth, routeDeps),
cursorControlPlaneProcedure("/aiserver.v1.DashboardService/GetAvailableMcpServers", "dashboard_get_available_mcp_servers", server.ConnectUnary(), host.controlPlaneAuth, routeDeps),
cursorControlPlaneProcedure("/aiserver.v1.DashboardService/GetEffectiveUserPlugins", "dashboard_get_effective_user_plugins", server.ConnectUnary(), host.controlPlaneAuth, routeDeps),
cursorControlPlaneProcedure("/aiserver.v1.DashboardService/GetPlugin", "dashboard_get_plugin", server.ConnectUnary(), host.controlPlaneAuth, routeDeps),
cursorControlPlaneProcedure("/aiserver.v1.DashboardService/GetPluginMcpConfig", "dashboard_get_plugin_mcp_config", server.ConnectUnary(), host.controlPlaneAuth, routeDeps),
cursorControlPlaneProcedure("/aiserver.v1.DashboardService/InstallUserPlugin", "dashboard_install_user_plugin", server.ConnectUnary(), host.controlPlaneAuth, routeDeps),
cursorControlPlaneProcedure("/aiserver.v1.DashboardService/ListMarketplacePlugins", "dashboard_list_marketplace_plugins", server.ConnectUnary(), host.controlPlaneAuth, routeDeps),
cursorControlPlaneProcedure("/aiserver.v1.DashboardService/ListMarketplaces", "dashboard_list_marketplaces", server.ConnectUnary(), host.controlPlaneAuth, routeDeps),
cursorControlPlaneProcedure("/aiserver.v1.DashboardService/ListUserPluginInstalls", "dashboard_list_user_plugin_installs", server.ConnectUnary(), host.controlPlaneAuth, routeDeps),
cursorControlPlaneProcedure("/aiserver.v1.DashboardService/RefreshMarketplace", "dashboard_refresh_marketplace", server.ConnectUnary(), host.controlPlaneAuth, routeDeps),
cursorControlPlaneProcedure("/aiserver.v1.DashboardService/RegisterMarketplaceAndPlugins", "dashboard_register_marketplace_and_plugins", server.ConnectUnary(), host.controlPlaneAuth, routeDeps),
cursorControlPlaneProcedure("/aiserver.v1.DashboardService/RemoveMarketplace", "dashboard_remove_marketplace", server.ConnectUnary(), host.controlPlaneAuth, routeDeps),
cursorControlPlaneProcedure("/aiserver.v1.DashboardService/ResolvePluginsByRef", "dashboard_resolve_plugins_by_ref", server.ConnectUnary(), host.controlPlaneAuth, routeDeps),
cursorControlPlaneProcedure("/aiserver.v1.DashboardService/UninstallUserPlugin", "dashboard_uninstall_user_plugin", server.ConnectUnary(), host.controlPlaneAuth, routeDeps),
cursorControlPlaneProcedure("/aiserver.v1.DashboardService/UpdateUserPluginInstall", "dashboard_update_user_plugin_install", server.ConnectUnary(), host.controlPlaneAuth, routeDeps),
cursorControlPlaneProcedure("/aiserver.v1.MCPRegistryService/GetKnownServers", "mcp_registry_get_known_servers", server.ConnectUnary(), host.controlPlaneAuth, routeDeps),
server.Any("/aiserver.v1.DashboardService/*",
server.Name("dashboard"),
server.Any("/*",
server.Name("upstream_fallback"),
server.HTTP(),
server.Local(func(ctx *server.Context) error {
http.NotFound(ctx.Writer, ctx.Request)
return nil
}),
server.Local(fallbackForward),
),
// The always-local extension probes NetworkService/IsConnected roughly 10s
// after any slow request starts. A 404 here is treated as "network
// disconnected" and aborts in-flight work (e.g. commit message
// generation) even while the model is still streaming. Always answer OK.
server.POST("/aiserver.v1.NetworkService/IsConnected",
server.Name("network_is_connected"),
server.ConnectUnary(),
server.Local(upstream.MockProtoAction(routeDeps, upstream.CompatRouteConfig{
Name: "network_is_connected",
StatusCode: http.StatusOK,
MockProtoType: "aiserver.v1.IsConnectedResponse",
MockBuilder: upstream.EmptyMockBuilder,
})),
),
server.Any("/aiserver.v1.NetworkService/*",
server.Name("network_service"),
server.HTTP(),
server.Local(func(ctx *server.Context) error {
http.NotFound(ctx.Writer, ctx.Request)
return nil
}),
),
server.Any("/aiserver.v1.InAppAdService/*",
server.Name("in_app_ad"),
server.HTTP(),
server.Local(func(ctx *server.Context) error {
http.NotFound(ctx.Writer, ctx.Request)
return nil
}),
),
server.GET("/auth/full_stripe_profile",
server.Name("auth_full_stripe_profile"),
server.HTTP(),
server.Local(upstream.MockAuthFullStripeProfileAction(routeDeps, upstream.CompatRouteConfig{
Name: "auth_full_stripe_profile",
StatusCode: http.StatusOK,
})),
),
server.GET("/auth/stripe_profile",
server.Name("auth_stripe_profile"),
server.HTTP(),
server.Local(upstream.MockAuthStripeProfileAction(routeDeps, upstream.CompatRouteConfig{
Name: "auth_stripe_profile",
StatusCode: http.StatusOK,
})),
),
server.GET("/auth/has_valid_payment_method",
server.Name("auth_has_valid_payment_method"),
server.HTTP(),
server.Local(upstream.MockJSONAction(routeDeps, upstream.CompatRouteConfig{
Name: "auth_has_valid_payment_method",
StatusCode: http.StatusOK,
JSONBody: map[string]any{
"hasValidPaymentMethod": true,
},
})),
),
server.Any("/auth/poll",
server.Name("auth_poll"),
server.HTTP(),
server.Local(upstream.MockAuthPollAction(routeDeps, upstream.CompatRouteConfig{
Name: "auth_poll",
StatusCode: http.StatusOK,
})),
),
server.POST("/auth/logout",
server.Name("auth_logout"),
server.HTTP(),
server.Local(upstream.FixedStatusAction(routeDeps, upstream.CompatRouteConfig{
Name: "auth_logout",
StatusCode: http.StatusNoContent,
})),
),
server.Any("/auth/*",
server.Name("auth_proxy"),
server.HTTP(),
server.Local(func(ctx *server.Context) error {
http.NotFound(ctx.Writer, ctx.Request)
return nil
}),
),
)
))
return nil
}
func withLocalBackendCORS(next http.Handler) http.Handler {
return http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {
writer.Header().Set("Access-Control-Allow-Origin", "*")
writer.Header().Del("Access-Control-Allow-Credentials")
if strings.EqualFold(request.Method, http.MethodOptions) && strings.TrimSpace(request.Header.Get("Access-Control-Request-Method")) != "" {
writer.Header().Set("Access-Control-Allow-Methods", "GET,POST,PUT,PATCH,DELETE,OPTIONS")
requestedHeaders := strings.TrimSpace(request.Header.Get("Access-Control-Request-Headers"))
if requestedHeaders == "" {
requestedHeaders = "authorization,content-type,x-cursor-client-type"
}
writer.Header().Set("Access-Control-Allow-Headers", requestedHeaders)
writer.Header().Set("Access-Control-Max-Age", "86400")
writer.WriteHeader(http.StatusNoContent)
return
}
next.ServeHTTP(writer, request)
})
}
func repositoryServiceProcedure(pattern string, name string, protocol server.RouteOption, module *forwarder.Module) server.Option {
localAction := server.HTTPHandlerAction(module.RepositoryServiceHandler)
return server.POST(pattern,
@@ -817,46 +698,6 @@ func tabServerProcedure(pattern string, name string, protocol server.RouteOption
)
}
func cursorControlPlaneProcedure(
pattern string,
name string,
protocol server.RouteOption,
authorizationProvider upstream.AuthorizationProvider,
deps upstream.Dependencies,
) server.Option {
notFound := func(ctx *server.Context) error {
http.NotFound(ctx.Writer, ctx.Request)
return nil
}
return server.POST(pattern,
server.Name(name),
protocol,
server.Local(cursorControlPlaneAction(authorizationProvider, deps, name, notFound)),
)
}
func cursorControlPlaneAction(
authorizationProvider upstream.AuthorizationProvider,
deps upstream.Dependencies,
name string,
fallback server.HandlerFunc,
) server.HandlerFunc {
forward := upstream.AuthenticatedForwardAction(deps, upstream.CompatRouteConfig{Name: name}, authorizationProvider)
return func(ctx *server.Context) error {
if authorizationProvider == nil || !authorizationProvider.SignedIn() {
return fallback(ctx)
}
if ctx == nil || ctx.Request == nil || ctx.Request.URL == nil {
return fmt.Errorf("Cursor 控制面请求上下文无效")
}
targetURL := *ctx.Request.URL
targetURL.Scheme = "https"
targetURL.Host = "api2.cursor.sh:443"
ctx.UpstreamURL = &targetURL
return forward(ctx)
}
}
type serverSystemSettings struct {
configs *serverconfig.Manager
}
+175
View File
@@ -0,0 +1,175 @@
package backend
import (
"context"
"crypto/tls"
"crypto/x509"
"encoding/json"
"net"
"net/http"
"net/http/httptest"
"path/filepath"
"testing"
"time"
"cursor/gen/aiserverv1"
serverconfig "cursor/internal/backend/server/config"
"cursor/internal/certs"
"google.golang.org/protobuf/proto"
)
func TestHostServesDevLoginAndLocalTeamsRoute(t *testing.T) {
store := serverconfig.NewStore(filepath.Join(t.TempDir(), "config.yaml"), t.TempDir())
host, err := NewHost(store)
if err != nil {
t.Fatalf("new host: %v", err)
}
loginRequest := httptest.NewRequest(http.MethodGet, "http://local/auth/cursor_dev_session_token?plan=enterprise&email=enterprise%40example.com", nil)
loginRecorder := httptest.NewRecorder()
host.mux.ServeHTTP(loginRecorder, loginRequest)
if loginRecorder.Code != http.StatusOK {
t.Fatalf("dev login status: got %d, want %d; body=%s", loginRecorder.Code, http.StatusOK, loginRecorder.Body.String())
}
var loginResponse struct {
AccessToken string `json:"accessToken"`
}
if err := json.Unmarshal(loginRecorder.Body.Bytes(), &loginResponse); err != nil {
t.Fatalf("decode dev login: %v", err)
}
if loginResponse.AccessToken == "" {
t.Fatal("dev login returned an empty access token")
}
teamsRequest := httptest.NewRequest(http.MethodPost, "http://local/aiserver.v1.DashboardService/GetTeams", nil)
teamsRequest.Header.Set("Authorization", "Bearer "+loginResponse.AccessToken)
teamsRecorder := httptest.NewRecorder()
host.mux.ServeHTTP(teamsRecorder, teamsRequest)
if teamsRecorder.Code != http.StatusOK {
t.Fatalf("teams status: got %d, want %d", teamsRecorder.Code, http.StatusOK)
}
teams := &aiserverv1.GetTeamsResponse{}
if err := proto.Unmarshal(teamsRecorder.Body.Bytes(), teams); err != nil {
t.Fatalf("decode teams response: %v", err)
}
if len(teams.GetTeams()) != 1 || !teams.GetTeams()[0].GetIsEnterprise() {
t.Fatalf("unexpected teams response: %v", teams.GetTeams())
}
}
func TestHostAllowsWildcardCORS(t *testing.T) {
store := serverconfig.NewStore(filepath.Join(t.TempDir(), "config.yaml"), t.TempDir())
host, err := NewHost(store)
if err != nil {
t.Fatalf("new host: %v", err)
}
preflightRequest := httptest.NewRequest(http.MethodOptions, "http://local/auth/cursor_dev_session_token?plan=free", nil)
preflightRequest.Header.Set("Origin", "vscode-file://vscode-app")
preflightRequest.Header.Set("Access-Control-Request-Method", http.MethodGet)
preflightRequest.Header.Set("Access-Control-Request-Headers", "x-cursor-client-type")
preflightRecorder := httptest.NewRecorder()
host.mux.ServeHTTP(preflightRecorder, preflightRequest)
if preflightRecorder.Code != http.StatusNoContent {
t.Fatalf("preflight status: got %d, want %d", preflightRecorder.Code, http.StatusNoContent)
}
if got := preflightRecorder.Header().Get("Access-Control-Allow-Origin"); got != "*" {
t.Fatalf("preflight allow origin: got %q", got)
}
if got := preflightRecorder.Header().Get("Access-Control-Allow-Credentials"); got != "" {
t.Fatalf("preflight allow credentials: got %q, want empty", got)
}
if got := preflightRecorder.Header().Get("Access-Control-Allow-Headers"); got != "x-cursor-client-type" {
t.Fatalf("preflight allow headers: got %q", got)
}
loginRequest := httptest.NewRequest(http.MethodGet, "http://local/auth/cursor_dev_session_token?plan=free", nil)
loginRequest.Header.Set("Origin", "vscode-file://vscode-app")
loginRequest.Header.Set("x-cursor-client-type", "ide")
loginRecorder := httptest.NewRecorder()
host.mux.ServeHTTP(loginRecorder, loginRequest)
if loginRecorder.Code != http.StatusOK {
t.Fatalf("dev login status: got %d, want %d", loginRecorder.Code, http.StatusOK)
}
if got := loginRecorder.Header().Get("Access-Control-Allow-Origin"); got != "*" {
t.Fatalf("dev login allow origin: got %q", got)
}
}
func TestHostAllowsRemoteWebOriginWithWildcard(t *testing.T) {
store := serverconfig.NewStore(filepath.Join(t.TempDir(), "config.yaml"), t.TempDir())
host, err := NewHost(store)
if err != nil {
t.Fatalf("new host: %v", err)
}
request := httptest.NewRequest(http.MethodOptions, "http://local/auth/cursor_dev_session_token", nil)
request.Header.Set("Origin", "https://example.com")
request.Header.Set("Access-Control-Request-Method", http.MethodGet)
recorder := httptest.NewRecorder()
host.mux.ServeHTTP(recorder, request)
if got := recorder.Header().Get("Access-Control-Allow-Origin"); got != "*" {
t.Fatalf("remote origin allow origin: got %q, want wildcard", got)
}
}
func TestHostServesDevLoginOverTrustedLocalhostTLS(t *testing.T) {
listener, err := net.Listen("tcp", "127.0.0.1:0")
if err != nil {
t.Fatalf("reserve backend port: %v", err)
}
listenAddr := listener.Addr().String()
if err := listener.Close(); err != nil {
t.Fatalf("release backend port: %v", err)
}
store := serverconfig.NewStore(filepath.Join(t.TempDir(), "config.yaml"), t.TempDir())
config := serverconfig.DefaultConfig()
config.BackendListenAddr = listenAddr
if _, err := store.Save(context.Background(), config); err != nil {
t.Fatalf("save backend config: %v", err)
}
certificateManager, err := certs.NewEmbeddedManager()
if err != nil {
t.Fatalf("new certificate manager: %v", err)
}
serverCertificate, err := certificateManager.CertificateForServerName("localhost")
if err != nil {
t.Fatalf("create localhost certificate: %v", err)
}
host, err := NewHost(store, WithTLSCertificate(serverCertificate))
if err != nil {
t.Fatalf("new TLS host: %v", err)
}
if err := host.Start(); err != nil {
t.Fatalf("start TLS host: %v", err)
}
defer func() {
stopContext, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
if err := host.Stop(stopContext); err != nil {
t.Errorf("stop TLS host: %v", err)
}
}()
caCertificate, err := certificateManager.CATLSCertificate()
if err != nil {
t.Fatalf("load CA certificate: %v", err)
}
roots := x509.NewCertPool()
roots.AddCert(caCertificate.Leaf)
client := &http.Client{Transport: &http.Transport{TLSClientConfig: &tls.Config{
MinVersion: tls.VersionTLS12,
RootCAs: roots,
ServerName: "localhost",
}}}
response, err := client.Get(host.BaseURL() + "/auth/cursor_dev_session_token?plan=pro&trial=true")
if err != nil {
t.Fatalf("request dev login over TLS: %v", err)
}
defer response.Body.Close()
if response.StatusCode != http.StatusOK {
t.Fatalf("dev login TLS status: got %d, want %d", response.StatusCode, http.StatusOK)
}
}
+127
View File
@@ -0,0 +1,127 @@
package backend
import (
"io"
"net/http"
"net/http/httptest"
"path/filepath"
"strings"
"sync/atomic"
"testing"
"cursor/internal/backend/server"
serverconfig "cursor/internal/backend/server/config"
)
func TestHostForwardsUnhandledRoutesToOriginalUpstream(t *testing.T) {
var requestCount atomic.Int32
upstreamServer := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {
requestCount.Add(1)
body, err := io.ReadAll(request.Body)
if err != nil {
t.Errorf("read upstream request body: %v", err)
}
writer.Header().Set("X-Upstream-Path", request.URL.RequestURI())
writer.WriteHeader(http.StatusMultiStatus)
_, _ = writer.Write(body)
}))
defer upstreamServer.Close()
store := serverconfig.NewStore(filepath.Join(t.TempDir(), "config.yaml"), t.TempDir())
host, err := NewHost(store)
if err != nil {
t.Fatalf("new host: %v", err)
}
testCases := []struct {
name string
method string
path string
}{
{name: "managed skills", path: "/aiserver.v1.DashboardService/GetManagedSkills?source=skills"},
{name: "effective plugins", path: "/aiserver.v1.DashboardService/GetEffectiveUserPlugins?source=plugins"},
{name: "MCP registry", path: "/aiserver.v1.MCPRegistryService/GetKnownServers?source=mcp"},
{name: "auth poll", path: "/auth/poll?uuid=local-login&verifier=test"},
{name: "OAuth token", path: "/oauth/token"},
{name: "auth email", path: "/aiserver.v1.AuthService/GetEmail"},
{name: "dashboard me", path: "/aiserver.v1.DashboardService/GetMe"},
{name: "full stripe profile", method: http.MethodGet, path: "/auth/full_stripe_profile"},
{name: "stripe profile", method: http.MethodGet, path: "/auth/stripe_profile"},
{name: "valid payment method", method: http.MethodGet, path: "/auth/has_valid_payment_method"},
{name: "auth logout", path: "/auth/logout"},
{name: "dashboard global commands", path: "/aiserver.v1.DashboardService/GetGlobalCommands"},
{name: "dashboard CLI download", path: "/aiserver.v1.DashboardService/GetCliDownloadUrl"},
{name: "dashboard privacy mode", path: "/aiserver.v1.DashboardService/GetUserPrivacyMode"},
{name: "service catch-all", path: "/aiserver.v1.NetworkService/UnknownProcedure?source=network"},
{name: "AI handler miss", path: "/aiserver.v1.AiService/UnknownProcedure?source=ai"},
{name: "global miss", path: "/unknown/service/path?source=global"},
}
for _, testCase := range testCases {
t.Run(testCase.name, func(t *testing.T) {
method := testCase.method
if method == "" {
method = http.MethodPost
}
body := "payload-" + testCase.name
request := httptest.NewRequest(method, "http://localhost:8000"+testCase.path, strings.NewReader(body))
request.Header.Set(server.HeaderServerUpstreamURL, upstreamServer.URL+testCase.path)
recorder := httptest.NewRecorder()
host.mux.ServeHTTP(recorder, request)
if got := recorder.Code; got != http.StatusMultiStatus {
t.Fatalf("status: got %d, want %d; body=%s", got, http.StatusMultiStatus, recorder.Body.String())
}
if got := recorder.Header().Get("X-Upstream-Path"); got != testCase.path {
t.Fatalf("upstream path: got %q, want %q", got, testCase.path)
}
wantBody := body
if method == http.MethodGet {
wantBody = ""
}
if got := recorder.Body.String(); got != wantBody {
t.Fatalf("response body: got %q, want %q", got, wantBody)
}
})
}
requestsBeforeHealthCheck := requestCount.Load()
healthRequest := httptest.NewRequest(http.MethodGet, "http://localhost:8000"+healthPath, nil)
healthRecorder := httptest.NewRecorder()
host.mux.ServeHTTP(healthRecorder, healthRequest)
if got := healthRecorder.Code; got != http.StatusOK {
t.Fatalf("health status: got %d, want %d", got, http.StatusOK)
}
if got := requestCount.Load(); got != requestsBeforeHealthCheck {
t.Fatalf("local health route unexpectedly reached upstream: requests before=%d after=%d", requestsBeforeHealthCheck, got)
}
}
func TestHostFallbackKeepsWildcardCORSWhenUpstreamReturnsCORSHeaders(t *testing.T) {
upstreamServer := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {
writer.Header().Set("Access-Control-Allow-Origin", "vscode-file://vscode-app")
writer.Header().Set("Access-Control-Allow-Credentials", "true")
writer.WriteHeader(http.StatusOK)
}))
defer upstreamServer.Close()
store := serverconfig.NewStore(filepath.Join(t.TempDir(), "config.yaml"), t.TempDir())
host, err := NewHost(store)
if err != nil {
t.Fatalf("new host: %v", err)
}
request := httptest.NewRequest(http.MethodGet, "http://localhost:8000/auth/poll?uuid=test", nil)
request.Header.Set("Origin", "vscode-file://vscode-app")
request.Header.Set(server.HeaderServerUpstreamURL, upstreamServer.URL+request.URL.RequestURI())
recorder := httptest.NewRecorder()
host.mux.ServeHTTP(recorder, request)
if got := recorder.Header().Values("Access-Control-Allow-Origin"); len(got) != 1 || got[0] != "*" {
t.Fatalf("allow origin values: got %q, want [*]", got)
}
if got := recorder.Header().Get("Access-Control-Allow-Credentials"); got != "" {
t.Fatalf("allow credentials: got %q, want empty", got)
}
}
+9 -11
View File
@@ -13,7 +13,7 @@ import (
)
const (
DefaultBackendListenAddr = "127.0.0.1:18090"
DefaultBackendListenAddr = "127.0.0.1:8000"
DefaultProxyListenAddr = "127.0.0.1:18080"
DefaultFrontendBaseURL = "http://127.0.0.1"
DefaultProviderStreamIdleTimeoutSeconds = 240
@@ -141,8 +141,8 @@ func NormalizeModelAdapterConfigs(input []ModelAdapterConfig) ([]ModelAdapterCon
return nil, errors.New("模型适配器 tooltipData 不能为空")
case next.ModelID == "":
return nil, errors.New("模型适配器 modelID 不能为空")
case next.Type == "openai" && !isSupportedReasoningEffort(next.ReasoningEffort):
return nil, errors.New("模型适配器 reasoningEffort 仅支持空值、low、medium、high、xhigh、max")
case next.Type == "openai" && next.ReasoningEffort == "":
return nil, errors.New("模型适配器 reasoningEffort 仅支持 low、medium、high、xhigh、max")
case next.Type == "openai" && next.OpenAIEndpoint == "":
return nil, errors.New("模型适配器 openAIEndpoint 仅支持 /v1/responses、/v1/chat/completions 或 /custom(自定义路径)")
case next.Type == "openai" && next.OpenAIExtraParamsEnabled:
@@ -224,15 +224,13 @@ func validateHeadersJSON(value string) error {
}
func normalizeReasoningEffort(value string) string {
return strings.ToLower(strings.TrimSpace(value))
}
func isSupportedReasoningEffort(value string) bool {
switch value {
case "", "low", "medium", "high", "xhigh", "max":
return true
switch strings.ToLower(strings.TrimSpace(value)) {
case "", "medium":
return "medium"
case "low", "high", "xhigh", "max":
return strings.ToLower(strings.TrimSpace(value))
default:
return false
return ""
}
}
@@ -58,25 +58,3 @@ func TestNormalizeModelAdapterConfigsUsesStableExplicitSort(t *testing.T) {
}
}
}
func TestNormalizeModelAdapterConfigsAllowsBlankReasoningEffort(t *testing.T) {
adapter := testModelAdapter("non-reasoning-model", 1)
adapter.ReasoningEffort = ""
adapters, err := NormalizeModelAdapterConfigs([]ModelAdapterConfig{adapter})
if err != nil {
t.Fatalf("NormalizeModelAdapterConfigs returned error: %v", err)
}
if got := adapters[0].ReasoningEffort; got != "" {
t.Fatalf("ReasoningEffort = %q, want blank", got)
}
}
func TestNormalizeModelAdapterConfigsRejectsUnknownReasoningEffort(t *testing.T) {
adapter := testModelAdapter("invalid-reasoning-effort", 1)
adapter.ReasoningEffort = "unsupported"
if _, err := NormalizeModelAdapterConfigs([]ModelAdapterConfig{adapter}); err == nil {
t.Fatal("NormalizeModelAdapterConfigs should reject an unknown reasoning effort")
}
}
+22 -88
View File
@@ -14,12 +14,13 @@ import (
type CompatRouteConfig struct {
Name string
StatusCode int
JSONBody map[string]any
MockProtoType string
MockBuilder func(*RequestContext) (map[string]any, error)
ConsoleLog bool
}
const DefaultCursorUpstreamBaseURL = "https://api2.cursor.sh:443"
func ForwardAction(deps Dependencies, cfg CompatRouteConfig) server.HandlerFunc {
return func(ctx *server.Context) error {
reqCtx, route, err := newCompatRouteObjects(ctx, deps, cfg)
@@ -30,31 +31,27 @@ func ForwardAction(deps Dependencies, cfg CompatRouteConfig) server.HandlerFunc
}
}
// AuthenticatedForwardAction forwards a Cursor control-plane request with the
// independent desktop account after the local-mode identity rewrite has run.
func AuthenticatedForwardAction(deps Dependencies, cfg CompatRouteConfig, authorizationProvider AuthorizationProvider) server.HandlerFunc {
// FallbackForwardAction preserves an MITM request's original upstream URL. A
// native request has no original host metadata, so it is resolved against the
// configured default upstream while retaining its path and query string.
func FallbackForwardAction(deps Dependencies, cfg CompatRouteConfig, defaultBaseURL string) server.HandlerFunc {
forward := ForwardAction(deps, cfg)
return func(ctx *server.Context) error {
reqCtx, _, err := newCompatRouteObjects(ctx, deps, cfg)
if err != nil {
return err
if ctx == nil || ctx.Request == nil || ctx.Request.URL == nil {
return fmt.Errorf("fallback upstream request context is invalid")
}
if reqCtx == nil || reqCtx.Request == nil {
return fmt.Errorf("Cursor 控制面请求上下文无效")
if ctx.UpstreamURL == nil {
baseURL, err := ParseAndValidateRawURL(defaultBaseURL)
if err != nil {
return fmt.Errorf("parse fallback upstream URL: %w", err)
}
targetURL := *ctx.Request.URL
targetURL.Scheme = baseURL.Scheme
targetURL.Host = baseURL.Host
targetURL.User = baseURL.User
ctx.UpstreamURL = &targetURL
}
if authorizationProvider == nil {
return fmt.Errorf("Cursor 账号服务未初始化")
}
authorization, err := authorizationProvider.Authorization(reqCtx.Request.Context())
if err != nil {
return err
}
_, err = ForwardToUpstream(reqCtx, ForwardOptions{
PatchHeaders: func(headers http.Header) {
headers.Set("Authorization", authorization)
headers.Set("x-cursor-checksum", BuildCursorChecksum(authorization))
},
})
return err
return forward(ctx)
}
}
@@ -68,63 +65,13 @@ func FixedStatusAction(deps Dependencies, cfg CompatRouteConfig) server.HandlerF
}
}
func MockJSONAction(deps Dependencies, cfg CompatRouteConfig) server.HandlerFunc {
func MockDevSessionTokenAction(deps Dependencies, cfg CompatRouteConfig) server.HandlerFunc {
return func(ctx *server.Context) error {
reqCtx, route, err := newCompatRouteObjects(ctx, deps, cfg)
if err != nil {
return err
}
return handleMockJSON(reqCtx, route)
}
}
func MockOAuthAction(deps Dependencies, cfg CompatRouteConfig) server.HandlerFunc {
return func(ctx *server.Context) error {
reqCtx, route, err := newCompatRouteObjects(ctx, deps, cfg)
if err != nil {
return err
}
return handleMockOAuth(reqCtx, route)
}
}
func MockAuthFullStripeProfileAction(deps Dependencies, cfg CompatRouteConfig) server.HandlerFunc {
return func(ctx *server.Context) error {
reqCtx, route, err := newCompatRouteObjects(ctx, deps, cfg)
if err != nil {
return err
}
return handleMockAuthFullStripeProfile(reqCtx, route)
}
}
func MockAuthStripeProfileAction(deps Dependencies, cfg CompatRouteConfig) server.HandlerFunc {
return func(ctx *server.Context) error {
reqCtx, route, err := newCompatRouteObjects(ctx, deps, cfg)
if err != nil {
return err
}
return handleMockAuthStripeProfile(reqCtx, route)
}
}
func MockAuthPollAction(deps Dependencies, cfg CompatRouteConfig) server.HandlerFunc {
return func(ctx *server.Context) error {
reqCtx, route, err := newCompatRouteObjects(ctx, deps, cfg)
if err != nil {
return err
}
return handleMockAuthPoll(reqCtx, route)
}
}
func MockAuthEmailAction(deps Dependencies, cfg CompatRouteConfig) server.HandlerFunc {
return func(ctx *server.Context) error {
reqCtx, route, err := newCompatRouteObjects(ctx, deps, cfg)
if err != nil {
return err
}
return handleMockAuthEmail(reqCtx, route)
return handleMockDevSessionToken(reqCtx, route)
}
}
@@ -169,7 +116,6 @@ func newCompatRouteObjects(ctx *server.Context, deps Dependencies, cfg CompatRou
Name: cfg.Name,
Pattern: ctx.Request.URL.Path,
StatusCode: cfg.StatusCode,
JSONBody: cfg.JSONBody,
MockProtoType: cfg.MockProtoType,
MockPayloadBuilder: cfg.MockBuilder,
ConsoleLog: cfg.ConsoleLog,
@@ -221,10 +167,6 @@ func DashboardTeamsMockBuilder(reqCtx *RequestContext) (map[string]any, error) {
return buildDashboardTeamsPayload(reqCtx)
}
func DashboardManagedSkillsMockBuilder(reqCtx *RequestContext) (map[string]any, error) {
return buildDashboardManagedSkillsPayload(reqCtx)
}
// EmptyMockBuilder возвращает пустой proto-ответ для ручек, где клиенту
// достаточно успешного "пусто": нет team-настроек, нет репозиториев,
// нет маркетплейсов/плагинов/команд, телеметрия принята без обработки.
@@ -237,14 +179,6 @@ func SubmitLogsMockBuilder(reqCtx *RequestContext) (map[string]any, error) {
return map[string]any{"success": true}, nil
}
func DashboardGetMeMockBuilder(reqCtx *RequestContext) (map[string]any, error) {
return buildDashboardGetMePayload(reqCtx)
}
func DashboardUserPrivacyModeMockBuilder(reqCtx *RequestContext) (map[string]any, error) {
return buildDashboardUserPrivacyModePayload(reqCtx)
}
func DashboardPlanInfoMockBuilder(reqCtx *RequestContext) (map[string]any, error) {
return buildDashboardPlanInfoPayload(reqCtx)
}
+193
View File
@@ -0,0 +1,193 @@
package upstream
import (
"encoding/base64"
"encoding/json"
"fmt"
"net/http"
"strconv"
"strings"
"time"
legacyruntime "cursor/internal/runtime"
)
const (
localDevDefaultPlan = "ultra"
localDevTokenLifetime = 10 * 365 * 24 * time.Hour
localDevSubscriptionActive = "active"
)
var localDevPlans = map[string]struct{}{
"free": {},
"pro": {},
"pro_plus": {},
"ultra": {},
"enterprise": {},
}
type localDevSessionClaims struct {
Subject string `json:"sub"`
Email string `json:"email"`
Plan string `json:"cursor_local_plan"`
Trial bool `json:"cursor_local_trial"`
TokenType string `json:"type"`
Issuer string `json:"iss"`
Scope string `json:"scope"`
IssuedAt int64 `json:"iat"`
ExpiresAt int64 `json:"exp"`
}
func handleMockDevSessionToken(reqCtx *RequestContext, route *Route) error {
_ = route
if reqCtx == nil || reqCtx.Request == nil || reqCtx.ResponseWriter == nil {
return fmt.Errorf("dev session request context is invalid")
}
plan, trial, email, err := parseLocalDevSessionQuery(reqCtx.Request)
if err != nil {
writeJSONError(reqCtx.ResponseWriter, http.StatusBadRequest, err.Error())
return nil
}
token, claims, err := buildLocalDevSessionToken(plan, trial, email, time.Now())
if err != nil {
return err
}
responseBody, err := marshalJSONBody(map[string]any{
"accessToken": token,
"refreshToken": token,
"authId": claims.Subject,
})
if err != nil {
return err
}
reqCtx.ResponseWriter.Header().Set("content-type", "application/json")
reqCtx.ResponseWriter.WriteHeader(http.StatusOK)
_, _ = reqCtx.ResponseWriter.Write(responseBody)
return nil
}
func parseLocalDevSessionQuery(request *http.Request) (string, bool, string, error) {
plan := localDevDefaultPlan
email := legacyruntime.InjectAccountEmail
if request == nil || request.URL == nil {
return plan, false, email, nil
}
query := request.URL.Query()
if requestedPlan := strings.TrimSpace(query.Get("plan")); requestedPlan != "" {
plan = requestedPlan
}
if _, ok := localDevPlans[plan]; !ok {
return "", false, "", fmt.Errorf("unsupported dev plan %q", plan)
}
trial := false
if rawTrial := strings.TrimSpace(query.Get("trial")); rawTrial != "" {
parsed, err := strconv.ParseBool(rawTrial)
if err != nil {
return "", false, "", fmt.Errorf("invalid trial value %q", rawTrial)
}
trial = parsed
}
if trial && plan != "pro" && plan != "pro_plus" {
return "", false, "", fmt.Errorf("trial is only supported for pro and pro_plus")
}
if requestedEmail := strings.TrimSpace(query.Get("email")); requestedEmail != "" {
email = requestedEmail
}
return plan, trial, email, nil
}
func buildLocalDevSessionToken(plan string, trial bool, email string, now time.Time) (string, localDevSessionClaims, error) {
authID := "local-dev-" + strings.ReplaceAll(plan, "_", "-")
if trial {
authID += "-trial"
}
claims := localDevSessionClaims{
Subject: authID,
Email: strings.TrimSpace(email),
Plan: plan,
Trial: trial,
TokenType: "session",
Issuer: "cursor-local-backend",
Scope: "openid profile email",
IssuedAt: now.Unix(),
ExpiresAt: now.Add(localDevTokenLifetime).Unix(),
}
headerJSON, err := json.Marshal(map[string]string{"alg": "HS256", "typ": "JWT"})
if err != nil {
return "", localDevSessionClaims{}, err
}
claimsJSON, err := json.Marshal(claims)
if err != nil {
return "", localDevSessionClaims{}, err
}
encode := base64.RawURLEncoding.EncodeToString
token := encode(headerJSON) + "." + encode(claimsJSON) + ".local-dev"
return token, claims, nil
}
func localDevClaimsFromRequest(reqCtx *RequestContext) (localDevSessionClaims, bool) {
if reqCtx == nil {
return localDevSessionClaims{}, false
}
return localDevClaimsFromAuthorization(reqCtx.Headers.Get("authorization"))
}
func localDevClaimsFromAuthorization(authorization string) (localDevSessionClaims, bool) {
authorization = strings.TrimSpace(authorization)
if len(authorization) >= len("Bearer ") && strings.EqualFold(authorization[:len("Bearer ")], "Bearer ") {
authorization = strings.TrimSpace(authorization[len("Bearer "):])
}
parts := strings.Split(authorization, ".")
if len(parts) != 3 {
return localDevSessionClaims{}, false
}
payload, err := base64.RawURLEncoding.DecodeString(parts[1])
if err != nil {
return localDevSessionClaims{}, false
}
claims := localDevSessionClaims{}
if err := json.Unmarshal(payload, &claims); err != nil {
return localDevSessionClaims{}, false
}
if claims.Issuer != "cursor-local-backend" {
return localDevSessionClaims{}, false
}
if _, ok := localDevPlans[claims.Plan]; !ok || strings.TrimSpace(claims.Subject) == "" {
return localDevSessionClaims{}, false
}
return claims, true
}
func localDevPlanFromRequest(reqCtx *RequestContext) string {
if claims, ok := localDevClaimsFromRequest(reqCtx); ok {
return claims.Plan
}
return localDevDefaultPlan
}
func localDevPlanDetails(plan string) (string, int) {
switch plan {
case "free":
return "Free Plan", 0
case "pro":
return "Pro Plan", 2000
case "pro_plus":
return "Pro+ Plan", 6000
case "enterprise":
return "Enterprise Plan", 0
default:
return "Ultra Plan", localUltraPlanIncludedCents
}
}
func writeJSONError(writer http.ResponseWriter, statusCode int, message string) {
writer.Header().Set("content-type", "application/json")
writer.WriteHeader(statusCode)
payload, _ := json.Marshal(map[string]string{"error": message})
_, _ = writer.Write(payload)
}
@@ -0,0 +1,138 @@
package upstream
import (
"encoding/json"
"net/http"
"net/http/httptest"
"strings"
"testing"
"time"
"cursor/gen/aiserverv1"
"cursor/internal/backend/server"
"google.golang.org/protobuf/proto"
)
func TestMockDevSessionTokenActionSupportsCursorDevLoginModes(t *testing.T) {
testCases := []struct {
name string
query string
plan string
trial bool
}{
{name: "default", query: "", plan: "ultra"},
{name: "free", query: "?plan=free", plan: "free"},
{name: "pro trial", query: "?plan=pro&trial=true", plan: "pro", trial: true},
{name: "pro", query: "?plan=pro", plan: "pro"},
{name: "pro plus trial", query: "?plan=pro_plus&trial=true", plan: "pro_plus", trial: true},
{name: "pro plus", query: "?plan=pro_plus", plan: "pro_plus"},
{name: "ultra", query: "?plan=ultra", plan: "ultra"},
{name: "enterprise", query: "?plan=enterprise", plan: "enterprise"},
}
for _, testCase := range testCases {
t.Run(testCase.name, func(t *testing.T) {
request := httptest.NewRequest(http.MethodGet, "http://local/auth/cursor_dev_session_token"+testCase.query, nil)
recorder := httptest.NewRecorder()
handler := MockDevSessionTokenAction(Dependencies{}, CompatRouteConfig{Name: "dev_login", StatusCode: http.StatusOK})
if err := handler(&server.Context{Writer: recorder, Request: request}); err != nil {
t.Fatalf("dev login handler: %v", err)
}
if recorder.Code != http.StatusOK {
t.Fatalf("status: got %d, want %d; body=%s", recorder.Code, http.StatusOK, recorder.Body.String())
}
var response struct {
AccessToken string `json:"accessToken"`
RefreshToken string `json:"refreshToken"`
AuthID string `json:"authId"`
}
if err := json.Unmarshal(recorder.Body.Bytes(), &response); err != nil {
t.Fatalf("decode response: %v", err)
}
if response.AccessToken == "" || response.RefreshToken != response.AccessToken {
t.Fatalf("unexpected tokens: access=%q refresh=%q", response.AccessToken, response.RefreshToken)
}
claims, ok := localDevClaimsFromAuthorization("Bearer " + response.AccessToken)
if !ok {
t.Fatal("response access token is not a local dev JWT")
}
if claims.Plan != testCase.plan || claims.Trial != testCase.trial {
t.Fatalf("claims: got plan=%q trial=%v, want plan=%q trial=%v", claims.Plan, claims.Trial, testCase.plan, testCase.trial)
}
if response.AuthID != claims.Subject || claims.ExpiresAt <= time.Now().Unix() {
t.Fatalf("unexpected identity claims: response=%+v claims=%+v", response, claims)
}
})
}
}
func TestMockDevSessionTokenActionUsesRequestedEmail(t *testing.T) {
request := httptest.NewRequest(http.MethodGet, "http://local/auth/cursor_dev_session_token?plan=pro&email=dev%2Bcursor%40example.com", nil)
recorder := httptest.NewRecorder()
handler := MockDevSessionTokenAction(Dependencies{}, CompatRouteConfig{Name: "dev_login", StatusCode: http.StatusOK})
if err := handler(&server.Context{Writer: recorder, Request: request}); err != nil {
t.Fatalf("dev login handler: %v", err)
}
var response map[string]string
if err := json.Unmarshal(recorder.Body.Bytes(), &response); err != nil {
t.Fatalf("decode response: %v", err)
}
claims, ok := localDevClaimsFromAuthorization(response["accessToken"])
if !ok || claims.Email != "dev+cursor@example.com" {
t.Fatalf("unexpected email claims: %+v", claims)
}
}
func TestMockDevSessionTokenActionRejectsUnsupportedOptions(t *testing.T) {
for _, query := range []string{"?plan=business", "?plan=ultra&trial=true", "?plan=pro&trial=maybe"} {
request := httptest.NewRequest(http.MethodGet, "http://local/auth/cursor_dev_session_token"+query, nil)
recorder := httptest.NewRecorder()
handler := MockDevSessionTokenAction(Dependencies{}, CompatRouteConfig{Name: "dev_login", StatusCode: http.StatusOK})
if err := handler(&server.Context{Writer: recorder, Request: request}); err != nil {
t.Fatalf("dev login handler for %q: %v", query, err)
}
if recorder.Code != http.StatusBadRequest {
t.Fatalf("status for %q: got %d, want %d", query, recorder.Code, http.StatusBadRequest)
}
}
}
func TestEnterpriseDevSessionProvidesBillableTeam(t *testing.T) {
token, _, err := buildLocalDevSessionToken("enterprise", false, "enterprise@example.com", time.Now())
if err != nil {
t.Fatalf("build token: %v", err)
}
reqCtx := authRequestContext(http.MethodPost, "/aiserver.v1.DashboardService/GetTeams", "", token)
payload, err := buildDashboardTeamsPayload(reqCtx)
if err != nil {
t.Fatalf("build teams: %v", err)
}
encoded, err := encodeMockProto("aiserver.v1.GetTeamsResponse", payload)
if err != nil {
t.Fatalf("encode teams: %v", err)
}
response := &aiserverv1.GetTeamsResponse{}
if err := proto.Unmarshal(encoded, response); err != nil {
t.Fatalf("decode teams: %v", err)
}
if len(response.Teams) != 1 || !response.Teams[0].GetHasBilling() || response.Teams[0].GetSeats() == 0 || !response.Teams[0].GetIsEnterprise() {
t.Fatalf("unexpected enterprise teams: %+v", response.Teams)
}
}
func authRequestContext(method string, path string, body string, token string) *RequestContext {
request := httptest.NewRequest(method, "http://local"+path, strings.NewReader(body))
if token != "" {
request.Header.Set("Authorization", "Bearer "+token)
}
return &RequestContext{
ResponseWriter: httptest.NewRecorder(),
Request: request,
Method: method,
Headers: request.Header.Clone(),
RequestBody: []byte(body),
}
}
+4 -161
View File
@@ -2,23 +2,18 @@ package upstream
import (
"bytes"
"crypto/sha256"
"encoding/base64"
"encoding/json"
"fmt"
"io"
"math/big"
"net/http"
"net/url"
"strconv"
"strings"
"time"
"cursor/gen/agentv1"
"cursor/gen/aiserverv1"
"cursor/internal/logger"
"cursor/internal/netproxy"
legacyruntime "cursor/internal/runtime"
"google.golang.org/protobuf/encoding/protojson"
"google.golang.org/protobuf/proto"
@@ -87,14 +82,6 @@ func buildUpstreamRequest(reqCtx *RequestContext, body []byte, options ForwardOp
}
upstreamRequest.Host = reqCtx.TargetURL.Host
if shouldRewriteHost(reqCtx.TargetURL.Hostname()) {
auth := formatBearerAuthorization(legacyruntime.LocalRelayToken)
if auth == "" {
return nil, nil, legacyruntime.ErrInvalidSystemSetting
}
upstreamRequest.Header.Set("Authorization", auth)
upstreamRequest.Header.Set("x-cursor-checksum", BuildCursorChecksum(auth))
}
if options.PatchHeaders != nil {
options.PatchHeaders(upstreamRequest.Header)
}
@@ -167,61 +154,21 @@ func copyRequestHeadersForUpstream(target http.Header, source http.Header) {
}
func copyResponseHeadersToClient(target http.Header, source http.Header) {
localWildcardCORS := target.Get("Access-Control-Allow-Origin") == "*"
for key, values := range source {
lowerKey := strings.ToLower(key)
if _, exists := hopByHopHeaders[lowerKey]; exists {
continue
}
if localWildcardCORS && (lowerKey == "access-control-allow-origin" || lowerKey == "access-control-allow-credentials") {
continue
}
for _, value := range values {
target.Add(key, value)
}
}
}
func shouldRewriteHost(host string) bool {
normalized := strings.TrimSuffix(strings.ToLower(strings.TrimSpace(host)), ".")
if normalized == "" {
return false
}
return normalized == "cursor.sh" || strings.HasSuffix(normalized, ".cursor.sh")
}
func BuildCursorChecksum(authorization string) string {
const (
checksumTimestampDivisor = 1_000_000
checksumInitialSeed = 165
)
timestamp := time.Now().UnixMilli() / checksumTimestampDivisor
timestampBytes := make([]byte, 6)
timestampBigInt := big.NewInt(timestamp)
for index := 0; index < len(timestampBytes); index++ {
shift := uint((len(timestampBytes) - 1 - index) * 8)
timestampBytes[index] = byte(new(big.Int).Rsh(timestampBigInt, shift).Uint64() & 0xff)
}
seed := checksumInitialSeed
for index := 0; index < len(timestampBytes); index++ {
current := int(timestampBytes[index]^byte(seed)) + (index % 256)
current &= 0xff
timestampBytes[index] = byte(current)
seed = current
}
prefix := strings.TrimRight(base64.StdEncoding.EncodeToString(timestampBytes), "=")
hashBytes := sha256.Sum256([]byte(strings.TrimSpace(authorization)))
hash := fmt.Sprintf("%x", hashBytes)
return prefix + hash[:32]
}
func formatBearerAuthorization(raw string) string {
value := strings.TrimSpace(raw)
if value == "" {
return ""
}
if strings.HasPrefix(strings.ToLower(value), "bearer ") {
return value
}
return "Bearer " + value
}
func shouldRequestCarryBody(method string) bool {
switch strings.ToUpper(strings.TrimSpace(method)) {
case http.MethodGet, http.MethodHead, http.MethodDelete:
@@ -238,17 +185,6 @@ func marshalJSONBody(payload map[string]any) ([]byte, error) {
return json.Marshal(payload)
}
func handleMockJSON(reqCtx *RequestContext, route *Route) error {
responseBody, err := marshalJSONBody(route.JSONBody)
if err != nil {
return err
}
reqCtx.ResponseWriter.Header().Set("content-type", "application/json")
reqCtx.ResponseWriter.WriteHeader(route.StatusCode)
_, _ = reqCtx.ResponseWriter.Write(responseBody)
return nil
}
func handleMockProto(reqCtx *RequestContext, route *Route) error {
payload := map[string]any{}
if route.MockPayloadBuilder != nil {
@@ -270,91 +206,6 @@ func handleMockProto(reqCtx *RequestContext, route *Route) error {
return nil
}
func handleMockOAuth(reqCtx *RequestContext, route *Route) error {
payload := struct {
RefreshToken string `json:"refresh_token"`
}{}
_ = json.Unmarshal(reqCtx.RequestBody, &payload)
responseBody, err := marshalJSONBody(map[string]any{
"access_token": payload.RefreshToken,
"id_token": payload.RefreshToken,
"shouldLogout": false,
})
if err != nil {
return err
}
reqCtx.ResponseWriter.Header().Set("content-type", "application/json")
reqCtx.ResponseWriter.WriteHeader(http.StatusOK)
_, _ = reqCtx.ResponseWriter.Write(responseBody)
return nil
}
func handleMockAuthFullStripeProfile(reqCtx *RequestContext, route *Route) error {
_ = route
responseBody, err := marshalJSONBody(map[string]any{
"membershipType": localUltraMembershipType,
"subscriptionStatus": localUltraSubscriptionStatus,
"lastPaymentFailed": false,
"pendingCancellationDate": "",
"daysRemainingOnTrial": 0,
"paymentId": localUltraPaymentID,
})
if err != nil {
return err
}
reqCtx.ResponseWriter.Header().Set("content-type", "application/json")
reqCtx.ResponseWriter.WriteHeader(http.StatusOK)
_, _ = reqCtx.ResponseWriter.Write(responseBody)
return nil
}
func handleMockAuthStripeProfile(reqCtx *RequestContext, route *Route) error {
_ = route
responseBody, err := json.Marshal(localUltraPaymentID)
if err != nil {
return err
}
reqCtx.ResponseWriter.Header().Set("content-type", "application/json")
reqCtx.ResponseWriter.WriteHeader(http.StatusOK)
_, _ = reqCtx.ResponseWriter.Write(responseBody)
return nil
}
func handleMockAuthPoll(reqCtx *RequestContext, route *Route) error {
_ = route
responseBody, err := marshalJSONBody(map[string]any{
"accessToken": legacyruntime.InjectAuthToken,
"refreshToken": legacyruntime.InjectAuthToken,
"authId": "local_auth",
})
if err != nil {
return err
}
reqCtx.ResponseWriter.Header().Set("content-type", "application/json")
reqCtx.ResponseWriter.WriteHeader(http.StatusOK)
_, _ = reqCtx.ResponseWriter.Write(responseBody)
return nil
}
func handleMockAuthEmail(reqCtx *RequestContext, route *Route) error {
_ = route
responseBody := encodeAuthGetEmailResponse(legacyruntime.InjectAccountEmail)
reqCtx.ResponseWriter.Header().Set("content-type", "application/proto")
reqCtx.ResponseWriter.Header().Set("content-length", strconv.Itoa(len(responseBody)))
reqCtx.ResponseWriter.WriteHeader(http.StatusOK)
_, _ = reqCtx.ResponseWriter.Write(responseBody)
return nil
}
func encodeAuthGetEmailResponse(email string) []byte {
output := make([]byte, 0, len(email)+8)
output = append(output, 0x0a)
output = appendProtoVarint(output, uint64(len(email)))
output = append(output, []byte(email)...)
output = append(output, 0x10, 0x03) // GetEmailResponse.SignUpType.SIGN_UP_TYPE_GOOGLE
return output
}
func appendProtoVarint(output []byte, value uint64) []byte {
for value >= 0x80 {
output = append(output, byte(value)|0x80)
@@ -431,8 +282,6 @@ func newProtoMessage(typeName string) (proto.Message, error) {
return &aiserverv1.GetTeamAdminSettingsResponse{}, nil
case "aiserver.v1.GetTeamReposResponse":
return &aiserverv1.GetTeamReposResponse{}, nil
case "aiserver.v1.ListMarketplacesResponse":
return &aiserverv1.ListMarketplacesResponse{}, nil
case "aiserver.v1.GetUsableModelsResponse":
return &agentv1.GetUsableModelsResponse{}, nil
case "aiserver.v1.GetDefaultModelForCliResponse":
@@ -441,18 +290,12 @@ func newProtoMessage(typeName string) (proto.Message, error) {
return &aiserverv1.GetDefaultModelResponse{}, nil
case "aiserver.v1.GetGlobalCommandsResponse":
return &aiserverv1.GetGlobalCommandsResponse{}, nil
case "aiserver.v1.GetEffectiveUserPluginsResponse":
return &aiserverv1.GetEffectiveUserPluginsResponse{}, nil
case "aiserver.v1.RegisterMarketplaceAndPluginsResponse":
return &aiserverv1.RegisterMarketplaceAndPluginsResponse{}, nil
case "aiserver.v1.GetCliDownloadUrlResponse":
return &aiserverv1.GetCliDownloadUrlResponse{}, nil
case "aiserver.v1.SubmitLogsResponse":
return &aiserverv1.SubmitLogsResponse{}, nil
case "aiserver.v1.TrackEventsResponse":
return &aiserverv1.TrackEventsResponse{}, nil
case "aiserver.v1.IsConnectedResponse":
return &aiserverv1.IsConnectedResponse{}, nil
default:
return nil, fmt.Errorf("unsupported proto message type %q", typeName)
}
@@ -0,0 +1,147 @@
package upstream
import (
"io"
"net/http"
"net/http/httptest"
"net/url"
"strings"
"testing"
"cursor/internal/backend/server"
)
type fallbackHTTPClientFunc func(*http.Request) (*http.Response, error)
func (fn fallbackHTTPClientFunc) Do(request *http.Request) (*http.Response, error) {
return fn(request)
}
func TestFallbackForwardActionUsesOriginalMITMUpstreamURL(t *testing.T) {
originalURL := "https://api3.cursor.sh/aiserver.v1.UnknownService/Call?mode=exact"
parsedURL, err := url.Parse(originalURL)
if err != nil {
t.Fatalf("parse original URL: %v", err)
}
client := fallbackHTTPClientFunc(func(request *http.Request) (*http.Response, error) {
if got := request.URL.String(); got != originalURL {
t.Fatalf("upstream URL: got %q, want %q", got, originalURL)
}
if got := request.Method; got != http.MethodPost {
t.Fatalf("method: got %q, want POST", got)
}
body, readErr := io.ReadAll(request.Body)
if readErr != nil {
t.Fatalf("read request body: %v", readErr)
}
if got := string(body); got != "request-body" {
t.Fatalf("body: got %q", got)
}
if got := request.Header.Get("X-Test-Header"); got != "preserved" {
t.Fatalf("custom header: got %q", got)
}
if got := request.Header.Get(server.HeaderServerUpstreamURL); got != "" {
t.Fatalf("internal upstream header leaked: %q", got)
}
return &http.Response{
StatusCode: http.StatusAccepted,
Status: "202 Accepted",
Header: http.Header{"X-Upstream-Response": []string{"preserved"}},
Body: io.NopCloser(strings.NewReader("upstream-body")),
}, nil
})
request := httptest.NewRequest(http.MethodPost, "http://localhost:8000/ignored", strings.NewReader("request-body"))
request.Header.Set("X-Test-Header", "preserved")
request.Header.Set(server.HeaderServerUpstreamURL, originalURL)
recorder := httptest.NewRecorder()
ctx := &server.Context{Writer: recorder, Request: request, UpstreamURL: parsedURL}
action := FallbackForwardAction(Dependencies{HTTPClient: client}, CompatRouteConfig{Name: "fallback"}, DefaultCursorUpstreamBaseURL)
if err := action(ctx); err != nil {
t.Fatalf("forward fallback request: %v", err)
}
if got := recorder.Code; got != http.StatusAccepted {
t.Fatalf("response status: got %d, want %d", got, http.StatusAccepted)
}
if got := recorder.Header().Get("X-Upstream-Response"); got != "preserved" {
t.Fatalf("response header: got %q", got)
}
if got := recorder.Body.String(); got != "upstream-body" {
t.Fatalf("response body: got %q", got)
}
}
func TestFallbackForwardActionUsesDefaultUpstreamForNativeRequest(t *testing.T) {
const defaultBaseURL = "https://fallback.example:8443"
wantURL := defaultBaseURL + "/aiserver.v1.UnknownService/Call?mode=native"
client := fallbackHTTPClientFunc(func(request *http.Request) (*http.Response, error) {
if got := request.URL.String(); got != wantURL {
t.Fatalf("upstream URL: got %q, want %q", got, wantURL)
}
return &http.Response{
StatusCode: http.StatusNoContent,
Status: "204 No Content",
Header: make(http.Header),
Body: io.NopCloser(strings.NewReader("")),
}, nil
})
request := httptest.NewRequest(http.MethodGet, "http://localhost:8000/aiserver.v1.UnknownService/Call?mode=native", nil)
recorder := httptest.NewRecorder()
ctx := &server.Context{Writer: recorder, Request: request}
action := FallbackForwardAction(Dependencies{HTTPClient: client}, CompatRouteConfig{Name: "fallback"}, defaultBaseURL)
if err := action(ctx); err != nil {
t.Fatalf("forward fallback request: %v", err)
}
if got := recorder.Code; got != http.StatusNoContent {
t.Fatalf("response status: got %d, want %d", got, http.StatusNoContent)
}
}
func TestFallbackForwardActionPreservesAuthorization(t *testing.T) {
const (
originalURL = "https://api2.cursor.sh/aiserver.v1.AuthService/GetEmail"
officialAuthorization = "Bearer official-access-token"
officialChecksum = "official-checksum"
)
parsedURL, err := url.Parse(originalURL)
if err != nil {
t.Fatalf("parse original URL: %v", err)
}
client := fallbackHTTPClientFunc(func(request *http.Request) (*http.Response, error) {
if got := request.Header.Get("Authorization"); got != officialAuthorization {
t.Fatalf("authorization: got %q, want %q", got, officialAuthorization)
}
if got := request.Header.Get("x-cursor-checksum"); got != officialChecksum {
t.Fatalf("checksum: got %q, want %q", got, officialChecksum)
}
return &http.Response{
StatusCode: http.StatusOK,
Status: "200 OK",
Header: make(http.Header),
Body: io.NopCloser(strings.NewReader("upstream-account")),
}, nil
})
request := httptest.NewRequest(http.MethodPost, "http://localhost:8000/aiserver.v1.AuthService/GetEmail", nil)
request.Header.Set("Authorization", officialAuthorization)
request.Header.Set("x-cursor-checksum", officialChecksum)
recorder := httptest.NewRecorder()
ctx := &server.Context{Writer: recorder, Request: request, UpstreamURL: parsedURL}
action := FallbackForwardAction(
Dependencies{HTTPClient: client},
CompatRouteConfig{Name: "fallback"},
DefaultCursorUpstreamBaseURL,
)
if err := action(ctx); err != nil {
t.Fatalf("forward authenticated fallback request: %v", err)
}
if got := recorder.Body.String(); got != "upstream-account" {
t.Fatalf("response body: got %q, want upstream-account", got)
}
}
+45 -60
View File
@@ -24,9 +24,7 @@ const (
// файловых инструментов падают с "[unimplemented] HTTP 404".
localPathEncryptionKey = "6f6e63652d6c6f63616c2d706174682d656e6372797074696f6e2d6b6579"
localUltraMembershipType = "ultra"
localUltraPaymentID = "local_ultra"
localUltraSubscriptionStatus = "active"
localUltraPlanIncludedCents = 20000
localUltraDashboardUserID = 1
localUltraBillingCycleDuration = 30 * 24 * time.Hour
@@ -431,7 +429,8 @@ func buildServerTimePayload(*RequestContext) (map[string]any, error) {
func buildServerConfigPayload(*RequestContext) (map[string]any, error) {
return map[string]any{
"configVersion": "local_cli_sandbox_defaults_disabled_v2",
"configVersion": "local_cli_sandbox_defaults_disabled_v2",
"isDevDoNotUseForSecretThingsBecauseCanBeSpoofedByUsers": true,
"http2Config": "HTTP2_CONFIG_FORCE_ALL_DISABLED",
"cliSandboxDefaultEnabled": true,
"indexingConfig": map[string]any{
@@ -547,26 +546,29 @@ func buildFirstWindowStatsigDecisionPayload(*RequestContext) (map[string]any, er
}, nil
}
func buildDashboardCurrentPeriodUsagePayload(*RequestContext) (map[string]any, error) {
func buildDashboardCurrentPeriodUsagePayload(reqCtx *RequestContext) (map[string]any, error) {
plan := localDevPlanFromRequest(reqCtx)
planName, includedSpend := localDevPlanDetails(plan)
billingCycleStart := time.Now().Add(-localUltraBillingCycleDuration).UnixMilli()
billingCycleEnd := time.Now().Add(10 * 365 * 24 * time.Hour).UnixMilli()
displayMessage := planName + " active"
return map[string]any{
"autoModelSelectedDisplayMessage": "Ultra plan active",
"autoModelSelectedDisplayMessage": displayMessage,
"billingCycleEnd": billingCycleEnd,
"billingCycleStart": billingCycleStart,
"displayMessage": "Ultra plan active",
"displayMessage": displayMessage,
"displayThreshold": 99999999,
"enabled": true,
"namedModelSelectedDisplayMessage": "Ultra plan active",
"namedModelSelectedDisplayMessage": displayMessage,
"planUsage": map[string]any{
"apiPercentUsed": 0,
"apiSpend": 0,
"autoPercentUsed": 0,
"autoSpend": 0,
"bonusTooltip": "Ultra local account mock is active.",
"includedSpend": localUltraPlanIncludedCents,
"limit": localUltraPlanIncludedCents,
"remaining": localUltraPlanIncludedCents,
"bonusTooltip": "Local account mock is active.",
"includedSpend": includedSpend,
"limit": includedSpend,
"remaining": includedSpend,
"remainingBonus": false,
"totalPercentUsed": 0,
"totalSpend": 0,
@@ -577,62 +579,45 @@ func buildDashboardCurrentPeriodUsagePayload(*RequestContext) (map[string]any, e
}, nil
}
func buildDashboardTeamsPayload(*RequestContext) (map[string]any, error) {
func buildDashboardTeamsPayload(reqCtx *RequestContext) (map[string]any, error) {
if claims, ok := localDevClaimsFromRequest(reqCtx); ok && claims.Plan == "enterprise" {
return map[string]any{
"teams": []map[string]any{{
"name": "Local Enterprise",
"id": 1,
"seats": 1,
"hasBilling": true,
"subscriptionStatus": localDevSubscriptionActive,
"verified": true,
"isEnterprise": true,
"membershipType": "enterprise",
}},
}, nil
}
return map[string]any{
"teams": []map[string]any{},
}, nil
}
func buildDashboardManagedSkillsPayload(*RequestContext) (map[string]any, error) {
return map[string]any{
"skills": []map[string]any{},
}, nil
}
func buildDashboardGetMePayload(reqCtx *RequestContext) (map[string]any, error) {
authID := ""
if reqCtx != nil {
authID = authIDFromBearer(reqCtx.Headers.Get("authorization"))
func buildDashboardPlanInfoPayload(reqCtx *RequestContext) (map[string]any, error) {
plan := localDevPlanFromRequest(reqCtx)
planName, includedAmountCents := localDevPlanDetails(plan)
price := "$200/mo"
switch plan {
case "free":
price = "$0/mo"
case "pro":
price = "$20/mo"
case "pro_plus":
price = "$60/mo"
case "enterprise":
price = "Custom"
}
if authID == "" {
authID = authIDFromJWT(legacyruntime.InjectAuthToken)
}
if authID == "" {
authID = localUltraPaymentID
}
return map[string]any{
"authId": authID,
"userId": localUltraDashboardUserID,
"email": legacyruntime.InjectAccountEmail,
"firstName": "Cursor",
"lastName": "Local",
"createdAt": time.Now().UTC().Format(time.RFC3339),
"isEnterpriseUser": false,
"teamName": "",
"emailDomainType": "personal",
"country": "US",
"profilePictureUrl": "",
}, nil
}
func buildDashboardUserPrivacyModePayload(*RequestContext) (map[string]any, error) {
return map[string]any{
"privacyMode": "PRIVACY_MODE_NO_STORAGE",
"hoursRemainingInGracePeriod": 0,
"isEnforcedByTeam": false,
"isNotMigratedToServerSourceOfTruth": false,
"partnerDataShare": false,
"hasAcknowledgedGracePeriodDisclaimer": true,
}, nil
}
func buildDashboardPlanInfoPayload(*RequestContext) (map[string]any, error) {
return map[string]any{
"planInfo": map[string]any{
"planName": "Ultra Plan",
"includedAmountCents": localUltraPlanIncludedCents,
"price": "$200/mo",
"planName": planName,
"includedAmountCents": includedAmountCents,
"price": price,
"billingCycleEnd": time.Now().Add(10 * 365 * 24 * time.Hour).UnixMilli(),
},
}, nil
@@ -835,7 +820,7 @@ func defaultThinkingEffortForAdapter(adapter legacyruntime.ModelAdapterConfig) s
if strings.EqualFold(strings.TrimSpace(adapter.Type), "anthropic") {
return normalizeAvailableModelThinkingEffort(adapter.AnthropicThinkingEffort, true, "xhigh")
}
return normalizeAvailableModelThinkingEffort(adapter.ReasoningEffort, true, "disabled")
return normalizeAvailableModelThinkingEffort(adapter.ReasoningEffort, true, "medium")
}
func normalizeAvailableModelThinkingEffort(raw string, allowMax bool, fallback string) string {
+21 -34
View File
@@ -6,6 +6,7 @@ import (
"testing"
"cursor/gen/agentv1"
"cursor/gen/aiserverv1"
legacyruntime "cursor/internal/runtime"
"google.golang.org/protobuf/proto"
@@ -28,40 +29,6 @@ func TestBuildCLIModelDetailsPreservesChannelMetadata(t *testing.T) {
}
}
func TestDefaultThinkingEffortForOpenAIAdapterUsesDisabledWhenUnset(t *testing.T) {
adapter := legacyruntime.ModelAdapterConfig{Type: "openai", ReasoningEffort: ""}
if got := defaultThinkingEffortForAdapter(adapter); got != "disabled" {
t.Fatalf("default thinking effort = %q, want disabled", got)
}
}
func TestBuildAvailableModelEntriesUsesDisabledVariantWhenReasoningEffortUnset(t *testing.T) {
entries := buildAvailableModelEntries([]legacyruntime.ModelAdapterConfig{{
ID: "channel-a",
DisplayName: "Model A",
ModelID: "model-a",
Type: "openai",
}})
if len(entries) != 1 {
t.Fatalf("entry count = %d, want 1", len(entries))
}
variants, ok := entries[0]["variants"].([]map[string]any)
if !ok {
t.Fatalf("variants type = %T, want []map[string]any", entries[0]["variants"])
}
if len(variants) == 0 {
t.Fatal("variants should not be empty")
}
if got := variants[0]["variantStringRepresentation"]; got != "channel-a:disabled" {
t.Fatalf("first variant representation = %#v, want channel-a:disabled", got)
}
if got := variants[0]["isDefaultNonMaxConfig"]; got != true {
t.Fatalf("disabled variant default flag = %#v, want true", got)
}
}
func TestEncodeCLIModelsUsesAgentModelDetailsWireFormat(t *testing.T) {
payload := map[string]any{"models": buildCLIModelDetails([]legacyruntime.ModelAdapterConfig{{ID: "channel-a", DisplayName: "Model A", APIKey: "provider-secret", BaseURL: "https://provider.example/v1"}})}
encoded, err := encodeMockProto("aiserver.v1.GetUsableModelsResponse", payload)
@@ -88,6 +55,26 @@ func TestEncodeCLIModelsUsesAgentModelDetailsWireFormat(t *testing.T) {
}
}
func TestBuildServerConfigEnablesDevUserBackendCommands(t *testing.T) {
payload, err := buildServerConfigPayload(nil)
if err != nil {
t.Fatalf("build server config: %v", err)
}
encoded, err := encodeMockProto("aiserver.v1.GetServerConfigResponse", payload)
if err != nil {
t.Fatalf("encode server config: %v", err)
}
response := &aiserverv1.GetServerConfigResponse{}
if err := proto.Unmarshal(encoded, response); err != nil {
t.Fatalf("decode server config: %v", err)
}
if !response.GetIsDevDoNotUseForSecretThingsBecauseCanBeSpoofedByUsers() {
t.Fatal("expected server config to enable dev-user backend commands")
}
}
func TestBuildBootstrapStatsigConfigJSONDisablesAlwaysLocalDecompositionGate(t *testing.T) {
payload, err := buildBootstrapStatsigConfigJSON(12345, "test-auth-id")
if err != nil {
@@ -20,13 +20,6 @@ type SystemSettingService interface {
ResolveModelAdapters(context.Context) ([]legacyruntime.ModelAdapterConfig, error)
}
// AuthorizationProvider supplies the independent Cursor account used only by
// official control-plane requests such as Plugins, Skills, and MCP registry.
type AuthorizationProvider interface {
Authorization(context.Context) (string, error)
SignedIn() bool
}
type HTTPClient interface {
Do(req *http.Request) (*http.Response, error)
}
@@ -91,7 +84,6 @@ type Route struct {
Matcher Matcher
ConsoleLog bool
StatusCode int
JSONBody map[string]any
MockProtoType string
MockPayloadBuilder func(*RequestContext) (map[string]any, error)
Handler RouteHandler
-28
View File
@@ -30,9 +30,6 @@ type ModelAdapterModelsRequest = client.ModelAdapterModelsRequest
// ModelAdapterModelsResult 定义模型列表查询结果。
type ModelAdapterModelsResult = client.ModelAdapterModelsResult
// CursorAccountStatus 是可安全展示给桌面前端的独立 Cursor 账号状态。
type CursorAccountStatus = client.CursorAccountStatus
// LicenseActionRequest 定义了当前模块中的 LicenseActionRequest 类型。
type LicenseActionRequest = client.LicenseActionRequest
@@ -100,31 +97,6 @@ func (s *ProxyService) SaveUserConfig(cfg UserConfig) error {
return s.core.SaveUserConfig(cfg)
}
// ExportUserConfig 将当前完整配置导出为 YAML 文件。
func (s *ProxyService) ExportUserConfig(path string) (string, error) {
return s.core.ExportUserConfig(path)
}
// ImportUserConfig 从 YAML 文件校验并替换当前完整配置。
func (s *ProxyService) ImportUserConfig(path string) (UserConfig, error) {
return s.core.ImportUserConfig(path)
}
// GetCursorAccountStatus 返回 cursor-byok 独立 Cursor 账号的脱敏状态。
func (s *ProxyService) GetCursorAccountStatus() CursorAccountStatus {
return s.core.GetCursorAccountStatus()
}
// StartCursorAccountLogin 打开官方浏览器登录并异步等待结果。
func (s *ProxyService) StartCursorAccountLogin() (CursorAccountStatus, error) {
return s.core.StartCursorAccountLogin()
}
// DisconnectCursorAccount 只断开 cursor-byok 自己的账号。
func (s *ProxyService) DisconnectCursorAccount() (CursorAccountStatus, error) {
return s.core.DisconnectCursorAccount()
}
// TestModelAdapter 用于处理与 TestModelAdapter 相关的逻辑。
func (s *ProxyService) TestModelAdapter(adapter ModelAdapterConfig) (ModelAdapterTestResult, error) {
return s.core.TestModelAdapter(adapter)
+44 -42
View File
@@ -1,7 +1,6 @@
package certs
import (
"bytes"
"crypto"
"crypto/ecdsa"
"crypto/ed25519"
@@ -10,24 +9,33 @@ import (
"crypto/tls"
"crypto/x509"
"crypto/x509/pkix"
_ "embed"
"encoding/pem"
"errors"
"fmt"
"math/big"
"net"
"os"
"strings"
"sync"
"time"
)
// embeddedCACertPEM 表示当前模块中的 embeddedCACertPEM 状态值。
//
//go:embed ca.crt
var embeddedCACertPEM []byte
// embeddedCAKeyPEM 表示当前模块中的 embeddedCAKeyPEM 状态值。
//
//go:embed ca.key
var embeddedCAKeyPEM []byte
// Manager 定义了当前模块中的 Manager 类型。
type Manager struct {
// caCert 表示当前声明中的 caCert。
caCert *x509.Certificate
// caKey 表示当前声明中的 caKey。
caKey crypto.PrivateKey
// caCertPEM 保存可注入宿主信任存储的 CA 证书,不包含私钥。
caCertPEM []byte
// mu 表示当前声明中的 mu。
mu sync.Mutex
@@ -44,26 +52,28 @@ func NewManager(caCertPath, caKeyPath string) (*Manager, error) {
return NewManagerFromPEM(certPEM, keyPEM)
}
// NewEmbeddedManager 用于处理与 NewEmbeddedManager 相关的逻辑。
func NewEmbeddedManager() (*Manager, error) {
return NewManagerFromPEM(embeddedCACertPEM, embeddedCAKeyPEM)
}
// EmbeddedCACertPEM 用于处理与 EmbeddedCACertPEM 相关的逻辑。
func EmbeddedCACertPEM() []byte {
return cloneBytes(embeddedCACertPEM)
}
// EmbeddedCAKeyPEM 用于处理与 EmbeddedCAKeyPEM 相关的逻辑。
func EmbeddedCAKeyPEM() []byte {
return cloneBytes(embeddedCAKeyPEM)
}
// NewManagerFromPEM 用于处理与 NewManagerFromPEM 相关的逻辑。
func NewManagerFromPEM(caCertPEM, caKeyPEM []byte) (*Manager, error) {
caCert, caKey, err := loadCAFromPEM(caCertPEM, caKeyPEM)
if err != nil {
return nil, err
}
return &Manager{
caCert: caCert,
caKey: caKey,
caCertPEM: cloneBytes(caCertPEM),
cache: make(map[string]*tls.Certificate),
}, nil
}
// CACertPEM 返回当前 Manager 使用的 CA 证书。返回值不包含私钥。
func (m *Manager) CACertPEM() []byte {
if m == nil {
return nil
}
return cloneBytes(m.caCertPEM)
return &Manager{caCert: caCert, caKey: caKey, cache: make(map[string]*tls.Certificate)}, nil
}
// CATLSCertificate 用于处理与 CATLSCertificate 相关的逻辑。
@@ -175,6 +185,19 @@ func marshalPrivateKeyPEM(key any) ([]byte, error) {
}
}
// loadCAPEMFromFiles 用于处理与 loadCAPEMFromFiles 相关的逻辑。
func loadCAPEMFromFiles(certPath, keyPath string) ([]byte, []byte, error) {
certPEM, err := os.ReadFile(certPath)
if err != nil {
return nil, nil, err
}
keyPEM, err := os.ReadFile(keyPath)
if err != nil {
return nil, nil, err
}
return certPEM, keyPEM, nil
}
// loadCAFromPEM 用于处理与 loadCAFromPEM 相关的逻辑。
func loadCAFromPEM(certPEM, keyPEM []byte) (*x509.Certificate, crypto.PrivateKey, error) {
certBlock, _ := pem.Decode(certPEM)
@@ -191,49 +214,28 @@ func loadCAFromPEM(certPEM, keyPEM []byte) (*x509.Certificate, crypto.PrivateKey
return nil, nil, errors.New("invalid CA key PEM")
}
var caKey crypto.PrivateKey
switch keyBlock.Type {
case "RSA PRIVATE KEY":
key, err := x509.ParsePKCS1PrivateKey(keyBlock.Bytes)
if err != nil {
return nil, nil, err
}
caKey = key
return caCert, key, nil
case "EC PRIVATE KEY":
key, err := x509.ParseECPrivateKey(keyBlock.Bytes)
if err != nil {
return nil, nil, err
}
caKey = key
return caCert, key, nil
case "PRIVATE KEY":
key, err := x509.ParsePKCS8PrivateKey(keyBlock.Bytes)
if err != nil {
return nil, nil, err
}
caKey = key
return caCert, key, nil
default:
return nil, nil, errors.New("unsupported CA key format")
}
if !caCert.IsCA || !caCert.BasicConstraintsValid || caCert.KeyUsage&x509.KeyUsageCertSign == 0 {
return nil, nil, errors.New("certificate is not a valid signing CA")
}
signer, ok := caKey.(crypto.Signer)
if !ok {
return nil, nil, errors.New("CA private key cannot sign certificates")
}
certPublicKey, err := x509.MarshalPKIXPublicKey(caCert.PublicKey)
if err != nil {
return nil, nil, fmt.Errorf("marshal CA certificate public key: %w", err)
}
privatePublicKey, err := x509.MarshalPKIXPublicKey(signer.Public())
if err != nil {
return nil, nil, fmt.Errorf("marshal CA private key public key: %w", err)
}
if !bytes.Equal(certPublicKey, privatePublicKey) {
return nil, nil, errors.New("CA certificate and private key do not match")
}
return caCert, caKey, nil
}
// normalizeHost 用于处理与 normalizeHost 相关的逻辑。
+27
View File
@@ -0,0 +1,27 @@
-----BEGIN RSA PRIVATE KEY-----
MIIEpAIBAAKCAQEAyh3jND/aFusuRjGTmhQtX2hkF1qroNjEEKCxWPlfprvdl8Tx
upxNv1TQcm+K9KsS7OxnKYP4Qtv068XLbaCCGuoA/xpor6enrT85KulBq0j8z/g1
y0VWxjz3xN9F9ND13h9yxDZCn76egbRkhFxpXow67jLsIlrqWDSltERlTKh2cJ1g
hTRuQNSr7jtKQgFAR3aQ6dTzOP4fCOtLYn63jL4+YcxdGoK66tx8eFHq8oBLYvsc
LNjjkaAHilZFwA4Jr7zHBofTRBg04eZum2UTaRqmSyT65ifXm4vRdQa4k1FS1gnq
hPMhIjST8b6RkxvbLNmFpkOfI9eCMRpFG7y+NQIDAQABAoIBAATU9ZVOcHmLSkop
zcBJerM09O2dAIziGb/XA55fqdJ728aQ0gGW0oIANlKCCaWjQFrTJP04VzNL/F01
l5EpnOqlTPxMRpPqc2cAI677sBL29fpH0gtnvzUSiI7Xkp3RcAtNH6qCrJGSlkn+
BMgoSGmW+yKuK3h/yWnt6kc2umA8fN+bHKhS3pI56PMW8qVnny9n92RaCA7Uf/4j
XDewIreiH5jRqRwrPbOpjDFmv+W18LWZQiTwwxfY6sRZiZpsfsHidzfFGUFZXMlq
2P3FCqoF4oMM1rRgBlHhDR7JHmSkFpZG639HJTXLllpyDbj3H86nsjIj9WZYKn+h
B9k9bcECgYEA1HDzRCqoRZh6cL46KiYjv+LwmEKMa3nPb4ljywgLbRkzTMyVs0MK
fDsDoTLFY9PBhveypU5gjTbQqtyBsDtbU+dH9Eks1FdL9P8bJf0ZOJFRiEB7uB4a
z3V9tcXHwH4l2bCWbWThQGFwRudDAoY0EH89oSA/WHayjOKe1Wi4ovUCgYEA848B
cYi+Qbkk+fOv9gSJn8KS1LH/jE28S/e7E4YkTfYUuu+7wr8bdRKUNpPIsLX0Fo9R
KpJX0Oyjjady9n/8ARZRmD8Upl+F7Sl7Ro6F8+nfqQUrxbVDiIL3b+aZF+cDfFrB
/xL5kyZZqTtFfP360tfYnlS6Sssd4E2Jsj0fJkECgYEA0gG+WbqZkgMDtwQ114jQ
elZLZRkUWwKVjzsQDZssQHNTBS6RJh619M0Z73aTLvYcL+IZFdT/GVoAuYc2JRLo
W28c8F6OFHMfwVeWbN1g20y8fqbQJtiLxF3vIYwcxStvG123tvisu8oXBeCDm7Ez
MsO2FtwcAsWECEXWojzdmSkCgYEA18GvPawtHnuszd+Z2Q5b/DKZb+Hex6N1Ura5
+qmyL333D0Kfyf0RjbxPn6l690+4UuPSuyu4r1Nx72KO7N6jlzL2RTBcUqX8NgOx
OOe4skJT5561EAdrM9sQ5wgYRpxW8ipUAGoGvNwUQV5ISFmVgIHFWz0jam5UoQcP
G94ZYgECgYBN4PB8BuAwMRqhTCLL5RJKtcdC/Ls9xdWrmswtQQ4OSBUTLGdeMhEN
E23F0d+NdtiTWDYRDJ7z6KAF8CcOwWMt2sNbrxgNRuDzyfLqbuVLOM8iY3xA50Nl
clZWujKILWp47+gA6/AqJLv2lA2LK8gM8zYzzm/nCuLCBLYH/hOiCA==
-----END RSA PRIVATE KEY-----
-191
View File
@@ -1,191 +0,0 @@
package certs
import (
"crypto/rand"
"crypto/rsa"
"crypto/sha256"
"crypto/x509"
"crypto/x509/pkix"
"encoding/hex"
"encoding/pem"
"errors"
"fmt"
"math/big"
"os"
"path/filepath"
"strings"
"time"
)
const legacySharedCASHA256 = "836E6BB84F6C3E63316DBB4EC257223AF09F7490E7AAE09030B8515ED61EE9FF"
// LoadOrCreateManager loads the installation-specific CA, generating it on
// first run. The private key is persisted only in the supplied local path.
func LoadOrCreateManager(certPath, keyPath string) (*Manager, []byte, error) {
certPEM, certErr := os.ReadFile(certPath)
keyPEM, keyErr := os.ReadFile(keyPath)
if certErr == nil && isLegacySharedCA(certPEM) {
return generateAndPersistManager(certPath, keyPath)
}
if certErr == nil && keyErr == nil {
manager, err := NewManagerFromPEM(certPEM, keyPEM)
if err != nil {
return nil, nil, fmt.Errorf("load installation CA: %w", err)
}
if err := os.Chmod(keyPath, 0o600); err != nil {
return nil, nil, fmt.Errorf("restrict installation CA private key permissions: %w", err)
}
return manager, manager.CACertPEM(), nil
}
if errors.Is(certErr, os.ErrNotExist) && errors.Is(keyErr, os.ErrNotExist) {
return generateAndPersistManager(certPath, keyPath)
}
// A key without a certificate cannot have been installed as a trusted root.
// This is safe to recover if the first-run write was interrupted.
if errors.Is(certErr, os.ErrNotExist) && keyErr == nil {
return generateAndPersistManager(certPath, keyPath)
}
if certErr != nil && !errors.Is(certErr, os.ErrNotExist) {
return nil, nil, fmt.Errorf("read installation CA certificate: %w", certErr)
}
if keyErr != nil && !errors.Is(keyErr, os.ErrNotExist) {
return nil, nil, fmt.Errorf("read installation CA private key: %w", keyErr)
}
return nil, nil, errors.New("installation CA is incomplete; both certificate and private key are required")
}
// NewGeneratedManager creates an in-memory CA suitable for short-lived tools.
func NewGeneratedManager() (*Manager, []byte, error) {
certPEM, keyPEM, err := generateCA()
if err != nil {
return nil, nil, err
}
manager, err := NewManagerFromPEM(certPEM, keyPEM)
if err != nil {
return nil, nil, err
}
return manager, manager.CACertPEM(), nil
}
func generateAndPersistManager(certPath, keyPath string) (*Manager, []byte, error) {
if filepath.Dir(certPath) != filepath.Dir(keyPath) {
return nil, nil, errors.New("installation CA certificate and key must share a directory")
}
certPEM, keyPEM, err := generateCA()
if err != nil {
return nil, nil, err
}
if err := os.MkdirAll(filepath.Dir(certPath), 0o700); err != nil {
return nil, nil, fmt.Errorf("create installation CA directory: %w", err)
}
// Write the private key first so a crash cannot leave a new certificate
// without the signing key needed by the proxy.
if err := writeLocalCAFile(keyPath, keyPEM, 0o600); err != nil {
return nil, nil, fmt.Errorf("persist installation CA private key: %w", err)
}
if err := writeLocalCAFile(certPath, certPEM, 0o644); err != nil {
return nil, nil, fmt.Errorf("persist installation CA certificate: %w", err)
}
manager, err := NewManagerFromPEM(certPEM, keyPEM)
if err != nil {
return nil, nil, err
}
return manager, manager.CACertPEM(), nil
}
func generateCA() ([]byte, []byte, error) {
privateKey, err := rsa.GenerateKey(rand.Reader, 3072)
if err != nil {
return nil, nil, fmt.Errorf("generate installation CA private key: %w", err)
}
serial, err := rand.Int(rand.Reader, new(big.Int).Lsh(big.NewInt(1), 128))
if err != nil {
return nil, nil, fmt.Errorf("generate installation CA serial: %w", err)
}
publicKeyDER, err := x509.MarshalPKIXPublicKey(&privateKey.PublicKey)
if err != nil {
return nil, nil, fmt.Errorf("marshal installation CA public key: %w", err)
}
subjectKeyID := sha256.Sum256(publicKeyDER)
now := time.Now()
template := &x509.Certificate{
SerialNumber: serial,
Subject: pkix.Name{
CommonName: "Cursor BYOK Local CA",
Organization: []string{"Cursor BYOK"},
},
NotBefore: now.Add(-5 * time.Minute),
NotAfter: now.AddDate(10, 0, 0),
KeyUsage: x509.KeyUsageDigitalSignature | x509.KeyUsageCertSign | x509.KeyUsageCRLSign,
BasicConstraintsValid: true,
IsCA: true,
MaxPathLen: 0,
MaxPathLenZero: true,
SubjectKeyId: append([]byte(nil), subjectKeyID[:20]...),
}
der, err := x509.CreateCertificate(rand.Reader, template, template, &privateKey.PublicKey, privateKey)
if err != nil {
return nil, nil, fmt.Errorf("create installation CA certificate: %w", err)
}
certPEM := pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: der})
keyPEM := pem.EncodeToMemory(&pem.Block{Type: "RSA PRIVATE KEY", Bytes: x509.MarshalPKCS1PrivateKey(privateKey)})
return certPEM, keyPEM, nil
}
func writeLocalCAFile(path string, data []byte, mode os.FileMode) error {
temp, err := os.CreateTemp(filepath.Dir(path), ".ca-*")
if err != nil {
return err
}
tempPath := temp.Name()
defer os.Remove(tempPath)
if err := temp.Chmod(mode); err != nil {
_ = temp.Close()
return err
}
if _, err := temp.Write(data); err != nil {
_ = temp.Close()
return err
}
if err := temp.Sync(); err != nil {
_ = temp.Close()
return err
}
if err := temp.Close(); err != nil {
return err
}
if err := os.Remove(path); err != nil && !errors.Is(err, os.ErrNotExist) {
return err
}
if err := os.Rename(tempPath, path); err != nil {
return err
}
return os.Chmod(path, mode)
}
func isLegacySharedCA(certPEM []byte) bool {
block, _ := pem.Decode(certPEM)
if block == nil {
return false
}
cert, err := x509.ParseCertificate(block.Bytes)
if err != nil {
return false
}
sum := sha256.Sum256(cert.Raw)
return strings.EqualFold(hex.EncodeToString(sum[:]), legacySharedCASHA256)
}
// loadCAPEMFromFiles reads an explicitly supplied CA pair.
func loadCAPEMFromFiles(certPath, keyPath string) ([]byte, []byte, error) {
certPEM, err := os.ReadFile(certPath)
if err != nil {
return nil, nil, err
}
keyPEM, err := os.ReadFile(keyPath)
if err != nil {
return nil, nil, err
}
return certPEM, keyPEM, nil
}
-125
View File
@@ -1,125 +0,0 @@
package certs
import (
"bytes"
"crypto/x509"
"encoding/pem"
"os"
"path/filepath"
"runtime"
"testing"
)
func TestLoadOrCreateManagerPersistsAndReusesInstallationCA(t *testing.T) {
dir := t.TempDir()
certPath := filepath.Join(dir, "ca.crt")
keyPath := filepath.Join(dir, "ca.key")
manager, certPEM, err := LoadOrCreateManager(certPath, keyPath)
if err != nil {
t.Fatalf("LoadOrCreateManager() error = %v", err)
}
if isLegacySharedCA(certPEM) {
t.Fatal("generated CA reused the legacy shared certificate")
}
keyInfo, err := os.Stat(keyPath)
if err != nil {
t.Fatalf("stat private key: %v", err)
}
if runtime.GOOS != "windows" && keyInfo.Mode().Perm() != 0o600 {
t.Fatalf("private key mode = %o, want 600", keyInfo.Mode().Perm())
}
leaf, err := manager.CertificateForServerName("api2.cursor.sh")
if err != nil {
t.Fatalf("CertificateForServerName() error = %v", err)
}
ca := parseCertificatePEM(t, certPEM)
roots := x509.NewCertPool()
roots.AddCert(ca)
if _, err := leaf.Leaf.Verify(x509.VerifyOptions{DNSName: "api2.cursor.sh", Roots: roots}); err != nil {
t.Fatalf("verify generated leaf: %v", err)
}
_, reusedCertPEM, err := LoadOrCreateManager(certPath, keyPath)
if err != nil {
t.Fatalf("second LoadOrCreateManager() error = %v", err)
}
if !bytes.Equal(certPEM, reusedCertPEM) {
t.Fatal("installation CA changed between loads")
}
}
func TestLoadOrCreateManagerCreatesUniqueCAsPerInstallation(t *testing.T) {
firstDir := t.TempDir()
secondDir := t.TempDir()
_, firstCert, err := LoadOrCreateManager(filepath.Join(firstDir, "ca.crt"), filepath.Join(firstDir, "ca.key"))
if err != nil {
t.Fatalf("create first CA: %v", err)
}
_, secondCert, err := LoadOrCreateManager(filepath.Join(secondDir, "ca.crt"), filepath.Join(secondDir, "ca.key"))
if err != nil {
t.Fatalf("create second CA: %v", err)
}
if bytes.Equal(firstCert, secondCert) {
t.Fatal("separate installations received the same CA certificate")
}
}
func TestLoadOrCreateManagerReplacesLegacySharedCertificate(t *testing.T) {
dir := t.TempDir()
certPath := filepath.Join(dir, "ca.crt")
keyPath := filepath.Join(dir, "ca.key")
legacyCert, err := os.ReadFile(filepath.Join("testdata", "legacy_shared_ca.crt"))
if err != nil {
t.Fatalf("read legacy certificate fixture: %v", err)
}
if !isLegacySharedCA(legacyCert) {
t.Fatal("legacy certificate fixture fingerprint changed")
}
if err := os.WriteFile(certPath, legacyCert, 0o644); err != nil {
t.Fatalf("write legacy certificate: %v", err)
}
_, generatedCert, err := LoadOrCreateManager(certPath, keyPath)
if err != nil {
t.Fatalf("migrate legacy certificate: %v", err)
}
if bytes.Equal(legacyCert, generatedCert) || isLegacySharedCA(generatedCert) {
t.Fatal("legacy shared certificate was not replaced")
}
if _, err := os.Stat(keyPath); err != nil {
t.Fatalf("generated private key missing: %v", err)
}
}
func TestNewManagerFromPEMRejectsMismatchedKey(t *testing.T) {
firstDir := t.TempDir()
secondDir := t.TempDir()
_, _, err := LoadOrCreateManager(filepath.Join(firstDir, "ca.crt"), filepath.Join(firstDir, "ca.key"))
if err != nil {
t.Fatalf("create first CA: %v", err)
}
_, _, err = LoadOrCreateManager(filepath.Join(secondDir, "ca.crt"), filepath.Join(secondDir, "ca.key"))
if err != nil {
t.Fatalf("create second CA: %v", err)
}
certPEM, _ := os.ReadFile(filepath.Join(firstDir, "ca.crt"))
keyPEM, _ := os.ReadFile(filepath.Join(secondDir, "ca.key"))
if _, err := NewManagerFromPEM(certPEM, keyPEM); err == nil {
t.Fatal("NewManagerFromPEM() accepted a mismatched private key")
}
}
func parseCertificatePEM(t *testing.T, certPEM []byte) *x509.Certificate {
t.Helper()
block, _ := pem.Decode(certPEM)
if block == nil {
t.Fatal("certificate PEM is invalid")
}
cert, err := x509.ParseCertificate(block.Bytes)
if err != nil {
t.Fatalf("parse certificate: %v", err)
}
return cert
}
-12
View File
@@ -17,12 +17,6 @@ func (s *ProxyService) LoadUserConfig() (UserConfig, error) {
if s == nil {
return serverconfig.DefaultConfig(), nil
}
s.configMu.Lock()
defer s.configMu.Unlock()
return s.loadUserConfig()
}
func (s *ProxyService) loadUserConfig() (UserConfig, error) {
app := application.Get()
ctx := context.Background()
if app != nil {
@@ -42,12 +36,6 @@ func (s *ProxyService) SaveUserConfig(cfg UserConfig) error {
if s == nil {
return nil
}
s.configMu.Lock()
defer s.configMu.Unlock()
return s.saveUserConfig(cfg)
}
func (s *ProxyService) saveUserConfig(cfg UserConfig) error {
app := application.Get()
ctx := context.Background()
if app != nil {
-195
View File
@@ -1,195 +0,0 @@
package client
import (
"bytes"
"errors"
"fmt"
"io"
"os"
"path/filepath"
"strings"
serverconfig "cursor/internal/backend/server/config"
"gopkg.in/yaml.v3"
)
const maxConfigTransferFileSize = 4 << 20
type importedUserConfigDocument struct {
serverconfig.Config `yaml:",inline"`
LegacyRouting any `yaml:"routing,omitempty"`
}
// ExportUserConfig 将当前完整配置导出为 YAML 文件。
func (s *ProxyService) ExportUserConfig(path string) (string, error) {
if s == nil {
return "", errors.New("配置服务未初始化")
}
targetPath := normalizeConfigExportPath(path)
if targetPath == "" {
return "", errors.New("导出路径不能为空")
}
s.configMu.Lock()
defer s.configMu.Unlock()
cfg, err := s.loadUserConfig()
if err != nil {
return "", fmt.Errorf("读取当前配置失败: %w", err)
}
data, err := yaml.Marshal(cfg)
if err != nil {
return "", fmt.Errorf("序列化导出配置失败: %w", err)
}
if err := writeExportedUserConfig(targetPath, data); err != nil {
return "", err
}
return targetPath, nil
}
// ImportUserConfig 从 YAML 文件校验并替换当前完整配置。
func (s *ProxyService) ImportUserConfig(path string) (UserConfig, error) {
if s == nil {
return serverconfig.Config{}, errors.New("配置服务未初始化")
}
s.lifecycleMu.Lock()
defer s.lifecycleMu.Unlock()
s.configMu.Lock()
defer s.configMu.Unlock()
if err := ensureConfigImportAllowed(s.GetState()); err != nil {
return serverconfig.Config{}, err
}
data, err := readImportedUserConfig(path)
if err != nil {
return serverconfig.Config{}, err
}
cfg, err := decodeImportedUserConfig(data)
if err != nil {
return serverconfig.Config{}, err
}
if err := s.saveUserConfig(cfg); err != nil {
return serverconfig.Config{}, fmt.Errorf("保存导入配置失败: %w", err)
}
persisted, err := s.loadUserConfig()
if err != nil {
return serverconfig.Config{}, fmt.Errorf("重新读取导入配置失败: %w", err)
}
return persisted, nil
}
func normalizeConfigExportPath(path string) string {
trimmed := strings.TrimSpace(path)
if trimmed == "" {
return ""
}
extension := strings.ToLower(filepath.Ext(trimmed))
if extension != ".yaml" && extension != ".yml" {
trimmed += ".yaml"
}
return filepath.Clean(trimmed)
}
func writeExportedUserConfig(path string, data []byte) error {
directory := filepath.Dir(path)
file, err := os.CreateTemp(directory, ".cursor-byok-config-*.tmp")
if err != nil {
return fmt.Errorf("创建导出配置临时文件失败: %w", err)
}
temporaryPath := file.Name()
defer os.Remove(temporaryPath)
if err := file.Chmod(0o600); err != nil {
_ = file.Close()
return fmt.Errorf("设置导出配置权限失败: %w", err)
}
if _, err := file.Write(data); err != nil {
_ = file.Close()
return fmt.Errorf("写入导出配置失败: %w", err)
}
if err := file.Sync(); err != nil {
_ = file.Close()
return fmt.Errorf("同步导出配置失败: %w", err)
}
if err := file.Close(); err != nil {
return fmt.Errorf("关闭导出配置失败: %w", err)
}
if err := replaceExportFile(temporaryPath, path); err != nil {
return fmt.Errorf("替换导出配置失败: %w", err)
}
return nil
}
func ensureConfigImportAllowed(state ProxyState) error {
if state.BackendRunning || state.ProxyRunning || state.Running {
return errors.New("服务运行中不能导入完整配置,请先停止服务")
}
return nil
}
func readImportedUserConfig(path string) ([]byte, error) {
sourcePath := strings.TrimSpace(path)
if sourcePath == "" {
return nil, errors.New("导入路径不能为空")
}
file, err := os.Open(sourcePath)
if err != nil {
return nil, fmt.Errorf("打开导入配置失败: %w", err)
}
defer file.Close()
info, err := file.Stat()
if err != nil {
return nil, fmt.Errorf("读取导入配置信息失败: %w", err)
}
if !info.Mode().IsRegular() {
return nil, errors.New("导入配置必须是普通文件")
}
if info.Size() > maxConfigTransferFileSize {
return nil, fmt.Errorf("导入配置不能超过 %d MiB", maxConfigTransferFileSize>>20)
}
data, err := io.ReadAll(io.LimitReader(file, maxConfigTransferFileSize+1))
if err != nil {
return nil, fmt.Errorf("读取导入配置失败: %w", err)
}
if len(data) > maxConfigTransferFileSize {
return nil, fmt.Errorf("导入配置不能超过 %d MiB", maxConfigTransferFileSize>>20)
}
return data, nil
}
func decodeImportedUserConfig(data []byte) (serverconfig.Config, error) {
if err := validateImportedUserConfigDocument(data); err != nil {
return serverconfig.Config{}, err
}
decoder := yaml.NewDecoder(bytes.NewReader(data))
decoder.KnownFields(true)
var document importedUserConfigDocument
if err := decoder.Decode(&document); err != nil {
return serverconfig.Config{}, fmt.Errorf("导入配置包含未知字段或无效 YAML: %w", err)
}
var trailing any
if err := decoder.Decode(&trailing); err == nil {
return serverconfig.Config{}, errors.New("导入配置只能包含单个 YAML 文档")
} else if !errors.Is(err, io.EOF) {
return serverconfig.Config{}, fmt.Errorf("解析导入配置尾部失败: %w", err)
}
normalized, err := serverconfig.NormalizeConfig(document.Config)
if err != nil {
return serverconfig.Config{}, fmt.Errorf("导入配置校验失败: %w", err)
}
return normalized, nil
}
func validateImportedUserConfigDocument(data []byte) error {
decoder := yaml.NewDecoder(bytes.NewReader(data))
var document yaml.Node
if err := decoder.Decode(&document); err != nil {
if errors.Is(err, io.EOF) {
return errors.New("导入配置不能为空")
}
return fmt.Errorf("导入配置不是有效 YAML: %w", err)
}
if len(document.Content) != 1 || document.Content[0].Kind != yaml.MappingNode {
return errors.New("导入配置顶层必须是 YAML 对象")
}
return nil
}
@@ -1,9 +0,0 @@
//go:build !windows
package client
import "os"
func replaceExportFile(sourcePath, targetPath string) error {
return os.Rename(sourcePath, targetPath)
}
@@ -1,21 +0,0 @@
//go:build windows
package client
import "golang.org/x/sys/windows"
func replaceExportFile(sourcePath, targetPath string) error {
source, err := windows.UTF16PtrFromString(sourcePath)
if err != nil {
return err
}
target, err := windows.UTF16PtrFromString(targetPath)
if err != nil {
return err
}
return windows.MoveFileEx(
source,
target,
windows.MOVEFILE_REPLACE_EXISTING|windows.MOVEFILE_WRITE_THROUGH,
)
}
-245
View File
@@ -1,245 +0,0 @@
package client
import (
"os"
"path/filepath"
"runtime"
"strings"
"testing"
"time"
serverconfig "cursor/internal/backend/server/config"
"gopkg.in/yaml.v3"
)
func TestExportAndImportUserConfigRoundTrip(t *testing.T) {
source := newConfigTransferTestService(t)
want := serverconfig.DefaultConfig()
want.Log = true
want.ModelAdapters = []serverconfig.ModelAdapterConfig{{
DisplayName: "迁移模型",
Type: "openai",
BaseURL: "https://provider.example/v1",
APIKey: "migration-secret",
TooltipData: "迁移备注",
ModelID: "model-a",
ReasoningEffort: "medium",
OpenAIEndpoint: "/v1/responses",
}}
if err := source.SaveUserConfig(want); err != nil {
t.Fatalf("SaveUserConfig() error = %v", err)
}
exportPath, err := source.ExportUserConfig(filepath.Join(t.TempDir(), "cursor-byok-backup"))
if err != nil {
t.Fatalf("ExportUserConfig() error = %v", err)
}
if filepath.Ext(exportPath) != ".yaml" {
t.Fatalf("ExportUserConfig() path = %q, want .yaml extension", exportPath)
}
if runtime.GOOS != "windows" {
info, statErr := os.Stat(exportPath)
if statErr != nil {
t.Fatalf("Stat() error = %v", statErr)
}
if gotMode := info.Mode().Perm(); gotMode != 0o600 {
t.Fatalf("export mode = %o, want 600", gotMode)
}
}
target := newConfigTransferTestService(t)
got, err := target.ImportUserConfig(exportPath)
if err != nil {
t.Fatalf("ImportUserConfig() error = %v", err)
}
if !got.Log || len(got.ModelAdapters) != 1 {
t.Fatalf("ImportUserConfig() = %#v", got)
}
adapter := got.ModelAdapters[0]
if adapter.DisplayName != "迁移模型" || adapter.APIKey != "migration-secret" || adapter.ModelID != "model-a" {
t.Fatalf("imported adapter = %#v", adapter)
}
persisted, err := target.LoadUserConfig()
if err != nil {
t.Fatalf("LoadUserConfig() error = %v", err)
}
if len(persisted.ModelAdapters) != 1 || persisted.ModelAdapters[0].APIKey != "migration-secret" {
t.Fatalf("persisted config = %#v", persisted)
}
}
func TestWriteExportedUserConfigReplacesExistingFile(t *testing.T) {
directory := t.TempDir()
path := filepath.Join(directory, "config.yaml")
if err := os.WriteFile(path, []byte("old"), 0o600); err != nil {
t.Fatalf("WriteFile() error = %v", err)
}
if err := writeExportedUserConfig(path, []byte("new")); err != nil {
t.Fatalf("writeExportedUserConfig() error = %v", err)
}
data, err := os.ReadFile(path)
if err != nil {
t.Fatalf("ReadFile() error = %v", err)
}
if string(data) != "new" {
t.Fatalf("exported content = %q, want new", data)
}
matches, err := filepath.Glob(filepath.Join(directory, ".cursor-byok-config-*.tmp"))
if err != nil || len(matches) != 0 {
t.Fatalf("temporary exports = %v, error = %v", matches, err)
}
}
func TestEnsureConfigImportAllowedRejectsRunningService(t *testing.T) {
for _, state := range []ProxyState{
{BackendRunning: true},
{ProxyRunning: true},
{Running: true},
} {
if err := ensureConfigImportAllowed(state); err == nil {
t.Fatalf("ensureConfigImportAllowed(%+v) error = nil", state)
}
}
if err := ensureConfigImportAllowed(ProxyState{}); err != nil {
t.Fatalf("ensureConfigImportAllowed(stopped) error = %v", err)
}
}
func TestImportUserConfigWaitsForLifecycleTransition(t *testing.T) {
service := newConfigTransferTestService(t)
path := filepath.Join(t.TempDir(), "config.yaml")
if err := os.WriteFile(path, []byte("modelAdapters: []\n"), 0o600); err != nil {
t.Fatalf("WriteFile() error = %v", err)
}
service.lifecycleMu.Lock()
done := make(chan error, 1)
go func() {
_, err := service.ImportUserConfig(path)
done <- err
}()
select {
case err := <-done:
service.lifecycleMu.Unlock()
t.Fatalf("ImportUserConfig() completed during lifecycle transition: %v", err)
case <-time.After(100 * time.Millisecond):
}
service.lifecycleMu.Unlock()
select {
case err := <-done:
if err != nil {
t.Fatalf("ImportUserConfig() error = %v", err)
}
case <-time.After(5 * time.Second):
t.Fatal("ImportUserConfig() did not resume after lifecycle transition")
}
}
func TestDecodeImportedUserConfigRejectsNonMappingDocuments(t *testing.T) {
for _, raw := range []string{"", "null\n", "[]\n", "value\n"} {
if _, err := decodeImportedUserConfig([]byte(raw)); err == nil {
t.Fatalf("decodeImportedUserConfig(%q) error = nil", raw)
}
}
}
func TestDecodeImportedUserConfigAcceptsEmptyReasoningEffort(t *testing.T) {
raw := []byte("modelAdapters:\n - displayName: model\n type: openai\n baseURL: https://example.com/v1\n apiKey: secret\n tooltipData: migrated model\n modelID: model-a\n reasoningEffort: ''\n openAIEndpoint: /v1/responses\n")
got, err := decodeImportedUserConfig(raw)
if err != nil {
t.Fatalf("decodeImportedUserConfig() error = %v", err)
}
if got.ModelAdapters[0].ReasoningEffort != "" {
t.Fatalf("reasoningEffort = %q, want empty", got.ModelAdapters[0].ReasoningEffort)
}
}
func TestImportUserConfigRejectsUnknownFieldsWithoutOverwriting(t *testing.T) {
service := newConfigTransferTestService(t)
current := serverconfig.DefaultConfig()
current.Log = true
if err := service.SaveUserConfig(current); err != nil {
t.Fatalf("SaveUserConfig() error = %v", err)
}
path := filepath.Join(t.TempDir(), "unknown.yaml")
if err := os.WriteFile(path, []byte("modelAdapters: []\nunknownSetting: true\n"), 0o600); err != nil {
t.Fatalf("WriteFile() error = %v", err)
}
if _, err := service.ImportUserConfig(path); err == nil || !strings.Contains(err.Error(), "未知字段") {
t.Fatalf("ImportUserConfig() error = %v, want unknown field error", err)
}
persisted, err := service.LoadUserConfig()
if err != nil {
t.Fatalf("LoadUserConfig() error = %v", err)
}
if !persisted.Log {
t.Fatal("invalid import overwrote the existing config")
}
}
func TestImportUserConfigRejectsMultipleDocuments(t *testing.T) {
service := newConfigTransferTestService(t)
path := filepath.Join(t.TempDir(), "multiple.yaml")
content := "modelAdapters: []\n---\nmodelAdapters: []\n"
if err := os.WriteFile(path, []byte(content), 0o600); err != nil {
t.Fatalf("WriteFile() error = %v", err)
}
if _, err := service.ImportUserConfig(path); err == nil || !strings.Contains(err.Error(), "单个 YAML 文档") {
t.Fatalf("ImportUserConfig() error = %v, want multiple document error", err)
}
}
func TestImportUserConfigRejectsOversizedFile(t *testing.T) {
service := newConfigTransferTestService(t)
path := filepath.Join(t.TempDir(), "oversized.yaml")
content := make([]byte, maxConfigTransferFileSize+1)
if err := os.WriteFile(path, content, 0o600); err != nil {
t.Fatalf("WriteFile() error = %v", err)
}
if _, err := service.ImportUserConfig(path); err == nil || !strings.Contains(err.Error(), "不能超过") {
t.Fatalf("ImportUserConfig() error = %v, want size limit error", err)
}
}
func TestDecodeImportedUserConfigNormalizesValues(t *testing.T) {
raw := []byte("backendListenAddr: ' 127.0.0.1:12345 '\nproxyListenAddr: '127.0.0.1:12346'\nmodelAdapters: []\n")
got, err := decodeImportedUserConfig(raw)
if err != nil {
t.Fatalf("decodeImportedUserConfig() error = %v", err)
}
if got.BackendListenAddr != "127.0.0.1:12345" || got.ProxyListenAddr != "127.0.0.1:12346" {
t.Fatalf("decodeImportedUserConfig() = %#v", got)
}
encoded, err := yaml.Marshal(got)
if err != nil {
t.Fatalf("yaml.Marshal() error = %v", err)
}
if !strings.Contains(string(encoded), "backendListenAddr: 127.0.0.1:12345") {
t.Fatalf("encoded config = %s", encoded)
}
}
func TestDecodeImportedUserConfigAcceptsLegacyRouting(t *testing.T) {
raw := []byte("modelAdapters: []\nrouting:\n strategy: legacy\n")
got, err := decodeImportedUserConfig(raw)
if err != nil {
t.Fatalf("decodeImportedUserConfig() error = %v", err)
}
if len(got.ModelAdapters) != 0 {
t.Fatalf("decodeImportedUserConfig() = %#v", got)
}
}
func newConfigTransferTestService(t *testing.T) *ProxyService {
t.Helper()
root := t.TempDir()
return &ProxyService{
store: serverconfig.NewStore(filepath.Join(root, "config.yaml"), filepath.Join(root, "logs")),
}
}
-4
View File
@@ -5,7 +5,6 @@ import (
goruntime "runtime"
"cursor/internal/cursor"
"cursor/internal/logger"
)
// ApplyCursorSettings 用于处理与 ApplyCursorSettings 相关的逻辑。
@@ -22,9 +21,6 @@ func (s *ProxyService) ApplyCursorSettings() error {
if err != nil {
return fmt.Errorf("ensure ca cert file: %w", err)
}
if err := cursor.EnsureLegacySharedCACertRemoved(); err != nil {
logger.Errorf("remove legacy shared ca cert failed, continuing with installation CA: %v", err)
}
switch goruntime.GOOS {
case "windows":
-35
View File
@@ -1,35 +0,0 @@
package client
import (
"context"
"fmt"
"time"
"cursor/internal/cursoraccount"
)
type CursorAccountStatus = cursoraccount.Status
func (s *ProxyService) GetCursorAccountStatus() CursorAccountStatus {
if s == nil || s.cursorAccount == nil {
return CursorAccountStatus{State: cursoraccount.StateSignedOut}
}
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
s.cursorAccount.EnsureEmail(ctx)
return s.cursorAccount.Status()
}
func (s *ProxyService) StartCursorAccountLogin() (CursorAccountStatus, error) {
if s == nil || s.cursorAccount == nil {
return CursorAccountStatus{State: cursoraccount.StateError}, fmt.Errorf("Cursor 账号服务未初始化")
}
return s.cursorAccount.StartLogin()
}
func (s *ProxyService) DisconnectCursorAccount() (CursorAccountStatus, error) {
if s == nil || s.cursorAccount == nil {
return CursorAccountStatus{State: cursoraccount.StateSignedOut}, nil
}
return s.cursorAccount.Disconnect()
}
+2 -13
View File
@@ -11,7 +11,6 @@ import (
"cursor/internal/logger"
"cursor/internal/mitm"
"cursor/internal/netproxy"
localruntime "cursor/internal/runtime"
"github.com/wailsapp/wails/v3/pkg/application"
)
@@ -54,8 +53,6 @@ type ProxyState struct {
// StartProxy 用于处理与 StartProxy 相关的逻辑。
func (s *ProxyService) StartProxy() (ProxyState, error) {
s.lifecycleMu.Lock()
defer s.lifecycleMu.Unlock()
logger.Infof("start service requested config_path=%s logs_root=%s", s.configPath, s.logsRoot)
fail := func(step string, err error) (ProxyState, error) {
logger.Errorf("start service failed step=%s err=%v", step, err)
@@ -87,11 +84,8 @@ func (s *ProxyService) StartProxy() (ProxyState, error) {
if err := s.ensureProxy(cfg); err != nil {
return fail("ensure_proxy", err)
}
// 启动时注入账号信息
if err := cursor.InjectCursorUserInfo(localruntime.InjectAccountEmail, localruntime.InjectAuthToken); err != nil {
logger.Errorf("injectCursorUserInfo failed: %v", err)
// 不阻断启动,仅记录日志
if err := cursor.DisableCursorStatsigGates(); err != nil {
logger.Errorf("disableCursorStatsigGates failed: %v", err)
}
if s.proxy != nil && !s.proxy.IsRunning() {
@@ -129,8 +123,6 @@ func (s *ProxyService) StartProxy() (ProxyState, error) {
// StopProxy 用于处理与 StopProxy 相关的逻辑。
func (s *ProxyService) StopProxy() (ProxyState, error) {
s.lifecycleMu.Lock()
defer s.lifecycleMu.Unlock()
logger.Infof("stop service requested")
fail := func(step string, err error) (ProxyState, error) {
logger.Errorf("stop service failed step=%s err=%v", step, err)
@@ -269,9 +261,6 @@ func (s *ProxyService) ShutdownForQuit() {
finalErr = errors.Join(finalErr, err)
}
}
if s.cursorAccount != nil {
s.cursorAccount.Shutdown()
}
if finalErr != nil {
s.setLastError(finalErr)
}
@@ -950,8 +950,6 @@ func normalizeModelAdapterTestType(value string) string {
func normalizeModelAdapterTestReasoning(value string) string {
switch strings.ToLower(strings.TrimSpace(value)) {
case "":
return ""
case "low", "medium", "high", "xhigh", "max":
return strings.ToLower(strings.TrimSpace(value))
default:
@@ -1,15 +0,0 @@
package client
import (
"testing"
serverconfig "cursor/internal/backend/server/config"
)
func TestNormalizeModelAdapterTestProviderReasoningPreservesBlank(t *testing.T) {
adapter := serverconfig.ModelAdapterConfig{Type: "openai", ReasoningEffort: ""}
if got := normalizeModelAdapterTestProviderReasoning(adapter); got != "" {
t.Fatalf("reasoning effort = %q, want blank", got)
}
}
+14 -12
View File
@@ -4,7 +4,6 @@ import (
"context"
"fmt"
"net/http"
"path/filepath"
"sync"
"time"
@@ -12,7 +11,6 @@ import (
backend "cursor/internal/backend"
serverconfig "cursor/internal/backend/server/config"
"cursor/internal/certs"
"cursor/internal/cursoraccount"
"cursor/internal/logger"
"cursor/internal/mitm"
"cursor/internal/netproxy"
@@ -37,8 +35,6 @@ type ProxyService struct {
certManager *certs.Manager
// backendHost 表示当前嵌入式 backend 服务。
backendHost *backend.Host
// cursorAccount 持有仅供插件、Skills 和 MCP 控制面使用的真实 Cursor 身份。
cursorAccount *cursoraccount.Manager
// mu 表示当前声明中的 mu。
mu sync.RWMutex
@@ -49,8 +45,6 @@ type ProxyService struct {
// configMu 表示当前声明中的 configMu。
configMu sync.Mutex
// lifecycleMu 串行化服务启停与完整配置导入,避免过渡状态下切换配置。
lifecycleMu sync.Mutex
// configPath 表示当前声明中的 configPath。
configPath string
// store 表示统一的配置存储。
@@ -90,12 +84,8 @@ func NewProxyService(proxy *mitm.ProxyServer, certManager *certs.Manager, caCert
publicClient: netproxy.NewHTTPClient(publicAPITimeout),
modelTestResults: make(map[string]ModelAdapterTestResult),
}
service.cursorAccount = cursoraccount.NewManager(
filepath.Join(appdata.DataRootPath(), "cursor-account.json"),
netproxy.NewHTTPClient(publicAPITimeout),
)
service.store = serverconfig.NewStore(service.configPath, service.logsRoot)
host, err := backend.NewHost(service.store, service.cursorAccount)
host, err := service.newBackendHost()
if err != nil {
logger.Errorf("init backend host failed: %v", err)
} else {
@@ -111,7 +101,7 @@ func (s *ProxyService) ensureBackendHost() error {
if s.backendHost != nil {
return nil
}
host, err := backend.NewHost(s.store, s.cursorAccount)
host, err := s.newBackendHost()
if err != nil {
return err
}
@@ -119,6 +109,18 @@ func (s *ProxyService) ensureBackendHost() error {
return nil
}
func (s *ProxyService) newBackendHost() (*backend.Host, error) {
options := []backend.HostOption{}
if s != nil && s.certManager != nil {
certificate, err := s.certManager.CertificateForServerName("localhost")
if err != nil {
return nil, fmt.Errorf("create localhost backend certificate: %w", err)
}
options = append(options, backend.WithTLSCertificate(certificate))
}
return backend.NewHost(s.store, options...)
}
func (s *ProxyService) ensureProxy(cfg serverconfig.Config) error {
if s == nil {
return nil
-34
View File
@@ -16,7 +16,6 @@ import (
const (
darwinSecurityExe = "security"
darwinLoginKeychainName = "login.keychain-db"
legacySharedCASHA1 = "C14B7488C5AB83F098BEB2603F1135595A381FC0"
)
func getCertSHA1Fingerprint(certPEM []byte) (string, error) {
@@ -37,10 +36,7 @@ func isCACertInstalled(certPEM []byte) (bool, error) {
if err != nil {
return false, fmt.Errorf("获取证书指纹失败: %w", err)
}
return isCACertFingerprintInstalled(fingerprint)
}
func isCACertFingerprintInstalled(fingerprint string) (bool, error) {
out, err := exec.Command(darwinSecurityExe, "find-certificate", "-a", "-Z", darwinLoginKeychainName).CombinedOutput()
if err != nil {
return false, fmt.Errorf("检查 macOS 登录钥匙串失败: %w: %s", err, strings.TrimSpace(string(out)))
@@ -54,36 +50,6 @@ func isCACertFingerprintInstalled(fingerprint string) (bool, error) {
return installed, nil
}
// EnsureLegacySharedCACertRemoved removes the compromised CA shipped by older versions.
func EnsureLegacySharedCACertRemoved() error {
installed, err := isCACertFingerprintInstalled(legacySharedCASHA1)
if err != nil {
return fmt.Errorf("检查旧版共享 CA 失败: %w", err)
}
if !installed {
return nil
}
out, err := exec.Command(
darwinSecurityExe,
"delete-certificate",
"-Z", legacySharedCASHA1,
"-t",
darwinLoginKeychainName,
).CombinedOutput()
if err != nil {
return fmt.Errorf("从 macOS 登录钥匙串删除旧版共享 CA 失败: %w: %s", err, strings.TrimSpace(string(out)))
}
installed, err = isCACertFingerprintInstalled(legacySharedCASHA1)
if err != nil {
return fmt.Errorf("验证旧版共享 CA 删除状态失败: %w", err)
}
if installed {
return fmt.Errorf("删除命令已执行,但 macOS 登录钥匙串中仍存在旧版共享 CA")
}
logger.Infof("ensureLegacySharedCACertRemoved: legacy shared CA removed from macOS login keychain")
return nil
}
func installCACertToDarwinKeychain(certPEM []byte, certPath string) error {
fingerprint, err := getCertSHA1Fingerprint(certPEM)
if err != nil {
+66 -26
View File
@@ -33,14 +33,6 @@ var cursorStateDisabledStatsigGates = []string{
"disable_terminal_output_ui_streaming",
}
// cursorStateEnabledStatsigGates forces gates to true. Disabling the network
// change monitor stops the always-local extension from probing
// NetworkService/IsConnected and aborting slow in-flight requests (such as
// commit message generation) with "Network disconnected".
var cursorStateEnabledStatsigGates = []string{
"disable_network_change_monitor_local",
}
// InjectCursorUserInfo synchronizes the Cursor user-level auth cache used by the
// Settings page. It does not modify the installed Cursor app bundle.
func InjectCursorUserInfo(email, token string) error {
@@ -58,17 +50,33 @@ func InjectCursorUserInfo(email, token string) error {
}
logger.Infof(
"injectCursorUserInfo synced path=%s email=%s membership=%s subscription=%s disabled_statsig_gates=%s enabled_statsig_gates=%s",
"injectCursorUserInfo synced path=%s email=%s membership=%s subscription=%s disabled_statsig_gates=%s",
stateDBPath,
values["cursorAuth/cachedEmail"],
values["cursorAuth/stripeMembershipType"],
values["cursorAuth/stripeSubscriptionStatus"],
strings.Join(cursorStateDisabledStatsigGates, ","),
strings.Join(cursorStateEnabledStatsigGates, ","),
)
return nil
}
// DisableCursorStatsigGates preserves the local-mode feature gates without
// injecting or replacing Cursor account state.
func DisableCursorStatsigGates() error {
stateDBPath, err := resolveCursorStateDBPath()
if err != nil {
return err
}
if err := os.MkdirAll(filepath.Dir(stateDBPath), 0o755); err != nil {
return fmt.Errorf("创建 Cursor 状态目录失败: %w", err)
}
if err := disableCursorStatsigGatesInDB(stateDBPath); err != nil {
return fmt.Errorf("同步 Cursor Statsig gates 失败 path=%s: %w", stateDBPath, err)
}
logger.Infof("disableCursorStatsigGates synced path=%s gates=%s", stateDBPath, strings.Join(cursorStateDisabledStatsigGates, ","))
return nil
}
func buildCursorAuthStateValues(email, token string) map[string]string {
email = strings.TrimSpace(email)
token = strings.TrimSpace(token)
@@ -129,7 +137,7 @@ func syncCursorAuthStateDB(path string, values map[string]string) error {
}
}
if err := syncCursorStatsigGateOverrides(ctx, tx); err != nil {
if err := disableCursorStatsigGates(ctx, tx); err != nil {
return err
}
@@ -140,7 +148,45 @@ func syncCursorAuthStateDB(path string, values map[string]string) error {
return nil
}
func syncCursorStatsigGateOverrides(ctx context.Context, tx *sql.Tx) error {
func disableCursorStatsigGatesInDB(path string) error {
db, err := sql.Open("sqlite", path)
if err != nil {
return err
}
defer db.Close()
db.SetMaxOpenConns(1)
db.SetMaxIdleConns(1)
ctx := context.Background()
if _, err := db.ExecContext(ctx, fmt.Sprintf("PRAGMA busy_timeout = %d", cursorStateSQLiteBusyTimeoutMS)); err != nil {
return err
}
if _, err := db.ExecContext(ctx, "CREATE TABLE IF NOT EXISTS ItemTable (key TEXT UNIQUE ON CONFLICT REPLACE, value BLOB)"); err != nil {
return err
}
tx, err := db.BeginTx(ctx, &sql.TxOptions{})
if err != nil {
return err
}
committed := false
defer func() {
if !committed {
_ = tx.Rollback()
}
}()
if err := disableCursorStatsigGates(ctx, tx); err != nil {
return err
}
if err := tx.Commit(); err != nil {
return err
}
committed = true
return nil
}
func disableCursorStatsigGates(ctx context.Context, tx *sql.Tx) error {
var raw []byte
err := tx.QueryRowContext(ctx, "SELECT value FROM ItemTable WHERE key = ?", cursorStateStatsigBootstrapKey).Scan(&raw)
if err != nil {
@@ -163,15 +209,9 @@ func syncCursorStatsigGateOverrides(ctx context.Context, tx *sql.Tx) error {
hashUsed, _ := payload["hash_used"].(string)
for _, gate := range cursorStateDisabledStatsigGates {
setCursorStatsigGate(featureGates, gate, false, "local_disabled")
disableCursorStatsigGate(featureGates, gate)
if strings.EqualFold(hashUsed, "djb2") {
setCursorStatsigGate(featureGates, cursorStateDJB2Hash(gate), false, "local_disabled")
}
}
for _, gate := range cursorStateEnabledStatsigGates {
setCursorStatsigGate(featureGates, gate, true, "local_enabled")
if strings.EqualFold(hashUsed, "djb2") {
setCursorStatsigGate(featureGates, cursorStateDJB2Hash(gate), true, "local_enabled")
disableCursorStatsigGate(featureGates, cursorStateDJB2Hash(gate))
}
}
@@ -185,21 +225,21 @@ func syncCursorStatsigGateOverrides(ctx context.Context, tx *sql.Tx) error {
return nil
}
func setCursorStatsigGate(featureGates map[string]any, key string, value bool, rule string) {
func disableCursorStatsigGate(featureGates map[string]any, key string) {
gate, _ := featureGates[key].(map[string]any)
if gate == nil {
gate = map[string]any{
"name": key,
"rule_id": rule,
"ruleID": rule,
"group_name": rule,
"groupName": rule,
"rule_id": "local_disabled",
"ruleID": "local_disabled",
"group_name": "local_disabled",
"groupName": "local_disabled",
"id_type": "userID",
"idType": "userID",
}
featureGates[key] = gate
}
gate["value"] = value
gate["value"] = false
}
func cursorStateDJB2Hash(value string) string {
+49
View File
@@ -59,6 +59,55 @@ func TestSyncCursorAuthStateDBDisablesCachedTerminalOutputUIStreamingIdempotentl
}
}
func TestDisableCursorStatsigGatesInDBDoesNotInjectAuthState(t *testing.T) {
path := filepath.Join(t.TempDir(), "state.vscdb")
db, err := sql.Open("sqlite", path)
if err != nil {
t.Fatalf("open temporary state db: %v", err)
}
if _, err := db.Exec("CREATE TABLE ItemTable (key TEXT UNIQUE ON CONFLICT REPLACE, value BLOB)"); err != nil {
db.Close()
t.Fatalf("create ItemTable: %v", err)
}
bootstrap := map[string]any{
"feature_gates": map[string]any{},
"hash_used": "none",
}
raw, err := json.Marshal(bootstrap)
if err != nil {
db.Close()
t.Fatalf("encode bootstrap: %v", err)
}
if _, err := db.Exec("INSERT INTO ItemTable(key, value) VALUES(?, ?)", cursorStateStatsigBootstrapKey, raw); err != nil {
db.Close()
t.Fatalf("insert bootstrap: %v", err)
}
if err := db.Close(); err != nil {
t.Fatalf("close setup db: %v", err)
}
if err := disableCursorStatsigGatesInDB(path); err != nil {
t.Fatalf("disable statsig gates: %v", err)
}
updated := readCursorStatsigBootstrapForTest(t, path)
for _, gate := range cursorStateDisabledStatsigGates {
assertCursorStatsigGateValueForTest(t, updated, gate, false)
}
db, err = sql.Open("sqlite", path)
if err != nil {
t.Fatalf("reopen state db: %v", err)
}
defer db.Close()
var authKeyCount int
if err := db.QueryRow("SELECT COUNT(*) FROM ItemTable WHERE key LIKE 'cursorAuth/%'").Scan(&authKeyCount); err != nil {
t.Fatalf("count auth keys: %v", err)
}
if authKeyCount != 0 {
t.Fatalf("statsig sync injected %d auth keys", authKeyCount)
}
}
func readCursorStatsigBootstrapForTest(t *testing.T, path string) []byte {
t.Helper()
db, err := sql.Open("sqlite", path)
-27
View File
@@ -20,7 +20,6 @@ const (
windowsCertutilExe = "certutil.exe"
windowsPowerShellExe = "powershell.exe"
windowsUserCancelCode = 1223
legacySharedCASHA1 = "C14B7488C5AB83F098BEB2603F1135595A381FC0"
)
// getCertThumbprint 获取证书的SHA1指纹,用于唯一标识证书
@@ -52,10 +51,7 @@ func isCACertInstalled(certPEM []byte) (bool, error) {
if err != nil {
return false, fmt.Errorf("获取证书指纹失败: %w", err)
}
return isCACertThumbprintInstalled(thumbprint)
}
func isCACertThumbprintInstalled(thumbprint string) (bool, error) {
cmd := exec.Command(windowsCertutilExe, "-verifystore", windowsRootStoreName, thumbprint)
cmd.SysProcAttr = hideWindow()
output, err := cmd.CombinedOutput()
@@ -80,29 +76,6 @@ func isCACertThumbprintInstalled(thumbprint string) (bool, error) {
return false, nil
}
// EnsureLegacySharedCACertRemoved removes the compromised CA shipped by older versions.
func EnsureLegacySharedCACertRemoved() error {
installed, err := isCACertThumbprintInstalled(legacySharedCASHA1)
if err != nil {
return fmt.Errorf("检查旧版共享 CA 失败: %w", err)
}
if !installed {
return nil
}
if err := runElevatedCertutil("-delstore", windowsRootStoreName, legacySharedCASHA1); err != nil {
return fmt.Errorf("从 Windows 系统信任存储删除旧版共享 CA 失败: %w", err)
}
installed, err = isCACertThumbprintInstalled(legacySharedCASHA1)
if err != nil {
return fmt.Errorf("验证旧版共享 CA 删除状态失败: %w", err)
}
if installed {
return fmt.Errorf("删除命令已执行,但 Windows 系统信任存储中仍存在旧版共享 CA")
}
logger.Infof("ensureLegacySharedCACertRemoved: legacy shared CA removed from Windows system store")
return nil
}
func quotePowerShellLiteral(value string) string {
return "'" + strings.ReplaceAll(value, "'", "''") + "'"
}
-5
View File
@@ -8,8 +8,3 @@ import "fmt"
func EnsureCACertInstalled(_ []byte, certPath string) error {
return fmt.Errorf("ensureCACertInstalled: 当前平台暂不支持,certPath=%s", certPath)
}
// EnsureLegacySharedCACertRemoved is a no-op on unsupported platforms.
func EnsureLegacySharedCACertRemoved() error {
return nil
}
-589
View File
@@ -1,589 +0,0 @@
package cursoraccount
import (
"bytes"
"context"
"crypto/rand"
"crypto/sha256"
"encoding/base64"
"encoding/json"
"errors"
"fmt"
"io"
"net/http"
"net/url"
"os"
"path/filepath"
"strings"
"sync"
"time"
"cursor/gen/aiserverv1"
"cursor/internal/backend/server/upstream"
"github.com/google/uuid"
"github.com/pkg/browser"
"google.golang.org/protobuf/proto"
)
const (
StateSignedOut = "signed_out"
StateWaiting = "waiting"
StateSignedIn = "signed_in"
StateError = "error"
websiteURL = "https://cursor.com"
backendURL = "https://api2.cursor.sh"
authClientID = "KbZUR41cY7W6zRSdpSUJ7I7mLYBKOCmB"
loginTimeout = 10 * time.Minute
pollInterval = time.Second
refreshMargin = 2 * time.Minute
)
var ErrNotSignedIn = errors.New("尚未在 cursor-byok 中登录 Cursor 账号")
// Status 是可安全返回给前端的脱敏账号状态。
type Status struct {
State string `json:"state"`
AuthID string `json:"authId"`
Email string `json:"email"`
Error string `json:"error"`
}
type credentials struct {
AccessToken string `json:"accessToken"`
RefreshToken string `json:"refreshToken"`
AuthID string `json:"authId"`
Email string `json:"email,omitempty"`
}
type pollResponse struct {
AccessToken string `json:"accessToken"`
RefreshToken string `json:"refreshToken"`
AuthID string `json:"authId"`
}
type refreshResponse struct {
AccessToken string `json:"access_token"`
RefreshToken string `json:"refresh_token"`
ShouldLogout bool `json:"shouldLogout"`
}
// Manager 持有 cursor-byok 自己的 Cursor 登录态,不读写 Cursor 客户端状态库。
type Manager struct {
path string
client *http.Client
mu sync.RWMutex
credentials credentials
state string
lastError string
loginCancel context.CancelFunc
loginGeneration uint64
refreshMu sync.Mutex
}
func NewManager(path string, client *http.Client) *Manager {
if client == nil {
client = &http.Client{Timeout: 15 * time.Second}
}
manager := &Manager{
path: strings.TrimSpace(path),
client: client,
state: StateSignedOut,
}
if err := manager.load(); err != nil {
manager.state = StateError
manager.lastError = fmt.Sprintf("读取 Cursor 账号凭据失败: %v", err)
}
return manager
}
func (manager *Manager) Status() Status {
if manager == nil {
return Status{State: StateSignedOut}
}
manager.mu.RLock()
defer manager.mu.RUnlock()
return Status{
State: manager.state,
AuthID: manager.credentials.AuthID,
Email: manager.credentials.Email,
Error: manager.lastError,
}
}
// EnsureEmail backfills a human-readable identity for credentials saved by
// builds that only persisted authId. Profile lookup failure does not invalidate
// an otherwise usable control-plane login.
func (manager *Manager) EnsureEmail(ctx context.Context) {
if manager == nil || !manager.SignedIn() {
return
}
current, generation := manager.snapshotCredentials()
if strings.TrimSpace(current.Email) != "" {
return
}
authorization, err := manager.Authorization(ctx)
if err != nil {
return
}
profile, err := manager.fetchProfile(ctx, authorization)
if err != nil || strings.TrimSpace(profile.GetEmail()) == "" {
return
}
current, currentGeneration := manager.snapshotCredentials()
if currentGeneration != generation {
return
}
current.Email = strings.TrimSpace(profile.GetEmail())
_ = manager.commitCredentials(generation, current)
}
func (manager *Manager) SignedIn() bool {
if manager == nil {
return false
}
manager.mu.RLock()
defer manager.mu.RUnlock()
return manager.state == StateSignedIn && strings.TrimSpace(manager.credentials.AccessToken) != ""
}
// StartLogin 启动官方浏览器 PKCE 登录,并在后台等待登录结果。
func (manager *Manager) StartLogin() (Status, error) {
if manager == nil {
return Status{State: StateError}, fmt.Errorf("Cursor 账号服务未初始化")
}
verifierBytes := make([]byte, 32)
if _, err := rand.Read(verifierBytes); err != nil {
return manager.Status(), fmt.Errorf("生成 Cursor 登录校验码失败: %w", err)
}
verifier := base64.RawURLEncoding.EncodeToString(verifierBytes)
challengeBytes := sha256.Sum256([]byte(verifier))
challenge := base64.RawURLEncoding.EncodeToString(challengeBytes[:])
loginID := uuid.NewString()
loginURL, err := buildLoginURL(loginID, challenge)
if err != nil {
return manager.Status(), err
}
ctx, cancel := context.WithTimeout(context.Background(), loginTimeout)
manager.mu.Lock()
if manager.loginCancel != nil {
manager.loginCancel()
}
manager.loginGeneration++
generation := manager.loginGeneration
manager.loginCancel = cancel
manager.state = StateWaiting
manager.lastError = ""
manager.mu.Unlock()
if err := browser.OpenURL(loginURL); err != nil {
cancel()
manager.finishWithError(generation, fmt.Sprintf("打开 Cursor 登录页面失败: %v", err))
return manager.Status(), err
}
go manager.pollLogin(ctx, generation, loginID, verifier)
return manager.Status(), nil
}
// Disconnect 只清除 cursor-byok 自己保存的账号,不调用 Cursor 客户端 logout。
func (manager *Manager) Disconnect() (Status, error) {
if manager == nil {
return Status{State: StateSignedOut}, nil
}
manager.mu.Lock()
manager.loginGeneration++
if manager.loginCancel != nil {
manager.loginCancel()
manager.loginCancel = nil
}
manager.credentials = credentials{}
manager.state = StateSignedOut
manager.lastError = ""
manager.mu.Unlock()
err := os.Remove(manager.path)
if err != nil && !errors.Is(err, os.ErrNotExist) {
manager.mu.Lock()
manager.state = StateError
manager.lastError = fmt.Sprintf("清除 Cursor 账号凭据失败: %v", err)
manager.mu.Unlock()
return manager.Status(), err
}
return manager.Status(), nil
}
func (manager *Manager) Shutdown() {
if manager == nil {
return
}
manager.mu.Lock()
manager.loginGeneration++
if manager.loginCancel != nil {
manager.loginCancel()
manager.loginCancel = nil
}
manager.mu.Unlock()
}
// Authorization 返回官方控制面请求使用的真实 Cursor Bearer 身份。
func (manager *Manager) Authorization(ctx context.Context) (string, error) {
if manager == nil {
return "", ErrNotSignedIn
}
manager.refreshMu.Lock()
defer manager.refreshMu.Unlock()
creds, generation := manager.snapshotCredentials()
if strings.TrimSpace(creds.AccessToken) == "" {
return "", ErrNotSignedIn
}
if !tokenNeedsRefresh(creds.AccessToken, time.Now()) {
return bearer(creds.AccessToken), nil
}
if strings.TrimSpace(creds.RefreshToken) == "" {
manager.setAuthorizationError(generation, "Cursor 登录已过期,请重新登录")
return "", fmt.Errorf("Cursor 登录已过期且没有刷新令牌")
}
updated, shouldLogout, err := manager.refresh(ctx, creds)
if err != nil {
manager.setAuthorizationError(generation, fmt.Sprintf("刷新 Cursor 登录失败: %v", err))
return "", err
}
if shouldLogout {
manager.invalidateAuthorization(generation, "Cursor 登录已失效,请重新登录")
return "", ErrNotSignedIn
}
if err := manager.commitCredentials(generation, updated); err != nil {
return "", err
}
return bearer(updated.AccessToken), nil
}
func (manager *Manager) pollLogin(ctx context.Context, generation uint64, loginID string, verifier string) {
defer func() {
manager.mu.Lock()
if manager.loginGeneration == generation {
manager.loginCancel = nil
}
manager.mu.Unlock()
}()
for {
result, pending, err := manager.pollOnce(ctx, loginID, verifier)
if err == nil && !pending {
creds := credentials{
AccessToken: strings.TrimSpace(result.AccessToken),
RefreshToken: strings.TrimSpace(result.RefreshToken),
AuthID: strings.TrimSpace(result.AuthID),
}
if creds.AccessToken == "" {
manager.finishWithError(generation, "Cursor 登录响应缺少 access token")
return
}
if profile, profileErr := manager.fetchProfile(ctx, bearer(creds.AccessToken)); profileErr == nil {
creds.Email = strings.TrimSpace(profile.GetEmail())
}
_ = manager.commitCredentials(generation, creds)
return
}
if err != nil && !isRetryablePollError(err) {
manager.finishWithError(generation, fmt.Sprintf("Cursor 登录失败: %v", err))
return
}
select {
case <-ctx.Done():
if errors.Is(ctx.Err(), context.DeadlineExceeded) {
manager.finishWithError(generation, "Cursor 登录等待超时,请重试")
}
return
case <-time.After(pollInterval):
}
}
}
func (manager *Manager) fetchProfile(ctx context.Context, authorization string) (*aiserverv1.GetMeResponse, error) {
body, err := proto.Marshal(&aiserverv1.GetMeRequest{})
if err != nil {
return nil, err
}
req, err := http.NewRequestWithContext(ctx, http.MethodPost, backendURL+"/aiserver.v1.DashboardService/GetMe", bytes.NewReader(body))
if err != nil {
return nil, err
}
req.Header.Set("authorization", authorization)
req.Header.Set("x-cursor-checksum", upstream.BuildCursorChecksum(authorization))
req.Header.Set("content-type", "application/proto")
req.Header.Set("accept", "application/proto")
req.Header.Set("connect-protocol-version", "1")
resp, err := manager.client.Do(req)
if err != nil {
return nil, err
}
defer resp.Body.Close()
responseBody, err := io.ReadAll(io.LimitReader(resp.Body, 1024*1024))
if err != nil {
return nil, err
}
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
return nil, fmt.Errorf("GetMe 返回 HTTP %d", resp.StatusCode)
}
profile := &aiserverv1.GetMeResponse{}
if err := proto.Unmarshal(responseBody, profile); err != nil {
return nil, err
}
return profile, nil
}
func (manager *Manager) pollOnce(ctx context.Context, loginID string, verifier string) (pollResponse, bool, error) {
endpoint, err := url.Parse(backendURL + "/auth/poll")
if err != nil {
return pollResponse{}, false, err
}
query := endpoint.Query()
query.Set("uuid", loginID)
query.Set("verifier", verifier)
endpoint.RawQuery = query.Encode()
req, err := http.NewRequestWithContext(ctx, http.MethodGet, endpoint.String(), nil)
if err != nil {
return pollResponse{}, false, err
}
resp, err := manager.client.Do(req)
if err != nil {
return pollResponse{}, false, err
}
defer resp.Body.Close()
if resp.StatusCode == http.StatusNotFound {
_, _ = io.Copy(io.Discard, io.LimitReader(resp.Body, 64*1024))
return pollResponse{}, true, nil
}
body, err := io.ReadAll(io.LimitReader(resp.Body, 1024*1024))
if err != nil {
return pollResponse{}, false, err
}
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
return pollResponse{}, false, fmt.Errorf("登录服务返回 HTTP %d", resp.StatusCode)
}
result := pollResponse{}
if err := json.Unmarshal(body, &result); err != nil {
return pollResponse{}, false, fmt.Errorf("解析登录响应失败: %w", err)
}
return result, false, nil
}
func (manager *Manager) refresh(ctx context.Context, current credentials) (credentials, bool, error) {
payload, err := json.Marshal(map[string]string{
"grant_type": "refresh_token",
"client_id": authClientID,
"refresh_token": current.RefreshToken,
})
if err != nil {
return credentials{}, false, err
}
req, err := http.NewRequestWithContext(ctx, http.MethodPost, backendURL+"/oauth/token", bytes.NewReader(payload))
if err != nil {
return credentials{}, false, err
}
req.Header.Set("content-type", "application/json")
resp, err := manager.client.Do(req)
if err != nil {
return credentials{}, false, err
}
defer resp.Body.Close()
body, err := io.ReadAll(io.LimitReader(resp.Body, 1024*1024))
if err != nil {
return credentials{}, false, err
}
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
return credentials{}, false, fmt.Errorf("刷新服务返回 HTTP %d", resp.StatusCode)
}
result := refreshResponse{}
if err := json.Unmarshal(body, &result); err != nil {
return credentials{}, false, fmt.Errorf("解析刷新响应失败: %w", err)
}
if result.ShouldLogout {
return credentials{}, true, nil
}
if strings.TrimSpace(result.AccessToken) == "" {
return credentials{}, false, fmt.Errorf("刷新响应缺少 access token")
}
current.AccessToken = strings.TrimSpace(result.AccessToken)
if strings.TrimSpace(result.RefreshToken) != "" {
current.RefreshToken = strings.TrimSpace(result.RefreshToken)
}
return current, false, nil
}
func (manager *Manager) load() error {
if manager.path == "" {
return fmt.Errorf("Cursor 账号凭据路径为空")
}
data, err := os.ReadFile(manager.path)
if errors.Is(err, os.ErrNotExist) {
return nil
}
if err != nil {
return err
}
loaded := credentials{}
if err := json.Unmarshal(data, &loaded); err != nil {
return err
}
loaded.AccessToken = strings.TrimSpace(loaded.AccessToken)
loaded.RefreshToken = strings.TrimSpace(loaded.RefreshToken)
loaded.AuthID = strings.TrimSpace(loaded.AuthID)
loaded.Email = strings.TrimSpace(loaded.Email)
if loaded.AccessToken == "" {
return nil
}
manager.credentials = loaded
manager.state = StateSignedIn
return nil
}
func (manager *Manager) save(value credentials) error {
if manager.path == "" {
return fmt.Errorf("Cursor 账号凭据路径为空")
}
if err := os.MkdirAll(filepath.Dir(manager.path), 0o700); err != nil {
return err
}
data, err := json.MarshalIndent(value, "", " ")
if err != nil {
return err
}
tempPath := manager.path + ".tmp"
if err := os.WriteFile(tempPath, append(data, '\n'), 0o600); err != nil {
return err
}
if err := os.Chmod(tempPath, 0o600); err != nil {
_ = os.Remove(tempPath)
return err
}
if err := os.Rename(tempPath, manager.path); err != nil {
_ = os.Remove(tempPath)
return err
}
return os.Chmod(manager.path, 0o600)
}
func (manager *Manager) snapshotCredentials() (credentials, uint64) {
manager.mu.RLock()
defer manager.mu.RUnlock()
return manager.credentials, manager.loginGeneration
}
func (manager *Manager) finishWithError(generation uint64, message string) {
manager.mu.Lock()
defer manager.mu.Unlock()
if manager.loginGeneration != generation {
return
}
manager.state = StateError
manager.lastError = strings.TrimSpace(message)
}
func (manager *Manager) commitCredentials(generation uint64, value credentials) error {
manager.mu.Lock()
defer manager.mu.Unlock()
if manager.loginGeneration != generation {
return ErrNotSignedIn
}
if err := manager.save(value); err != nil {
manager.state = StateError
manager.lastError = fmt.Sprintf("保存 Cursor 登录凭据失败: %v", err)
return err
}
manager.credentials = value
manager.state = StateSignedIn
manager.lastError = ""
return nil
}
func (manager *Manager) setAuthorizationError(generation uint64, message string) {
manager.mu.Lock()
defer manager.mu.Unlock()
if manager.loginGeneration != generation {
return
}
manager.state = StateError
manager.lastError = strings.TrimSpace(message)
}
func (manager *Manager) invalidateAuthorization(generation uint64, message string) {
manager.mu.Lock()
defer manager.mu.Unlock()
if manager.loginGeneration != generation {
return
}
manager.loginGeneration++
manager.credentials = credentials{}
manager.state = StateError
manager.lastError = strings.TrimSpace(message)
_ = os.Remove(manager.path)
}
func buildLoginURL(loginID string, challenge string) (string, error) {
parsed, err := url.Parse(websiteURL + "/loginDeepControl")
if err != nil {
return "", err
}
query := parsed.Query()
query.Set("challenge", challenge)
query.Set("uuid", loginID)
query.Set("mode", "login")
query.Set("supportsSelectedTeamLogin", "true")
parsed.RawQuery = query.Encode()
return parsed.String(), nil
}
func bearer(token string) string {
value := strings.TrimSpace(token)
if strings.HasPrefix(strings.ToLower(value), "bearer ") {
return value
}
return "Bearer " + value
}
func tokenNeedsRefresh(token string, now time.Time) bool {
parts := strings.Split(strings.TrimSpace(token), ".")
if len(parts) < 2 {
return false
}
payload, err := base64.RawURLEncoding.DecodeString(parts[1])
if err != nil {
return false
}
claims := struct {
ExpiresAt json.Number `json:"exp"`
}{}
decoder := json.NewDecoder(bytes.NewReader(payload))
decoder.UseNumber()
if err := decoder.Decode(&claims); err != nil || claims.ExpiresAt == "" {
return false
}
expiresAt, err := claims.ExpiresAt.Int64()
if err != nil {
return false
}
return !now.Add(refreshMargin).Before(time.Unix(expiresAt, 0))
}
func isRetryablePollError(err error) bool {
if err == nil {
return false
}
var urlErr *url.Error
if errors.As(err, &urlErr) {
return true
}
message := strings.ToLower(err.Error())
return strings.Contains(message, "http 429") || strings.Contains(message, "http 5")
}
+14
View File
@@ -184,6 +184,19 @@ func NewProxyServer(addr, baseURL, _ string, _ string, certManager *certs.Manage
return nil, err
}
tlsConfig := &tls.Config{MinVersion: tls.VersionTLS12}
if certManager != nil {
caCertificate, err := certManager.CATLSCertificate()
if err != nil {
return nil, fmt.Errorf("load proxy backend CA: %w", err)
}
roots := x509.NewCertPool()
if caCertificate.Leaf != nil {
roots.AddCert(caCertificate.Leaf)
}
tlsConfig.RootCAs = roots
}
s := &ProxyServer{
addr: addr,
baseURL: normalizedBaseURL,
@@ -198,6 +211,7 @@ func NewProxyServer(addr, baseURL, _ string, _ string, certManager *certs.Manage
TLSHandshakeTimeout: 10 * time.Second,
ExpectContinueTimeout: 1 * time.Second,
ResponseHeaderTimeout: 60 * time.Second,
TLSClientConfig: tlsConfig,
},
},
}

Some files were not shown because too many files have changed in this diff Show More