mirror of
https://wget.la/https://github.com/leookun/cursor-byok
synced 2026-08-18 03:57:06 +08:00
Compare commits
24
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
2b8d1c9e3d | ||
|
|
7a67e76617 | ||
|
|
08ce12aa55 | ||
|
|
8f8d28880d | ||
|
|
988ba63d40 | ||
|
|
edd59dc684 | ||
|
|
da5fa34a4d | ||
|
|
4bd2359282 | ||
|
|
7838ffc6a2 | ||
|
|
ee30775be1 | ||
|
|
80a5093aa9 | ||
|
|
b9742b1667 | ||
|
|
7d3e74be59 | ||
|
|
1285bf9d62 | ||
|
|
7a724595eb | ||
|
|
00aec40e50 | ||
|
|
c10a2d475d | ||
|
|
4864d3675b | ||
|
|
3f95318a49 | ||
|
|
9373e57ebf | ||
|
|
f1992b0cfe | ||
|
|
67a9c27931 | ||
|
|
85a43115c7 | ||
|
|
b475166ba8 |
+13
-1
@@ -2,6 +2,11 @@
|
||||
|
||||
# cursor-byok
|
||||
|
||||
cursor-byok 是 Cursor 后端的本地实现。
|
||||
<br>
|
||||
<br>
|
||||
<a href="https://trendshift.io/repositories/39260?utm_source=repository-badge&utm_medium=badge&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)
|
||||
|
||||
[](https://github.com/leookun/cursor-byok/releases/latest)
|
||||
@@ -84,12 +89,19 @@ cursor-byok 在本机负责协议适配、模型请求转发、工具调用衔
|
||||
- [Telegram 交流群](https://t.me/cursor_byok)
|
||||
- QQ 交流群:`1095916242`、`1094411438`、`1095918002`、`1094419321`
|
||||
|
||||
<a href="https://trendshift.io/repositories/39260?utm_source=repository-badge&utm_medium=badge&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) 开源。
|
||||
|
||||
@@ -1,14 +1,20 @@
|
||||
<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&utm_medium=badge&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) · [Latest Release](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) · [Download](https://github.com/leookun/cursor-byok/releases/latest) · [Report an Issue](https://github.com/leookun/cursor-byok/issues) · [中文版本说明](./README-CN.md)
|
||||
|
||||
[](https://github.com/leookun/cursor-byok/releases/latest)
|
||||
[](https://github.com/leookun/cursor-byok/releases)
|
||||
[](./LICENSE)
|
||||
[](https://github.com/leookun/cursor-byok/releases/latest)
|
||||
|
||||
|
||||
|
||||
</div>
|
||||
|
||||

|
||||
@@ -84,12 +90,21 @@ 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&utm_medium=badge&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
@@ -8,7 +8,7 @@ info:
|
||||
description: "Cursor助手"
|
||||
copyright: "© 2026, Cursor助手"
|
||||
comments: "Cursor助手"
|
||||
version: "0.0.46"
|
||||
version: "0.0.47"
|
||||
|
||||
dev_mode:
|
||||
root_path: .
|
||||
|
||||
@@ -17,9 +17,9 @@
|
||||
<key>CFBundlePackageType</key>
|
||||
<string>APPL</string>
|
||||
<key>CFBundleShortVersionString</key>
|
||||
<string>0.0.46</string>
|
||||
<string>0.0.47</string>
|
||||
<key>CFBundleVersion</key>
|
||||
<string>0.0.46</string>
|
||||
<string>0.0.47</string>
|
||||
<key>LSMinimumSystemVersion</key>
|
||||
<string>12.0.0</string>
|
||||
<key>LSUIElement</key>
|
||||
|
||||
@@ -17,9 +17,9 @@
|
||||
<key>CFBundlePackageType</key>
|
||||
<string>APPL</string>
|
||||
<key>CFBundleShortVersionString</key>
|
||||
<string>0.0.46</string>
|
||||
<string>0.0.47</string>
|
||||
<key>CFBundleVersion</key>
|
||||
<string>0.0.46</string>
|
||||
<string>0.0.47</string>
|
||||
<key>LSMinimumSystemVersion</key>
|
||||
<string>12.0.0</string>
|
||||
<key>LSUIElement</key>
|
||||
|
||||
@@ -6,7 +6,7 @@
|
||||
name: "Cursor助手"
|
||||
arch: ${GOARCH}
|
||||
platform: "linux"
|
||||
version: "0.0.46"
|
||||
version: "0.0.47"
|
||||
section: "default"
|
||||
priority: "extra"
|
||||
maintainer: ${GIT_COMMITTER_NAME} <${GIT_COMMITTER_EMAIL}>
|
||||
|
||||
@@ -1,10 +1,10 @@
|
||||
{
|
||||
"fixed": {
|
||||
"file_version": "0.0.46"
|
||||
"file_version": "0.0.47"
|
||||
},
|
||||
"info": {
|
||||
"0000": {
|
||||
"ProductVersion": "0.0.46",
|
||||
"ProductVersion": "0.0.47",
|
||||
"CompanyName": "Cursor助手",
|
||||
"FileDescription": "Cursor助手",
|
||||
"LegalCopyright": "© 2026, Cursor助手",
|
||||
|
||||
@@ -14,7 +14,7 @@
|
||||
!define INFO_PRODUCTNAME "Cursor助手"
|
||||
!endif
|
||||
!ifndef INFO_PRODUCTVERSION
|
||||
!define INFO_PRODUCTVERSION "0.0.46"
|
||||
!define INFO_PRODUCTVERSION "0.0.47"
|
||||
!endif
|
||||
!ifndef INFO_COPYRIGHT
|
||||
!define INFO_COPYRIGHT "© 2026, Cursor助手"
|
||||
|
||||
@@ -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.46" processorArchitecture="*"/>
|
||||
<assemblyIdentity type="win32" name="com.cursor.wuxianxubei" version="0.0.47" processorArchitecture="*"/>
|
||||
<dependency>
|
||||
<dependentAssembly>
|
||||
<assemblyIdentity type="win32" name="Microsoft.Windows.Common-Controls" version="6.0.0.0" processorArchitecture="*" publicKeyToken="6595b64144ccf1df" language="*"/>
|
||||
|
||||
@@ -28,6 +28,7 @@ type exchangeContext struct {
|
||||
type Server struct {
|
||||
config Config
|
||||
certManager *certs.Manager
|
||||
caCertPEM []byte
|
||||
store *exchangeStore
|
||||
counter atomic.Uint64
|
||||
proxyServer *http.Server
|
||||
@@ -43,13 +44,14 @@ func New(config Config) (*Server, error) {
|
||||
if err := validateLoopbackAddress(config.UIAddr); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
manager, err := certs.NewEmbeddedManager()
|
||||
manager, caCertPEM, err := certs.NewGeneratedManager()
|
||||
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()
|
||||
|
||||
@@ -8,8 +8,6 @@ import (
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"cursor/internal/certs"
|
||||
)
|
||||
|
||||
//go:embed web/*
|
||||
@@ -93,7 +91,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(certs.EmbeddedCACertPEM())
|
||||
_, _ = writer.Write(server.caCertPEM)
|
||||
}
|
||||
|
||||
func writeJSON(writer http.ResponseWriter, status int, payload any) {
|
||||
|
||||
@@ -0,0 +1,202 @@
|
||||
<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,6 +19,7 @@ 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]" },
|
||||
|
||||
@@ -35,6 +35,7 @@ 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]" },
|
||||
@@ -150,7 +151,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。",
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -19,6 +19,7 @@
|
||||
"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",
|
||||
"24343a2096988d42": "Failed to open",
|
||||
@@ -39,9 +40,11 @@
|
||||
"37d23612f78a2e63": "Restart Now to Update",
|
||||
"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",
|
||||
@@ -56,10 +59,12 @@
|
||||
"4c0a929bb86ce912": "Current: {0}",
|
||||
"4d2b6e53be6002e5": "Cache Statistics Strategy: {0} ({1})",
|
||||
"4d8c1c5b42830791": "Unknown",
|
||||
"4e30d7c9ed2b0eee": "Not set",
|
||||
"4f0982ba1d37e51b": "Current outbound requests use environment variable proxy",
|
||||
"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}",
|
||||
@@ -78,6 +83,7 @@
|
||||
"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",
|
||||
@@ -101,11 +107,13 @@
|
||||
"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",
|
||||
"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}",
|
||||
@@ -159,12 +167,15 @@
|
||||
"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",
|
||||
"c8a52b66651d294c": "Failed to log out",
|
||||
"c8c14507b2d37395": "Reasoning Effort",
|
||||
"c98e118e0a43f078": "Model",
|
||||
"c9dd59beefd7144f": "Cache Read / (Cache Read + Non-cache Input)",
|
||||
@@ -172,6 +183,7 @@
|
||||
"ca1d1059408b3837": "Invalid turns: {0}",
|
||||
"cc5049729a2c10f1": "Test failed. Check the raw details.",
|
||||
"cd7ca5fb221e1c53": "{0} cannot be empty",
|
||||
"cfa6c803eb3fc713": "Waiting for browser login",
|
||||
"d0325067fed88e5a": "Cache hit rate {0}",
|
||||
"d20ab96566d33f25": "{0} display name cannot be empty",
|
||||
"d2243e1d44b2a94e": "Edit Model Settings",
|
||||
@@ -179,20 +191,24 @@
|
||||
"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",
|
||||
"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.",
|
||||
"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}",
|
||||
"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}",
|
||||
@@ -201,7 +217,6 @@
|
||||
"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",
|
||||
|
||||
@@ -19,6 +19,7 @@
|
||||
"1afed6a81a2512d2": "モデルを選択",
|
||||
"1baddde657dd2720": "現在のアウトバウンドリクエストはシステムプロキシを使用しています",
|
||||
"1bc77f5ab979f4c1": "モデル設定を追加",
|
||||
"1c631615c1d85c9e": "Cursor にログイン",
|
||||
"1e238093b79b3165": "空欄で 65536",
|
||||
"21296ab18ad9af25": "追加パラメータ JSON",
|
||||
"24343a2096988d42": "開けませんでした",
|
||||
@@ -39,9 +40,11 @@
|
||||
"37d23612f78a2e63": "今すぐ再起動して更新",
|
||||
"392d0dceb45998d3": "最高",
|
||||
"393df9bb13ea4900": "ヒット",
|
||||
"3ab8cc15939f3b5c": "ログアウト",
|
||||
"3af7e5489e61ea51": "更新中",
|
||||
"3c2a9f9901109e75": "{0} のタイプは OpenAI または Anthropic のみサポートします",
|
||||
"3d13868593ae4eeb": "表示言語",
|
||||
"3d52574ce1500561": "未接続",
|
||||
"3ea83f9f55062582": "公開日時: {0}",
|
||||
"3edda85621fd03b2": "件のモデルアダプター",
|
||||
"3fd47edce45b3603": "閉じる",
|
||||
@@ -56,10 +59,12 @@
|
||||
"4c0a929bb86ce912": "現在:{0}",
|
||||
"4d2b6e53be6002e5": "キャッシュ統計ポリシー:{0}({1})",
|
||||
"4d8c1c5b42830791": "不明",
|
||||
"4e30d7c9ed2b0eee": "設定しない",
|
||||
"4f0982ba1d37e51b": "現在のアウトバウンドリクエストは環境変数プロキシを使用しています",
|
||||
"5205125c0e91d346": "Anthropic モデルが1回の応答で生成できる最大 Token 数。空欄の場合はデフォルト値を使用します。",
|
||||
"54e6745ff43c9c74": "並べ替えに失敗しました",
|
||||
"56627c94a9decee6": "最大出力 Token",
|
||||
"58c6b0935a7216da": "コントリビューターのプロフィールを開けませんでした",
|
||||
"593a972852ba0004": "Cursor アシスタント | 永久無料 | カスタム API",
|
||||
"59a2195a01a8b35b": "{0}は有効なJSONオブジェクトである必要があります",
|
||||
"5aa8f5590c940829": "非キャッシュ入力:{0}",
|
||||
@@ -78,6 +83,7 @@
|
||||
"66af574b8948fe83": "{0} の API キーは必須です",
|
||||
"6744b4c6a9aa0038": "無効化",
|
||||
"675109292da4eb36": "まだテストしていません",
|
||||
"688102a402ba015a": "ログインを待っています...",
|
||||
"6a7b96f399e58138": "例: sk-xxxxxx",
|
||||
"6aa8f49cc992dfd7": "テスト",
|
||||
"6ae23d6d7cb18592": "サービスエラー",
|
||||
@@ -101,11 +107,13 @@
|
||||
"8139cb3dd11f5a67": "有効にすると、JSONオブジェクトが最終的なリクエストヘッダーを上書きします。同名のヘッダーはこの設定が優先され、値は文字列である必要があります。",
|
||||
"8151e8704a7ca89e": "一致する項目がありません",
|
||||
"83913e71fcf7ff60": "更新しました",
|
||||
"83be9cac28873059": "Cursor コントロールプレーンアカウント",
|
||||
"83fcfb4c1f2c1641": "モデルを取得",
|
||||
"8672864e90417138": "最大",
|
||||
"86df7ec743047234": "サービス稼働中",
|
||||
"899add6275682210": "空欄で 200000",
|
||||
"8a4ef3e48e4e8a5a": "有効",
|
||||
"8b8428f714611458": "モデルが reasoning_effort に対応している場合のみ推論強度を選択してください。「設定しない」を選ぶと、リクエストにこのパラメータは含まれません。値が高いほど安定しやすい反面、遅くなることがあります。",
|
||||
"8c1935935600e336": "モデルテスト",
|
||||
"8cbcf741e727dbf7": "モデル設定",
|
||||
"8d1de152be6360ce": "有効率: {0}",
|
||||
@@ -159,12 +167,15 @@
|
||||
"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": "設定フォルダー",
|
||||
"c8a52b66651d294c": "ログアウトに失敗しました",
|
||||
"c8c14507b2d37395": "推論強度",
|
||||
"c98e118e0a43f078": "モデル",
|
||||
"c9dd59beefd7144f": "キャッシュ読み取り / (キャッシュ読み取り + 非キャッシュ入力)",
|
||||
@@ -172,6 +183,7 @@
|
||||
"ca1d1059408b3837": "異常ターン: {0}",
|
||||
"cc5049729a2c10f1": "テストに失敗しました。元の詳細情報を確認してください。",
|
||||
"cd7ca5fb221e1c53": "{0}は空にできません",
|
||||
"cfa6c803eb3fc713": "ブラウザでのログインを待っています",
|
||||
"d0325067fed88e5a": "キャッシュヒット率 {0}",
|
||||
"d20ab96566d33f25": "{0} の表示名は必須です",
|
||||
"d2243e1d44b2a94e": "モデル設定を編集",
|
||||
@@ -179,20 +191,24 @@
|
||||
"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": "設定済み",
|
||||
"e01c5dae36cf8c35": "有効にすると、JSONオブジェクトがOpenAIのリクエストボディを上書きします。同名のフィールドはこの設定が優先されます。OpenAIのservice_tierはauto、default、flex、scale、priorityをサポートしています。",
|
||||
"e14c41ef2b7253c9": "総リクエスト Token: {0}",
|
||||
"e406825e0a72d2c2": "ローカル設定",
|
||||
"e4343921c928a856": "ログインに失敗しました",
|
||||
"e4c0daa3c4bea691": "Cursor コントロールプレーンアカウント機能への @aike0210 の貢献に感謝します。",
|
||||
"e53580f8031f13c0": "ブラウザでログインを完了し、Cursor に戻ってプラグインマーケットを開き直してください",
|
||||
"e552c2accdbf5178": "モデルを追加",
|
||||
"e6faccfddce722e8": "キャッシュ読込 Token: {0}",
|
||||
"e8a0a6053998ebfa": "ログイン済み",
|
||||
"eaffd48cd2ea9f1a": "例: https://api.anthropic.com",
|
||||
"eb1be07f2ca6e506": "Claude Opus 4.7の価格に基づいて見積もられます。",
|
||||
"ec3b17a75db49e24": "{0} t/s | 初回 Token {1}",
|
||||
@@ -201,7 +217,6 @@
|
||||
"f0b6a23368dd47cc": "モデルIDを直接入力するか、サーバーから返された一覧から選択します。",
|
||||
"f1aa7326f38b4c09": "ドラッグして並べ替え",
|
||||
"f1e0fc261d42fe29": "モデル一覧にホバーしたときに表示されるメモです。",
|
||||
"f363622480699c52": "推論強度は reasoning_effort をサポートする一部のモデルでのみ有効です。すべてのモデルが対応しているわけではありません。値が高いほど安定しやすい反面、遅くなることがあります。",
|
||||
"f3a76d896853c1df": "ミス",
|
||||
"f3fae6cccb9004b1": "カスタムヘッダー名は空にできません",
|
||||
"f474a4108aba4c4c": "サービスを停止",
|
||||
|
||||
@@ -19,6 +19,7 @@
|
||||
"1afed6a81a2512d2": "Выберите модель",
|
||||
"1baddde657dd2720": "Исходящие запросы используют системный прокси",
|
||||
"1bc77f5ab979f4c1": "Добавить настройки модели",
|
||||
"1c631615c1d85c9e": "Войти в Cursor",
|
||||
"1e238093b79b3165": "Если оставить пустым, используется 65536",
|
||||
"21296ab18ad9af25": "Дополнительные параметры JSON",
|
||||
"24343a2096988d42": "Не удалось открыть",
|
||||
@@ -39,9 +40,11 @@
|
||||
"37d23612f78a2e63": "Перезапустить и обновить",
|
||||
"392d0dceb45998d3": "Очень высокая",
|
||||
"393df9bb13ea4900": "Попадание",
|
||||
"3ab8cc15939f3b5c": "Выйти",
|
||||
"3af7e5489e61ea51": "Обновление",
|
||||
"3c2a9f9901109e75": "Тип {0} поддерживает только OpenAI или Anthropic",
|
||||
"3d13868593ae4eeb": "Язык интерфейса",
|
||||
"3d52574ce1500561": "Не подключено",
|
||||
"3ea83f9f55062582": "Дата выпуска: {0}",
|
||||
"3edda85621fd03b2": "адаптеров моделей",
|
||||
"3fd47edce45b3603": "Закрыть",
|
||||
@@ -56,10 +59,12 @@
|
||||
"4c0a929bb86ce912": "Сейчас: {0}",
|
||||
"4d2b6e53be6002e5": "Стратегия статистики кеша: {0} ({1})",
|
||||
"4d8c1c5b42830791": "Неизвестно",
|
||||
"4e30d7c9ed2b0eee": "Не задано",
|
||||
"4f0982ba1d37e51b": "Исходящие запросы используют прокси из переменных окружения",
|
||||
"5205125c0e91d346": "Максимальное число токенов, которое модель Anthropic может сгенерировать за один ответ. Оставьте поле пустым для значения по умолчанию.",
|
||||
"54e6745ff43c9c74": "Не удалось изменить порядок",
|
||||
"56627c94a9decee6": "Макс. выходных токенов",
|
||||
"58c6b0935a7216da": "Не удалось открыть профиль участника",
|
||||
"593a972852ba0004": "Cursor Assistant | Всегда бесплатно | Пользовательский API",
|
||||
"59a2195a01a8b35b": "{0} должен быть допустимым объектом JSON",
|
||||
"5aa8f5590c940829": "Ввод без кеша: {0}",
|
||||
@@ -78,6 +83,7 @@
|
||||
"66af574b8948fe83": "Ключ API {0} не может быть пустым",
|
||||
"6744b4c6a9aa0038": "Выключено",
|
||||
"675109292da4eb36": "Еще не проверено",
|
||||
"688102a402ba015a": "Ожидание входа...",
|
||||
"6a7b96f399e58138": "например, sk-xxxxxx",
|
||||
"6aa8f49cc992dfd7": "Проверить",
|
||||
"6ae23d6d7cb18592": "Ошибка сервиса",
|
||||
@@ -101,11 +107,13 @@
|
||||
"8139cb3dd11f5a67": "Если включено, объект JSON переопределит итоговые заголовки запроса. При совпадении имен используются значения отсюда; все значения должны быть строками.",
|
||||
"8151e8704a7ca89e": "Совпадений нет",
|
||||
"83913e71fcf7ff60": "Обновление выполнено",
|
||||
"83be9cac28873059": "Аккаунт управляющего уровня Cursor",
|
||||
"83fcfb4c1f2c1641": "Получить модели",
|
||||
"8672864e90417138": "Максимальная",
|
||||
"86df7ec743047234": "Сервис запущен",
|
||||
"899add6275682210": "Если оставить пустым, используется 200000",
|
||||
"8a4ef3e48e4e8a5a": "Включено",
|
||||
"8b8428f714611458": "Выбирайте интенсивность рассуждений только для моделей с поддержкой reasoning_effort. Если выбрать «Не задано», этот параметр не будет добавлен в запрос. Более высокие значения обычно дают более стабильный результат, но могут замедлить ответ.",
|
||||
"8c1935935600e336": "Проверка модели",
|
||||
"8cbcf741e727dbf7": "Настройки модели",
|
||||
"8d1de152be6360ce": "Доля успешных: {0}",
|
||||
@@ -159,12 +167,15 @@
|
||||
"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": "Папка настроек",
|
||||
"c8a52b66651d294c": "Не удалось выйти",
|
||||
"c8c14507b2d37395": "Интенсивность рассуждений",
|
||||
"c98e118e0a43f078": "Модель",
|
||||
"c9dd59beefd7144f": "Чтение кеша / (Чтение кеша + Ввод без кеша)",
|
||||
@@ -172,6 +183,7 @@
|
||||
"ca1d1059408b3837": "Ошибочных ходов: {0}",
|
||||
"cc5049729a2c10f1": "Тест не пройден. Проверьте исходные сведения.",
|
||||
"cd7ca5fb221e1c53": "{0} не может быть пустым",
|
||||
"cfa6c803eb3fc713": "Ожидание входа в браузере",
|
||||
"d0325067fed88e5a": "Доля попаданий в кеш: {0}",
|
||||
"d20ab96566d33f25": "Отображаемое имя {0} не может быть пустым",
|
||||
"d2243e1d44b2a94e": "Изменить настройки модели",
|
||||
@@ -179,20 +191,24 @@
|
||||
"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": "Настроено",
|
||||
"e01c5dae36cf8c35": "Если включено, объект JSON переопределит тело запроса OpenAI. При совпадении полей используются значения отсюда. OpenAI service_tier поддерживает auto, default, flex, scale и priority.",
|
||||
"e14c41ef2b7253c9": "Всего токенов запроса: {0}",
|
||||
"e406825e0a72d2c2": "Локальные настройки",
|
||||
"e4343921c928a856": "Не удалось войти",
|
||||
"e4c0daa3c4bea691": "Спасибо @aike0210 за вклад в функцию аккаунта панели управления Cursor.",
|
||||
"e53580f8031f13c0": "Завершите вход в браузере, затем вернитесь в Cursor и снова откройте магазин плагинов",
|
||||
"e552c2accdbf5178": "Добавить модель",
|
||||
"e6faccfddce722e8": "Токены чтения из кеша: {0}",
|
||||
"e8a0a6053998ebfa": "Выполнен вход",
|
||||
"eaffd48cd2ea9f1a": "например, https://api.anthropic.com",
|
||||
"eb1be07f2ca6e506": "Расчет основан на тарифах Claude Opus 4.7.",
|
||||
"ec3b17a75db49e24": "{0} т/с | Первый токен {1}",
|
||||
@@ -201,7 +217,6 @@
|
||||
"f0b6a23368dd47cc": "Введите идентификатор модели вручную или выберите его из списка, полученного от сервера.",
|
||||
"f1aa7326f38b4c09": "Перетащите, чтобы изменить порядок",
|
||||
"f1e0fc261d42fe29": "Примечание, отображаемое при наведении на модель в списке.",
|
||||
"f363622480699c52": "Интенсивность рассуждений применяется только к моделям с поддержкой reasoning_effort. Чем выше значение, тем обычно стабильнее результат, но ответ может формироваться медленнее.",
|
||||
"f3a76d896853c1df": "Промах",
|
||||
"f3fae6cccb9004b1": "Имя пользовательского заголовка не может быть пустым",
|
||||
"f474a4108aba4c4c": "Остановить сервис",
|
||||
|
||||
@@ -19,6 +19,7 @@
|
||||
"1afed6a81a2512d2": "选择模型",
|
||||
"1baddde657dd2720": "当前出站请求使用系统代理",
|
||||
"1bc77f5ab979f4c1": "新增模型配置",
|
||||
"1c631615c1d85c9e": "登录 Cursor",
|
||||
"1e238093b79b3165": "留空时默认 65536",
|
||||
"21296ab18ad9af25": "额外参数 JSON",
|
||||
"24343a2096988d42": "打开失败",
|
||||
@@ -39,9 +40,11 @@
|
||||
"37d23612f78a2e63": "立即重启更新",
|
||||
"392d0dceb45998d3": "极高",
|
||||
"393df9bb13ea4900": "命中",
|
||||
"3ab8cc15939f3b5c": "退出登录",
|
||||
"3af7e5489e61ea51": "刷新中",
|
||||
"3c2a9f9901109e75": "{0} 的类型仅支持 OpenAI 或 Anthropic",
|
||||
"3d13868593ae4eeb": "界面语言",
|
||||
"3d52574ce1500561": "未连接",
|
||||
"3ea83f9f55062582": "发布时间:{0}",
|
||||
"3edda85621fd03b2": "个模型适配器",
|
||||
"3fd47edce45b3603": "关闭",
|
||||
@@ -56,10 +59,12 @@
|
||||
"4c0a929bb86ce912": "当前:{0}",
|
||||
"4d2b6e53be6002e5": "缓存统计策略:{0}({1})",
|
||||
"4d8c1c5b42830791": "未知",
|
||||
"4e30d7c9ed2b0eee": "不设置",
|
||||
"4f0982ba1d37e51b": "当前出站请求使用环境变量代理",
|
||||
"5205125c0e91d346": "Anthropic 模型单次回复允许生成的最大 Token 数。留空时使用默认值。",
|
||||
"54e6745ff43c9c74": "排序失败",
|
||||
"56627c94a9decee6": "最大输出 Token",
|
||||
"58c6b0935a7216da": "打开贡献者主页失败",
|
||||
"593a972852ba0004": "Cursor助手|永久免费|自定义API",
|
||||
"59a2195a01a8b35b": "{0}必须是合法 JSON 对象",
|
||||
"5aa8f5590c940829": "非缓存输入:{0}",
|
||||
@@ -78,6 +83,7 @@
|
||||
"66af574b8948fe83": "{0} 的访问密钥不能为空",
|
||||
"6744b4c6a9aa0038": "已关闭",
|
||||
"675109292da4eb36": "尚未测试",
|
||||
"688102a402ba015a": "等待登录...",
|
||||
"6a7b96f399e58138": "例如:sk-xxxxxx",
|
||||
"6aa8f49cc992dfd7": "测试",
|
||||
"6ae23d6d7cb18592": "服务错误",
|
||||
@@ -101,11 +107,13 @@
|
||||
"8139cb3dd11f5a67": "开启后会把 JSON 对象覆盖到最终请求头。同名请求头以这里为准,值必须是字符串。",
|
||||
"8151e8704a7ca89e": "没有匹配项",
|
||||
"83913e71fcf7ff60": "刷新成功",
|
||||
"83be9cac28873059": "Cursor 控制面账号",
|
||||
"83fcfb4c1f2c1641": "获取模型",
|
||||
"8672864e90417138": "最高",
|
||||
"86df7ec743047234": "服务运行中",
|
||||
"899add6275682210": "留空时默认 200000",
|
||||
"8a4ef3e48e4e8a5a": "已开启",
|
||||
"8b8428f714611458": "仅当模型支持 reasoning_effort 时才选择推理强度;选择“不设置”后,请求不会携带该参数。越高通常越稳,但也可能更慢。",
|
||||
"8c1935935600e336": "模型测试",
|
||||
"8cbcf741e727dbf7": "模型配置",
|
||||
"8d1de152be6360ce": "有效占比:{0}",
|
||||
@@ -159,12 +167,15 @@
|
||||
"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": "设置文件夹",
|
||||
"c8a52b66651d294c": "退出登录失败",
|
||||
"c8c14507b2d37395": "推理强度",
|
||||
"c98e118e0a43f078": "模型",
|
||||
"c9dd59beefd7144f": "缓存读取 /(缓存读取 + 非缓存输入)",
|
||||
@@ -172,6 +183,7 @@
|
||||
"ca1d1059408b3837": "异常轮次:{0}",
|
||||
"cc5049729a2c10f1": "测试失败,请查看原始信息",
|
||||
"cd7ca5fb221e1c53": "{0}不能为空",
|
||||
"cfa6c803eb3fc713": "等待浏览器登录",
|
||||
"d0325067fed88e5a": "缓存命中率 {0}",
|
||||
"d20ab96566d33f25": "{0} 的显示名称不能为空",
|
||||
"d2243e1d44b2a94e": "编辑模型配置",
|
||||
@@ -179,20 +191,24 @@
|
||||
"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": "已配置",
|
||||
"e01c5dae36cf8c35": "开启后会把 JSON 对象覆盖到 OpenAI 请求体。同名字段以这里为准。OpenAI service_tier 支持 auto、default、flex、scale、priority。",
|
||||
"e14c41ef2b7253c9": "总请求:{0}",
|
||||
"e406825e0a72d2c2": "本地配置",
|
||||
"e4343921c928a856": "登录失败",
|
||||
"e4c0daa3c4bea691": "感谢 @aike0210 对 Cursor 控制面账号功能的贡献。",
|
||||
"e53580f8031f13c0": "请在浏览器完成登录,完成后返回 Cursor 重新打开插件市场",
|
||||
"e552c2accdbf5178": "新增模型",
|
||||
"e6faccfddce722e8": "缓存读取:{0}",
|
||||
"e8a0a6053998ebfa": "已经登录",
|
||||
"eaffd48cd2ea9f1a": "例如:https://api.anthropic.com",
|
||||
"eb1be07f2ca6e506": "按 Claude Opus 4.7 价格估算。",
|
||||
"ec3b17a75db49e24": "{0} t/s | 首字 {1}",
|
||||
@@ -201,7 +217,6 @@
|
||||
"f0b6a23368dd47cc": "可以直接输入模型标识,或从服务端返回的列表中选择。",
|
||||
"f1aa7326f38b4c09": "拖拽排序",
|
||||
"f1e0fc261d42fe29": "模型列表 hover 时显示的备注说明。",
|
||||
"f363622480699c52": "推理强度仅对部分支持 reasoning_effort 的模型生效,并不是所有模型都支持。越高通常越稳,但也可能更慢。",
|
||||
"f3a76d896853c1df": "未命中",
|
||||
"f3fae6cccb9004b1": "自定义请求头名称不能为空",
|
||||
"f474a4108aba4c4c": "关闭服务",
|
||||
|
||||
@@ -1,7 +1,10 @@
|
||||
import {
|
||||
DisconnectCursorAccount,
|
||||
GetCursorAccountStatus,
|
||||
GetState,
|
||||
LoadUserConfig,
|
||||
SaveUserConfig,
|
||||
StartCursorAccountLogin,
|
||||
StartProxy,
|
||||
StopProxy,
|
||||
} from "@bindings/cursor/internal/bridge/proxyservice.js";
|
||||
@@ -60,6 +63,18 @@ export function saveUserConfig(payload) {
|
||||
return withApiLogging("SaveUserConfig", payload, () => SaveUserConfig(payload));
|
||||
}
|
||||
|
||||
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());
|
||||
}
|
||||
|
||||
@@ -18,11 +18,14 @@ 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";
|
||||
@@ -165,7 +168,7 @@ export function buildModelAdapterTestRequestHash(source) {
|
||||
normalizeBaseURL(adapter.baseURL),
|
||||
asString(adapter.apiKey),
|
||||
asString(adapter.modelID),
|
||||
adapter.type === "openai" ? asString(adapter.reasoningEffort || "medium") : "",
|
||||
adapter.type === "openai" ? asString(adapter.reasoningEffort) : "",
|
||||
adapter.type === "openai" ? normalizeOpenAIEndpoint(adapter.openAIEndpoint) : "",
|
||||
adapter.type === "openai" ? String(Boolean(adapter.openAIExtraParamsEnabled)) : "false",
|
||||
adapter.type === "openai" && adapter.openAIExtraParamsEnabled ? asString(adapter.openAIExtraParamsJSON) : "",
|
||||
@@ -254,7 +257,7 @@ export function createEmptyModelAdapter() {
|
||||
apiKey: "",
|
||||
tooltipData: "备注",
|
||||
modelID: "",
|
||||
reasoningEffort: "medium",
|
||||
reasoningEffort: "",
|
||||
openAIEndpoint: OPENAI_ENDPOINT_RESPONSES,
|
||||
openAIExtraParamsEnabled: false,
|
||||
openAIExtraParamsJSON: OPENAI_EXTRA_PARAMS_DEFAULT_JSON,
|
||||
@@ -329,7 +332,7 @@ function validateAnthropicExtraParamsJSON(value) {
|
||||
export function normalizeModelAdapter(source) {
|
||||
const raw = source && typeof source === "object" ? source : {};
|
||||
const normalizedType = asString(raw.type).toLowerCase();
|
||||
const normalizedReasoningEffort = asString(raw.reasoningEffort || raw.reasoning_effort).toLowerCase();
|
||||
const normalizedReasoningEffort = normalizeReasoningEffort(raw.reasoningEffort ?? raw.reasoning_effort);
|
||||
const normalizedAnthropicThinkingEffort = asString(
|
||||
raw.anthropicThinkingEffort
|
||||
?? raw.anthropic_thinking_effort
|
||||
@@ -362,9 +365,7 @@ export function normalizeModelAdapter(source) {
|
||||
apiKey: asString(raw.apiKey || raw.key),
|
||||
tooltipData: asString(raw.tooltipData),
|
||||
modelID: asString(raw.modelID),
|
||||
reasoningEffort: SUPPORTED_REASONING_EFFORTS.has(normalizedReasoningEffort)
|
||||
? normalizedReasoningEffort
|
||||
: "medium",
|
||||
reasoningEffort: normalizedReasoningEffort,
|
||||
openAIEndpoint: normalizedType === "openai" ? normalizedOpenAIEndpoint : "",
|
||||
openAIExtraParamsEnabled,
|
||||
openAIExtraParamsJSON,
|
||||
@@ -442,7 +443,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 必须为正整数`;
|
||||
@@ -1086,6 +1087,10 @@ 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) {
|
||||
|
||||
@@ -0,0 +1,14 @@
|
||||
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 "";
|
||||
}
|
||||
@@ -0,0 +1,20 @@
|
||||
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);
|
||||
});
|
||||
@@ -2,6 +2,7 @@
|
||||
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 {
|
||||
@@ -173,6 +174,8 @@ onBeforeUnmount(() => {
|
||||
</div>
|
||||
</Card>
|
||||
|
||||
<CursorAccountCard />
|
||||
|
||||
<Card class="">
|
||||
<div class="flex items-center justify-between gap-4">
|
||||
<div>
|
||||
|
||||
+17
-19
@@ -70,20 +70,21 @@ func Run(resources EmbeddedResources) error {
|
||||
logger.Init()
|
||||
netproxy.InstallDefaultTransport()
|
||||
|
||||
embeddedCACertPEM := certs.EmbeddedCACertPEM()
|
||||
logEmbeddedCAInfo(embeddedCACertPEM)
|
||||
|
||||
certManager, err := certs.NewEmbeddedManager()
|
||||
if err := appdata.EnsureAssistantHome(); err != nil {
|
||||
return err
|
||||
}
|
||||
certManager, caCertPEM, err := certs.LoadOrCreateManager(appdata.CACertFilePath(), appdata.CAKeyFilePath())
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
logCAInfo(caCertPEM)
|
||||
|
||||
defaultBackendBaseURL := browserReachableLoopbackBaseURL(serverconfig.DefaultBackendListenAddr)
|
||||
defaultBackendBaseURL := "http://" + serverconfig.DefaultBackendListenAddr
|
||||
proxyServer, err := mitm.NewProxyServer(serverconfig.DefaultProxyListenAddr, defaultBackendBaseURL, "", "", certManager)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
proxyService := bridge.NewProxyService(proxyServer, certManager, embeddedCACertPEM)
|
||||
proxyService := bridge.NewProxyService(proxyServer, certManager, caCertPEM)
|
||||
adAssetBaseURL := defaultBackendBaseURL
|
||||
if cfg, err := proxyService.LoadUserConfig(); err == nil {
|
||||
adAssetBaseURL = browserReachableLoopbackBaseURL(cfg.BackendListenAddr)
|
||||
@@ -433,32 +434,29 @@ func windowsAdditionalBrowserArgs() []string {
|
||||
func browserReachableLoopbackBaseURL(listenAddr string) string {
|
||||
host, port, err := net.SplitHostPort(strings.TrimSpace(listenAddr))
|
||||
if err != nil || strings.TrimSpace(port) == "" {
|
||||
return "https://localhost:8000"
|
||||
return "http://" + serverconfig.DefaultBackendListenAddr
|
||||
}
|
||||
host = strings.TrimSpace(host)
|
||||
if host == "" || host == "0.0.0.0" || host == "::" || host == "[::]" {
|
||||
host = "127.0.0.1"
|
||||
}
|
||||
if host == "127.0.0.1" || host == "::1" || host == "localhost" {
|
||||
host = "localhost"
|
||||
}
|
||||
return "https://" + net.JoinHostPort(host, port)
|
||||
return "http://" + net.JoinHostPort(host, port)
|
||||
}
|
||||
|
||||
// logEmbeddedCAInfo 用于处理与 logEmbeddedCAInfo 相关的逻辑。
|
||||
func logEmbeddedCAInfo(certPEM []byte) {
|
||||
// logCAInfo 记录当前安装专属 CA 的公开信息。
|
||||
func logCAInfo(certPEM []byte) {
|
||||
if len(certPEM) == 0 {
|
||||
logger.Errorf("embedded CA is empty")
|
||||
logger.Errorf("installation CA is empty")
|
||||
return
|
||||
}
|
||||
cert, err := parseEmbeddedCert(certPEM)
|
||||
cert, err := parseCert(certPEM)
|
||||
if err != nil {
|
||||
logger.Errorf("parse embedded CA failed: %v", err)
|
||||
logger.Errorf("parse installation CA failed: %v", err)
|
||||
return
|
||||
}
|
||||
sum := sha256.Sum256(cert.Raw)
|
||||
logger.Infof(
|
||||
"embedded CA loaded: sha256=%s subject=%s valid=%s~%s",
|
||||
"installation CA loaded: sha256=%s subject=%s valid=%s~%s",
|
||||
strings.ToUpper(hex.EncodeToString(sum[:])),
|
||||
cert.Subject.String(),
|
||||
cert.NotBefore.Format(time.RFC3339),
|
||||
@@ -466,8 +464,8 @@ func logEmbeddedCAInfo(certPEM []byte) {
|
||||
)
|
||||
}
|
||||
|
||||
// parseEmbeddedCert 用于处理与 parseEmbeddedCert 相关的逻辑。
|
||||
func parseEmbeddedCert(data []byte) (*x509.Certificate, error) {
|
||||
// parseCert 解析 DER 或 PEM 编码的证书。
|
||||
func parseCert(data []byte) (*x509.Certificate, error) {
|
||||
if block, _ := pem.Decode(data); block != nil {
|
||||
return x509.ParseCertificate(block.Bytes)
|
||||
}
|
||||
|
||||
@@ -70,3 +70,8 @@ func LogsRootPath() string {
|
||||
func CACertFilePath() string {
|
||||
return filepath.Join(DataRootPath(), "ca.crt")
|
||||
}
|
||||
|
||||
// CAKeyFilePath 返回仅供本安装使用的 CA 私钥路径。
|
||||
func CAKeyFilePath() string {
|
||||
return filepath.Join(DataRootPath(), "ca.key")
|
||||
}
|
||||
|
||||
@@ -97,6 +97,7 @@ 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/`
|
||||
@@ -104,7 +105,8 @@ internal/backend/
|
||||
约定:
|
||||
|
||||
- `config.yaml` 是用户配置
|
||||
- `data/ca.crt` 是注入给宿主的 CA 证书
|
||||
- `data/ca.crt` 是首次运行时为当前用户生成、注入给宿主的 CA 证书
|
||||
- `data/ca.key` 是与该证书配套的本地私钥,权限固定为 `0600`,不得打包或提交到仓库
|
||||
- `data/ads/` 是广告包与资源缓存目录
|
||||
- `history/` 是会话事实与全局 usage JSON 目录,不属于日志
|
||||
- `logs/` 只保留必要文本运行日志
|
||||
|
||||
@@ -2,8 +2,15 @@
|
||||
package execbridge
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"crypto/sha256"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"image"
|
||||
_ "image/gif"
|
||||
_ "image/jpeg"
|
||||
_ "image/png"
|
||||
"net/http"
|
||||
"strings"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
@@ -31,10 +38,18 @@ 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
|
||||
@@ -145,6 +160,9 @@ 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":
|
||||
@@ -2285,6 +2303,64 @@ 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 {
|
||||
@@ -2318,7 +2394,9 @@ func convertReadResultToReadToolResult(result *agentv1.ReadResult) *agentv1.Read
|
||||
if content != "" {
|
||||
toolSuccess.Output = &agentv1.ReadToolSuccess_Content{Content: content}
|
||||
} else if len(data) > 0 {
|
||||
if len(data) > readReplayBinaryLimit {
|
||||
if imageOutput := readImageDataBlobOutput(data); imageOutput != nil {
|
||||
toolSuccess.Output = imageOutput
|
||||
} else if len(data) > readReplayBinaryLimit {
|
||||
toolSuccess.ExceededLimit = true
|
||||
toolSuccess.Output = &agentv1.ReadToolSuccess_Content{
|
||||
Content: replayTruncationNotice("Read binary data", readReplayBinaryLimit, 0, len(data)),
|
||||
|
||||
@@ -0,0 +1,125 @@
|
||||
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()
|
||||
}
|
||||
@@ -1126,7 +1126,16 @@ func isAnthropicCacheableBlock(block map[string]any) bool {
|
||||
case contentPartTypeText:
|
||||
return strings.TrimSpace(anthropicStringField(block, "text")) != ""
|
||||
case "tool_result":
|
||||
return strings.TrimSpace(anthropicStringField(block, "content")) != ""
|
||||
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
|
||||
}
|
||||
case "tool_use":
|
||||
return strings.TrimSpace(anthropicStringField(block, "id")) != "" && strings.TrimSpace(anthropicStringField(block, "name")) != ""
|
||||
default:
|
||||
@@ -1178,10 +1187,18 @@ 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": message.Content,
|
||||
"content": content,
|
||||
})
|
||||
case "user", "assistant":
|
||||
flushToolResults()
|
||||
|
||||
@@ -49,7 +49,8 @@ type openAIResponsesRequestBody struct {
|
||||
}
|
||||
|
||||
type openAIResponsesReasoning struct {
|
||||
Effort string `json:"effort,omitempty"`
|
||||
Effort string `json:"effort,omitempty"`
|
||||
Summary string `json:"summary,omitempty"`
|
||||
}
|
||||
|
||||
type openAIToolAccumulator struct {
|
||||
@@ -944,7 +945,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}
|
||||
requestBody.Reasoning = &openAIResponsesReasoning{Effort: effort, Summary: "auto"}
|
||||
requestBody.Include = []string{"reasoning.encrypted_content"}
|
||||
}
|
||||
body = requestBody
|
||||
@@ -1967,10 +1968,18 @@ func normalizeOpenAIResponsesInput(messages []Message) (string, []map[string]any
|
||||
}
|
||||
if role == "tool" && strings.TrimSpace(message.ToolCallID) != "" {
|
||||
callID := openAIResponsesToolMessageCallID(message, responsesCallIDs)
|
||||
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": openAIResponsesMessageText(message),
|
||||
"output": output,
|
||||
})
|
||||
activeAssistantReasoningKey = ""
|
||||
continue
|
||||
|
||||
@@ -2,12 +2,140 @@ 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")
|
||||
|
||||
@@ -1,10 +1,64 @@
|
||||
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{
|
||||
{
|
||||
|
||||
@@ -0,0 +1,86 @@
|
||||
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)
|
||||
}
|
||||
if len(items) != 1 || items[0]["type"] != "function_call_output" {
|
||||
t.Fatalf("openai responses items = %#v", items)
|
||||
}
|
||||
content, ok := items[0]["output"].([]map[string]any)
|
||||
if !ok || len(content) != 2 {
|
||||
t.Fatalf("openai responses output = %#v", items[0]["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"),
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
@@ -730,6 +730,9 @@ 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)
|
||||
}
|
||||
|
||||
@@ -19,119 +19,87 @@ 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) *aiHandler {
|
||||
handler := newAIHandlerMux()
|
||||
handler.Handle(
|
||||
func newAIHandler(service *Service) http.Handler {
|
||||
mux := http.NewServeMux()
|
||||
mux.Handle(
|
||||
dashboardServiceGetTokenUsageProcedure,
|
||||
connect.NewUnaryHandler(dashboardServiceGetTokenUsageProcedure, service.GetTokenUsage),
|
||||
)
|
||||
handler.Handle(
|
||||
mux.Handle(
|
||||
dashboardServiceGetGlassEarlyPreviewEnrollmentProcedure,
|
||||
connect.NewUnaryHandler(dashboardServiceGetGlassEarlyPreviewEnrollmentProcedure, service.GetGlassEarlyPreviewEnrollment),
|
||||
)
|
||||
handler.Handle(
|
||||
mux.Handle(
|
||||
aiserverv1connect.AiServiceCountTokensProcedure,
|
||||
connect.NewUnaryHandler(aiserverv1connect.AiServiceCountTokensProcedure, service.CountTokens),
|
||||
)
|
||||
handler.Handle(
|
||||
mux.Handle(
|
||||
aiserverv1connect.AiServiceGetThoughtAnnotationProcedure,
|
||||
connect.NewUnaryHandler(aiserverv1connect.AiServiceGetThoughtAnnotationProcedure, service.GetThoughtAnnotation),
|
||||
)
|
||||
handler.Handle(
|
||||
mux.Handle(
|
||||
aiserverv1connect.AiServiceWriteGitCommitMessageProcedure,
|
||||
connect.NewUnaryHandler(aiserverv1connect.AiServiceWriteGitCommitMessageProcedure, service.WriteGitCommitMessage),
|
||||
)
|
||||
handler.Handle(
|
||||
mux.Handle(
|
||||
aiserverv1connect.AiServiceCreateExperimentalIndexProcedure,
|
||||
connect.NewUnaryHandler(aiserverv1connect.AiServiceCreateExperimentalIndexProcedure, service.CreateExperimentalIndex),
|
||||
)
|
||||
handler.Handle(
|
||||
mux.Handle(
|
||||
aiserverv1connect.AiServiceListExperimentalIndexFilesProcedure,
|
||||
connect.NewUnaryHandler(aiserverv1connect.AiServiceListExperimentalIndexFilesProcedure, service.ListExperimentalIndexFiles),
|
||||
)
|
||||
handler.Handle(
|
||||
mux.Handle(
|
||||
aiserverv1connect.AiServiceListenExperimentalIndexProcedure,
|
||||
connect.NewServerStreamHandler(aiserverv1connect.AiServiceListenExperimentalIndexProcedure, service.ListenExperimentalIndex),
|
||||
)
|
||||
handler.Handle(
|
||||
mux.Handle(
|
||||
aiserverv1connect.AiServiceRegisterFileToIndexProcedure,
|
||||
connect.NewUnaryHandler(aiserverv1connect.AiServiceRegisterFileToIndexProcedure, service.RegisterFileToIndex),
|
||||
)
|
||||
handler.Handle(
|
||||
mux.Handle(
|
||||
aiserverv1connect.AiServiceSetupIndexDependenciesProcedure,
|
||||
connect.NewUnaryHandler(aiserverv1connect.AiServiceSetupIndexDependenciesProcedure, service.SetupIndexDependencies),
|
||||
)
|
||||
handler.Handle(
|
||||
mux.Handle(
|
||||
aiserverv1connect.AiServiceComputeIndexTopoSortProcedure,
|
||||
connect.NewUnaryHandler(aiserverv1connect.AiServiceComputeIndexTopoSortProcedure, service.ComputeIndexTopoSort),
|
||||
)
|
||||
handler.Handle(
|
||||
mux.Handle(
|
||||
aiserverv1connect.AiServiceDocumentationQueryProcedure,
|
||||
connect.NewUnaryHandler(aiserverv1connect.AiServiceDocumentationQueryProcedure, service.DocumentationQuery),
|
||||
)
|
||||
handler.Handle(
|
||||
mux.Handle(
|
||||
aiserverv1connect.AiServiceAvailableDocsProcedure,
|
||||
connect.NewUnaryHandler(aiserverv1connect.AiServiceAvailableDocsProcedure, service.AvailableDocs),
|
||||
)
|
||||
handler.Handle(
|
||||
mux.Handle(
|
||||
aiserverv1connect.AiServiceKnowledgeBaseAddProcedure,
|
||||
connect.NewUnaryHandler(aiserverv1connect.AiServiceKnowledgeBaseAddProcedure, service.KnowledgeBaseAdd),
|
||||
)
|
||||
handler.Handle(
|
||||
mux.Handle(
|
||||
aiserverv1connect.AiServiceKnowledgeBaseListProcedure,
|
||||
connect.NewUnaryHandler(aiserverv1connect.AiServiceKnowledgeBaseListProcedure, service.KnowledgeBaseList),
|
||||
)
|
||||
handler.Handle(
|
||||
mux.Handle(
|
||||
aiserverv1connect.AiServiceKnowledgeBaseRemoveProcedure,
|
||||
connect.NewUnaryHandler(aiserverv1connect.AiServiceKnowledgeBaseRemoveProcedure, service.KnowledgeBaseRemove),
|
||||
)
|
||||
handler.Handle(
|
||||
mux.Handle(
|
||||
aiserverv1connect.AiServiceKnowledgeBaseUpdateProcedure,
|
||||
connect.NewUnaryHandler(aiserverv1connect.AiServiceKnowledgeBaseUpdateProcedure, service.KnowledgeBaseUpdate),
|
||||
)
|
||||
handler.Handle(
|
||||
mux.Handle(
|
||||
aiserverv1connect.AiServiceFetchRelevantKnowledgeForConversationProcedure,
|
||||
connect.NewUnaryHandler(aiserverv1connect.AiServiceFetchRelevantKnowledgeForConversationProcedure, service.FetchRelevantKnowledgeForConversation),
|
||||
)
|
||||
return handler
|
||||
mux.Handle("/", http.NotFoundHandler())
|
||||
return mux
|
||||
}
|
||||
|
||||
func (service *Service) GetThoughtAnnotation(_ context.Context, req *connect.Request[aiserverv1.GetThoughtAnnotationRequest]) (*connect.Response[aiserverv1.GetThoughtAnnotationResponse], error) {
|
||||
|
||||
@@ -1,20 +0,0 @@
|
||||
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")
|
||||
}
|
||||
}
|
||||
@@ -19,15 +19,29 @@ type pendingCheckpointBlobWrite struct {
|
||||
blob CheckpointBlob
|
||||
}
|
||||
|
||||
func clonePendingTurnCompletion(completion *pendingTurnCompletion) *pendingTurnCompletion {
|
||||
func successfulCheckpointTerminalAction(completion *pendingTurnCompletion) checkpointTerminalAction {
|
||||
if completion == nil {
|
||||
return 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),
|
||||
}
|
||||
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
|
||||
}
|
||||
@@ -43,8 +57,8 @@ func (service *Service) queueCheckpointProjection(stream *ActiveStream, projecti
|
||||
if stream.ConfirmedCheckpointBlobs == nil {
|
||||
stream.ConfirmedCheckpointBlobs = make(map[string]struct{})
|
||||
}
|
||||
if completion == nil && stream.PendingCheckpoint != nil {
|
||||
completion = stream.PendingCheckpoint.Completion
|
||||
if terminal.Kind == checkpointTerminalActionNone && stream.PendingCheckpoint != nil {
|
||||
terminal = stream.PendingCheckpoint.Terminal
|
||||
}
|
||||
required := make(map[string]struct{}, len(projection.Blobs))
|
||||
pendingKeys := make(map[string]struct{}, len(stream.PendingCheckpointBlobWrites))
|
||||
@@ -74,11 +88,11 @@ func (service *Service) queueCheckpointProjection(stream *ActiveStream, projecti
|
||||
toWrite = append(toWrite, pendingCheckpointBlobWrite{requestID: requestID, blob: blob})
|
||||
}
|
||||
stream.PendingCheckpoint = &pendingCheckpointPublish{
|
||||
State: state,
|
||||
Required: required,
|
||||
Completion: clonePendingTurnCompletion(completion),
|
||||
State: state,
|
||||
Required: required,
|
||||
Terminal: terminal,
|
||||
}
|
||||
if completion != nil {
|
||||
if terminal.Kind != checkpointTerminalActionNone {
|
||||
stream.Phase = TurnPhaseCheckpointing
|
||||
}
|
||||
stream.UpdatedAt = time.Now().UTC()
|
||||
@@ -94,13 +108,8 @@ func (service *Service) queueCheckpointProjection(stream *ActiveStream, projecti
|
||||
if service.checkpointProjectionReady(stream) {
|
||||
return service.publishReadyCheckpoint(stream)
|
||||
}
|
||||
// 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))
|
||||
}
|
||||
}
|
||||
// Checkpoints reference these Blob IDs, so the client must confirm every
|
||||
// required Blob before the checkpoint becomes visible.
|
||||
service.scheduleStreamTimer(
|
||||
stream,
|
||||
providerTimerKey(streamTimerCheckpointBlobs, ""),
|
||||
@@ -113,31 +122,6 @@ func (service *Service) queueCheckpointProjection(stream *ActiveStream, projecti
|
||||
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
|
||||
@@ -207,24 +191,18 @@ func (service *Service) publishReadyCheckpoint(stream *ActiveStream) error {
|
||||
}
|
||||
stream.PendingCheckpoint = nil
|
||||
state := pending.State
|
||||
completion := clonePendingTurnCompletion(pending.Completion)
|
||||
published := pending.Published
|
||||
terminal := pending.Terminal
|
||||
stream.UpdatedAt = time.Now().UTC()
|
||||
stream.mu.Unlock()
|
||||
clearStreamTimer(stream, providerTimerKey(streamTimerCheckpointBlobs, ""))
|
||||
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
|
||||
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)
|
||||
}
|
||||
return err
|
||||
}
|
||||
if completion != nil {
|
||||
return service.finishSuccessfulTurnAfterCheckpoint(stream, *completion)
|
||||
}
|
||||
return nil
|
||||
return service.finishCheckpointTerminalAction(stream, terminal)
|
||||
}
|
||||
|
||||
func (service *Service) handleCheckpointBlobTimeout(stream *ActiveStream) error {
|
||||
@@ -251,12 +229,23 @@ 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 && pending.Completion != nil {
|
||||
return service.finishSuccessfulTurnAfterCheckpoint(stream, *pending.Completion)
|
||||
if pending != nil {
|
||||
return service.finishCheckpointTerminalAction(stream, pending.Terminal)
|
||||
}
|
||||
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,22 +8,43 @@ import (
|
||||
"cursor/gen/agentv1"
|
||||
)
|
||||
|
||||
func TestCheckpointBlobSyncPublishesNonTerminalCheckpointBeforeAcknowledgements(t *testing.T) {
|
||||
func TestCheckpointBlobSyncWaitsForAcknowledgementsBeforePublishingNonTerminalCheckpoint(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)+1 {
|
||||
t.Fatalf("events before ACK = %d, want %d Blob writes and one checkpoint", len(events), len(projection.Blobs))
|
||||
if len(events) != len(projection.Blobs) {
|
||||
t.Fatalf("events before ACK = %d, want %d Blob writes", len(events), len(projection.Blobs))
|
||||
}
|
||||
for _, event := range events[:len(projection.Blobs)] {
|
||||
for _, event := range events {
|
||||
if event.Message.GetKvServerMessage().GetSetBlobArgs() == nil {
|
||||
t.Fatalf("event before ACK = %#v, want set_blob_args", event.Message)
|
||||
}
|
||||
}
|
||||
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)
|
||||
|
||||
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")
|
||||
}
|
||||
}
|
||||
|
||||
acknowledgeCheckpointBlobs(t, service, stream)
|
||||
@@ -98,7 +119,123 @@ func TestCheckpointBlobTimeoutDoesNotFailSuccessfulTurn(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestCancellationKeepsPublishedCheckpointAndIgnoresLateAcknowledgements(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) {
|
||||
service, stream, projection := testCheckpointBlobProjection(t)
|
||||
if err := service.queueCheckpointProjection(stream, projection, nil); err != nil {
|
||||
t.Fatalf("queueCheckpointProjection() error = %v", err)
|
||||
@@ -110,8 +247,8 @@ func TestCancellationKeepsPublishedCheckpointAndIgnoresLateAcknowledgements(t *t
|
||||
checkpointBeforeCancel++
|
||||
}
|
||||
}
|
||||
if checkpointBeforeCancel != 1 {
|
||||
t.Fatalf("checkpoints before cancel = %d, want 1", checkpointBeforeCancel)
|
||||
if checkpointBeforeCancel != 0 {
|
||||
t.Fatalf("checkpoints before cancel = %d, want 0", checkpointBeforeCancel)
|
||||
}
|
||||
stream.mu.Lock()
|
||||
requestIDs := make([]uint32, 0, len(stream.PendingCheckpointBlobWrites))
|
||||
@@ -149,7 +286,7 @@ func TestCancellationKeepsPublishedCheckpointAndIgnoresLateAcknowledgements(t *t
|
||||
stream.mu.Lock()
|
||||
pending := stream.PendingCheckpoint
|
||||
stream.mu.Unlock()
|
||||
if checkpointCount != 1 || !canceledEnd || pending != nil {
|
||||
if checkpointCount != 0 || !canceledEnd || pending != nil {
|
||||
t.Fatalf("cancel events checkpoints=%d canceled_end=%v pending=%v", checkpointCount, canceledEnd, pending != nil)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -234,7 +234,7 @@ func (service *Service) buildLegacyCompactionPlan(base *compactionPlan, conversa
|
||||
if conversation == nil || base == nil {
|
||||
return nil, nil
|
||||
}
|
||||
candidates := buildContextCompactionCandidates(checkpointProjectionEntries(conversation.Entries), base.CurrentTurnSeq, base.CurrentRequestID)
|
||||
candidates := buildContextCompactionCandidates(replayablePromptProjectionEntries(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(checkpointProjectionEntries(conversation.Entries), base.CurrentTurnSeq, base.CurrentRequestID)
|
||||
currentCandidate, hasCurrentCandidate := buildCurrentTurnCompactionCandidate(replayablePromptProjectionEntries(conversation.Entries), base.CurrentTurnSeq, base.CurrentRequestID)
|
||||
if !hasCurrentCandidate {
|
||||
return legacyPlan, nil
|
||||
}
|
||||
@@ -447,16 +447,8 @@ func (service *Service) handleCompactionEvent(stream *ActiveStream, payload *str
|
||||
if err := service.completeManualCompactionTurn(stream); err != nil {
|
||||
return service.failStream(stream, "unknown", err)
|
||||
}
|
||||
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
|
||||
completion := manualCompactionTurnCompletion(stream)
|
||||
return service.publishCheckpointWithCompletion(stream.RequestID, stream.ConversationID, &completion)
|
||||
}
|
||||
return service.requestProviderAction(stream, providerActionResume)
|
||||
}
|
||||
@@ -500,12 +492,8 @@ func (service *Service) finishManualCompactionNoop(stream *ActiveStream) error {
|
||||
if err := service.completeManualCompactionTurn(stream); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := service.broker.Publish(stream.RequestID, StreamEvent{
|
||||
Message: buildTurnEndedMessage(0, 0, 0, 0),
|
||||
}); err != nil {
|
||||
return err
|
||||
}
|
||||
return service.broker.Complete(stream.RequestID, "", "")
|
||||
completion := manualCompactionTurnCompletion(stream)
|
||||
return service.publishCheckpointWithCompletion(stream.RequestID, stream.ConversationID, &completion)
|
||||
}
|
||||
|
||||
func (service *Service) completeManualCompactionTurn(stream *ActiveStream) error {
|
||||
@@ -530,10 +518,21 @@ 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
|
||||
@@ -568,6 +567,7 @@ 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
|
||||
}
|
||||
replacementEntries := append([]HistoryEntry(nil), candidateConversation.Entries...)
|
||||
compactionEntries := append([]HistoryEntry(nil), candidateConversation.Entries[originalEntryCount:]...)
|
||||
if service.store != nil {
|
||||
persisted, err := service.store.ReplaceEntries(conversationID, replacementEntries, func(item *ConversationFile) error {
|
||||
persisted, _, err := service.store.AppendEntriesWithUpdate(conversationID, resetEntrySequences(compactionEntries), func(item *ConversationFile) error {
|
||||
if item == nil {
|
||||
return nil
|
||||
}
|
||||
@@ -605,10 +605,7 @@ func (service *Service) applyCompactionPlan(stream *ActiveStream, conversationID
|
||||
if item == nil {
|
||||
return nil
|
||||
}
|
||||
item.Entries = nil
|
||||
item.NextEntrySeq = 1
|
||||
item.NextTurnSeq = 1
|
||||
appendEntriesInPlace(item, resetEntrySequences(replacementEntries))
|
||||
appendEntriesInPlace(item, resetEntrySequences(compactionEntries))
|
||||
item.TokenDetailsUsedTokens = 0
|
||||
clearConversationAutoCompactionState(item)
|
||||
return nil
|
||||
@@ -643,14 +640,13 @@ func applyCompactionToConversation(conversation *ConversationFile, plan *Pending
|
||||
if conversation == nil || plan == nil {
|
||||
return nil
|
||||
}
|
||||
replacementEntries, err := buildCompactedContextEntries(conversation, plan, summaryText)
|
||||
compactionEntries, err := buildCompactedContextEntries(conversation, plan, summaryText)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
conversation.Entries = nil
|
||||
conversation.NextEntrySeq = 1
|
||||
conversation.NextTurnSeq = 1
|
||||
appendEntriesInPlace(conversation, resetEntrySequences(replacementEntries))
|
||||
// Canonical history stays append-only. The prompt projector applies the
|
||||
// latest summary marker when constructing model-visible replay.
|
||||
appendEntriesInPlace(conversation, resetEntrySequences(compactionEntries))
|
||||
conversation.TokenDetailsUsedTokens = 0
|
||||
clearConversationAutoCompactionState(conversation)
|
||||
if conversation.TokenDetailsMaxTokens == 0 {
|
||||
@@ -671,40 +667,9 @@ 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),
|
||||
|
||||
@@ -0,0 +1,206 @@
|
||||
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{}
|
||||
@@ -20,16 +20,21 @@ type DefaultPromptCompiler struct {
|
||||
catalog ToolCatalog
|
||||
reminders ReminderInjector
|
||||
rules *UserRuleStore
|
||||
blobs contentBlobReader
|
||||
}
|
||||
|
||||
// NewPromptCompiler 创建默认 prompt 编译器。
|
||||
func NewPromptCompiler(projector *HistoryProjector, catalog ToolCatalog, reminders ReminderInjector, rules *UserRuleStore) *DefaultPromptCompiler {
|
||||
return &DefaultPromptCompiler{
|
||||
func NewPromptCompiler(projector *HistoryProjector, catalog ToolCatalog, reminders ReminderInjector, rules *UserRuleStore, blobReaders ...contentBlobReader) *DefaultPromptCompiler {
|
||||
compiler := &DefaultPromptCompiler{
|
||||
projector: projector,
|
||||
catalog: catalog,
|
||||
reminders: reminders,
|
||||
rules: rules,
|
||||
}
|
||||
if len(blobReaders) > 0 {
|
||||
compiler.blobs = blobReaders[0]
|
||||
}
|
||||
return compiler
|
||||
}
|
||||
|
||||
// Compile 生成当前 turn 应发送给 provider 的消息和工具集合。
|
||||
@@ -86,6 +91,10 @@ 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,
|
||||
|
||||
@@ -0,0 +1,111 @@
|
||||
// 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
|
||||
}
|
||||
@@ -0,0 +1,41 @@
|
||||
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")
|
||||
}
|
||||
}
|
||||
@@ -121,10 +121,15 @@ 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 {
|
||||
if len(entries) == 0 && update == nil {
|
||||
conversation, err := store.LoadConversation(conversationID)
|
||||
return conversation, nil, err
|
||||
}
|
||||
@@ -162,6 +167,11 @@ func (store *ConversationFileStore) AppendEntries(conversationID string, entries
|
||||
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
|
||||
@@ -762,6 +772,7 @@ 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)) {
|
||||
@@ -894,6 +905,7 @@ 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...)
|
||||
|
||||
@@ -0,0 +1,167 @@
|
||||
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
|
||||
}
|
||||
@@ -0,0 +1,109 @@
|
||||
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,6 +2,7 @@ package forwarder
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
@@ -125,6 +126,53 @@ 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)
|
||||
|
||||
@@ -32,11 +32,3 @@ 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)
|
||||
}
|
||||
|
||||
@@ -326,20 +326,24 @@ 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)
|
||||
for index, entry := range entries {
|
||||
if index < compactionIndex && isPromptReplayEntryKind(entry.Kind) {
|
||||
if _, ok := preservedIndexes[index]; !ok {
|
||||
continue
|
||||
}
|
||||
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])
|
||||
}
|
||||
if index < compactionIndex {
|
||||
if rewritten, ok := compactedProjectionPreservedEntry(entry); ok {
|
||||
entry = rewritten
|
||||
}
|
||||
}
|
||||
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
|
||||
}
|
||||
filtered = append(filtered, entry)
|
||||
}
|
||||
filtered = append(filtered, entries[compactionIndex+1:]...)
|
||||
return filtered
|
||||
}
|
||||
|
||||
@@ -575,7 +579,7 @@ func (projector *HistoryProjector) ProjectCheckpointProjection(conversation *Con
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
state.Turns = turnIDs
|
||||
state.Turns = append(cloneByteSlices(conversation.ImportedTurnIDs), turnIDs...)
|
||||
replayMessages, err := projector.ProjectPromptReplay(conversation)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -1292,7 +1296,7 @@ func filterCheckpointPersistentToolReplay(messages []promptengine.Message) []pro
|
||||
return filtered
|
||||
}
|
||||
|
||||
func restoreImportedReplayUserMessages(messages []promptengine.Message, importedTurns [][]byte) []promptengine.Message {
|
||||
func restoreImportedReplayUserMessages(messages []promptengine.Message, importedTurns [][]byte, blobs importedBlobStore) []promptengine.Message {
|
||||
if len(messages) == 0 || len(importedTurns) == 0 {
|
||||
return messages
|
||||
}
|
||||
@@ -1301,16 +1305,16 @@ func restoreImportedReplayUserMessages(messages []promptengine.Message, imported
|
||||
if len(rawTurn) == 0 {
|
||||
continue
|
||||
}
|
||||
turn := &agentv1.ConversationTurnStructure{}
|
||||
if err := proto.Unmarshal(rawTurn, turn); err != nil {
|
||||
turn, _, err := decodeImportedTurn(rawTurn, blobs)
|
||||
if err != nil || turn == nil {
|
||||
continue
|
||||
}
|
||||
agentTurn := turn.GetAgentConversationTurn()
|
||||
if agentTurn == nil || len(agentTurn.GetUserMessage()) == 0 {
|
||||
continue
|
||||
}
|
||||
userMessage := &agentv1.UserMessage{}
|
||||
if err := proto.Unmarshal(agentTurn.GetUserMessage(), userMessage); err != nil {
|
||||
userMessage, err := decodeImportedUserMessage(agentTurn.GetUserMessage(), blobs)
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
replay, ok := promptengine.BuildUserMessageReplayMessage(userMessage)
|
||||
|
||||
@@ -0,0 +1,180 @@
|
||||
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()
|
||||
}
|
||||
@@ -0,0 +1,174 @@
|
||||
// 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 ""
|
||||
}
|
||||
@@ -224,6 +224,7 @@ 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)
|
||||
@@ -269,10 +270,30 @@ 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)
|
||||
importedEntries, err = service.importConversationState(conversation, intent.ConversationState, intent.PreFetchedBlobs)
|
||||
if err != nil {
|
||||
return nil, agentv1.AgentMode_AGENT_MODE_AGENT, 0, nil, err
|
||||
}
|
||||
@@ -138,6 +138,7 @@ 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
|
||||
|
||||
@@ -248,6 +248,7 @@ func subagentModelOverrideSummaries(overrides map[string]runtimecore.SubagentMod
|
||||
|
||||
type Service struct {
|
||||
store *ConversationFileStore
|
||||
contentBlobs *ContentBlobStore
|
||||
usageStore *UsageFileStore
|
||||
codebaseIndexStore *CodebaseIndexStore
|
||||
docsIndexStore *DocsIndexStore
|
||||
@@ -274,6 +275,7 @@ 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
|
||||
@@ -287,12 +289,13 @@ 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),
|
||||
compiler: NewPromptCompiler(projector, NewToolCatalog(), NewReminderInjector(), rules, contentBlobs),
|
||||
provider: NewProviderGateway(resolver),
|
||||
resolver: resolver,
|
||||
modelMemory: modelMemory,
|
||||
@@ -317,6 +320,7 @@ 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,
|
||||
@@ -559,6 +563,7 @@ 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) {
|
||||
@@ -606,6 +611,7 @@ 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
|
||||
@@ -1012,6 +1018,9 @@ 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) {
|
||||
@@ -1050,6 +1059,21 @@ 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)
|
||||
@@ -2246,6 +2270,15 @@ 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
|
||||
@@ -2265,6 +2298,10 @@ 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)
|
||||
@@ -2282,7 +2319,7 @@ func (service *Service) publishCheckpointWithCompletion(requestID string, _ stri
|
||||
}
|
||||
projection.State.PendingToolCalls = buildPendingToolCalls(pendingExecs, pendingInteractions)
|
||||
service.rewriteCheckpointTokenDetailsForClient(stream, conversation, projection.State)
|
||||
return service.queueCheckpointProjection(stream, projection, completion)
|
||||
return service.queueCheckpointProjectionWithTerminal(stream, projection, terminal)
|
||||
}
|
||||
|
||||
func (service *Service) rewriteCheckpointTokenDetailsForClient(stream *ActiveStream, conversation *ConversationFile, state *agentv1.ConversationStateStructure) {
|
||||
@@ -2422,18 +2459,20 @@ func (service *Service) failActiveStream(stream *ActiveStream, conversationID st
|
||||
if cancel != nil {
|
||||
cancel()
|
||||
}
|
||||
service.setTurnPhase(stream, TurnPhaseFailed)
|
||||
var firstErr error
|
||||
if err := service.syncSummaryCarryForward(conversationID, requestID, modelCallID); err != nil && firstErr == nil {
|
||||
firstErr = err
|
||||
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,
|
||||
)
|
||||
}
|
||||
if err := service.publishCheckpoint(requestID, conversationID); 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.broker.Fail(requestID, terminalCode, terminalMessage); err != nil && firstErr == nil {
|
||||
firstErr = err
|
||||
}
|
||||
return firstErr
|
||||
return nil
|
||||
}
|
||||
|
||||
// buildRunEntries 构造一次 run intent 需要写入 history 的首批 entry。
|
||||
|
||||
@@ -45,13 +45,25 @@ func (snapshot turnUsageSnapshot) requestTokensTotal() int64 {
|
||||
return snapshot.promptTokensTotal() + nonNegativeInt64(snapshot.OutputTokens)
|
||||
}
|
||||
|
||||
func (service *Service) importConversationState(item *ConversationFile, state *agentv1.ConversationStateStructure) ([]HistoryEntry, error) {
|
||||
func (service *Service) importConversationState(item *ConversationFile, state *agentv1.ConversationStateStructure, prefetchedBlobs []*agentv1.PreFetchedBlob) ([]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 := importedConversationStateModelMessages(state); err != nil {
|
||||
if messages, err := importedConversationStateModelMessagesWithBlobs(state, blobs); err != nil {
|
||||
return nil, err
|
||||
} else {
|
||||
for _, message := range messages {
|
||||
@@ -105,6 +117,10 @@ 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
|
||||
}
|
||||
@@ -113,7 +129,7 @@ func importedConversationStateModelMessages(state *agentv1.ConversationStateStru
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("decode imported replay messages: %w", err)
|
||||
}
|
||||
decoded = restoreImportedReplayUserMessages(decoded, state.GetTurns())
|
||||
decoded = restoreImportedReplayUserMessages(decoded, state.GetTurns(), blobs)
|
||||
decoded = filterLegacyPlainWriteReplay(decoded)
|
||||
decoded = filterInternalPromptContextReplay(decoded)
|
||||
messages := make([]modeladapter.Message, 0, len(decoded))
|
||||
@@ -133,35 +149,18 @@ func importedConversationStateModelMessages(state *agentv1.ConversationStateStru
|
||||
if len(rawTurn) == 0 {
|
||||
continue
|
||||
}
|
||||
turn := &agentv1.ConversationTurnStructure{}
|
||||
if err := proto.Unmarshal(rawTurn, turn); err != nil {
|
||||
return nil, fmt.Errorf("decode imported turn: %w", err)
|
||||
turn, turnID, err := decodeImportedTurn(rawTurn, blobs)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
agentTurn := turn.GetAgentConversationTurn()
|
||||
if agentTurn == nil {
|
||||
continue
|
||||
if turn == nil && len(turnID) > 0 {
|
||||
return nil, fmt.Errorf("missing prefetched turn blob %x", turnID)
|
||||
}
|
||||
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))
|
||||
}
|
||||
turnMessages, err := importedBlobTurnMessages(turn, blobs)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
messages = append(messages, turnMessages...)
|
||||
}
|
||||
return normalizeReplayMessageSequence(messages), nil
|
||||
}
|
||||
|
||||
@@ -39,6 +39,7 @@ 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"`
|
||||
@@ -224,11 +225,25 @@ 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{}
|
||||
Completion *pendingTurnCompletion
|
||||
Published bool
|
||||
State *agentv1.ConversationStateStructure
|
||||
Required map[string]struct{}
|
||||
Terminal checkpointTerminalAction
|
||||
}
|
||||
|
||||
type PendingCompaction struct {
|
||||
@@ -430,6 +445,7 @@ 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
|
||||
|
||||
+264
-119
@@ -2,8 +2,6 @@ package backend
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/tls"
|
||||
"crypto/x509"
|
||||
"fmt"
|
||||
"net"
|
||||
"net/http"
|
||||
@@ -29,11 +27,11 @@ const healthPath = "/healthz"
|
||||
const tabServerBaseURL = "https://tab.leokun.cn"
|
||||
|
||||
type Host struct {
|
||||
store *serverconfig.Store
|
||||
listenAddr string
|
||||
configs *serverconfig.Manager
|
||||
healthHTTP *http.Client
|
||||
tlsCertificate *tls.Certificate
|
||||
store *serverconfig.Store
|
||||
listenAddr string
|
||||
configs *serverconfig.Manager
|
||||
healthHTTP *http.Client
|
||||
controlPlaneAuth upstream.AuthorizationProvider
|
||||
|
||||
runMu sync.RWMutex
|
||||
httpServer *http.Server
|
||||
@@ -43,20 +41,7 @@ type Host struct {
|
||||
mux http.Handler
|
||||
}
|
||||
|
||||
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) {
|
||||
func NewHost(store *serverconfig.Store, controlPlaneAuth upstream.AuthorizationProvider) (*Host, error) {
|
||||
if store == nil {
|
||||
return nil, fmt.Errorf("backend config store is required")
|
||||
}
|
||||
@@ -66,19 +51,12 @@ func NewHost(store *serverconfig.Store, options ...HostOption) (*Host, error) {
|
||||
}
|
||||
cfg := configs.Current()
|
||||
host := &Host{
|
||||
store: store,
|
||||
listenAddr: cfg.BackendListenAddr,
|
||||
configs: configs,
|
||||
store: store,
|
||||
listenAddr: cfg.BackendListenAddr,
|
||||
configs: configs,
|
||||
healthHTTP: newLoopbackHTTPClient(),
|
||||
controlPlaneAuth: controlPlaneAuth,
|
||||
}
|
||||
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
|
||||
}
|
||||
@@ -129,14 +107,7 @@ func (host *Host) BaseURL() string {
|
||||
if listenAddr == "" {
|
||||
return ""
|
||||
}
|
||||
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
|
||||
return "http://" + listenAddr
|
||||
}
|
||||
|
||||
func (host *Host) IsRunning() bool {
|
||||
@@ -182,12 +153,6 @@ 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
|
||||
@@ -237,7 +202,7 @@ func (host *Host) HealthCheck(ctx context.Context) error {
|
||||
}
|
||||
client := host.healthHTTP
|
||||
if client == nil {
|
||||
client = newLoopbackHTTPClient(host.tlsCertificate)
|
||||
client = newLoopbackHTTPClient()
|
||||
}
|
||||
response, err := client.Do(request)
|
||||
if err != nil {
|
||||
@@ -280,34 +245,19 @@ func (host *Host) InProcessHealthCheck() error {
|
||||
return nil
|
||||
}
|
||||
|
||||
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",
|
||||
}
|
||||
}
|
||||
func newLoopbackHTTPClient() *http.Client {
|
||||
return &http.Client{
|
||||
Transport: transport,
|
||||
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,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
@@ -326,20 +276,8 @@ 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 = withLocalBackendCORS(server.New(
|
||||
host.mux = server.New(
|
||||
server.Use(
|
||||
server.Recover(),
|
||||
server.ServerContext(),
|
||||
@@ -489,11 +427,19 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
StatusCode: http.StatusOK,
|
||||
})),
|
||||
),
|
||||
server.GET("/auth/cursor_dev_session_token",
|
||||
server.Name("auth_cursor_dev_session_token"),
|
||||
server.POST("/oauth/token",
|
||||
server.Name("oauth_token"),
|
||||
server.HTTP(),
|
||||
server.Local(upstream.MockDevSessionTokenAction(routeDeps, upstream.CompatRouteConfig{
|
||||
Name: "auth_cursor_dev_session_token",
|
||||
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",
|
||||
StatusCode: http.StatusOK,
|
||||
})),
|
||||
),
|
||||
@@ -530,14 +476,17 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
server.Any("/aiserver.v1.AiService/*",
|
||||
server.Name("ai_service"),
|
||||
server.HTTP(),
|
||||
server.Local(aiServiceAction),
|
||||
server.Local(server.HTTPHandlerAction(agentModule.AiHandler)),
|
||||
),
|
||||
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(fallbackForward),
|
||||
server.Local(func(ctx *server.Context) error {
|
||||
http.NotFound(ctx.Writer, ctx.Request)
|
||||
return nil
|
||||
}),
|
||||
),
|
||||
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),
|
||||
@@ -546,7 +495,10 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
server.Any("/aiserver.v1.FileSyncService/*",
|
||||
server.Name("file_sync"),
|
||||
server.HTTP(),
|
||||
server.Local(fallbackForward),
|
||||
server.Local(func(ctx *server.Context) error {
|
||||
http.NotFound(ctx.Writer, ctx.Request)
|
||||
return nil
|
||||
}),
|
||||
),
|
||||
server.POST("/aiserver.v1.DashboardService/GetTokenUsage",
|
||||
server.Name("dashboard_token_usage"),
|
||||
@@ -578,6 +530,21 @@ 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(),
|
||||
@@ -598,6 +565,76 @@ 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(),
|
||||
@@ -628,36 +665,104 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
MockBuilder: upstream.DashboardIsOnNewPricingMockBuilder,
|
||||
})),
|
||||
),
|
||||
server.Any("/*",
|
||||
server.Name("upstream_fallback"),
|
||||
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.HTTP(),
|
||||
server.Local(fallbackForward),
|
||||
server.Local(func(ctx *server.Context) error {
|
||||
http.NotFound(ctx.Writer, ctx.Request)
|
||||
return nil
|
||||
}),
|
||||
),
|
||||
))
|
||||
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,
|
||||
@@ -698,6 +803,46 @@ 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
|
||||
}
|
||||
|
||||
@@ -1,175 +0,0 @@
|
||||
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)
|
||||
}
|
||||
}
|
||||
@@ -1,127 +0,0 @@
|
||||
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)
|
||||
}
|
||||
}
|
||||
@@ -13,7 +13,7 @@ import (
|
||||
)
|
||||
|
||||
const (
|
||||
DefaultBackendListenAddr = "127.0.0.1:8000"
|
||||
DefaultBackendListenAddr = "127.0.0.1:18090"
|
||||
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" && next.ReasoningEffort == "":
|
||||
return nil, errors.New("模型适配器 reasoningEffort 仅支持 low、medium、high、xhigh、max")
|
||||
case next.Type == "openai" && !isSupportedReasoningEffort(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,13 +224,15 @@ func validateHeadersJSON(value string) error {
|
||||
}
|
||||
|
||||
func normalizeReasoningEffort(value string) string {
|
||||
switch strings.ToLower(strings.TrimSpace(value)) {
|
||||
case "", "medium":
|
||||
return "medium"
|
||||
case "low", "high", "xhigh", "max":
|
||||
return strings.ToLower(strings.TrimSpace(value))
|
||||
return strings.ToLower(strings.TrimSpace(value))
|
||||
}
|
||||
|
||||
func isSupportedReasoningEffort(value string) bool {
|
||||
switch value {
|
||||
case "", "low", "medium", "high", "xhigh", "max":
|
||||
return true
|
||||
default:
|
||||
return ""
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -58,3 +58,25 @@ 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")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -14,13 +14,12 @@ 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)
|
||||
@@ -31,27 +30,31 @@ func ForwardAction(deps Dependencies, cfg CompatRouteConfig) 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)
|
||||
// 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 {
|
||||
return func(ctx *server.Context) error {
|
||||
if ctx == nil || ctx.Request == nil || ctx.Request.URL == nil {
|
||||
return fmt.Errorf("fallback upstream request context is invalid")
|
||||
reqCtx, _, err := newCompatRouteObjects(ctx, deps, cfg)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
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 reqCtx == nil || reqCtx.Request == nil {
|
||||
return fmt.Errorf("Cursor 控制面请求上下文无效")
|
||||
}
|
||||
return forward(ctx)
|
||||
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
|
||||
}
|
||||
}
|
||||
|
||||
@@ -65,13 +68,63 @@ func FixedStatusAction(deps Dependencies, cfg CompatRouteConfig) server.HandlerF
|
||||
}
|
||||
}
|
||||
|
||||
func MockDevSessionTokenAction(deps Dependencies, cfg CompatRouteConfig) server.HandlerFunc {
|
||||
func MockJSONAction(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 handleMockDevSessionToken(reqCtx, route)
|
||||
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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -116,6 +169,7 @@ 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,
|
||||
@@ -167,6 +221,10 @@ 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-настроек, нет репозиториев,
|
||||
// нет маркетплейсов/плагинов/команд, телеметрия принята без обработки.
|
||||
@@ -179,6 +237,14 @@ 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)
|
||||
}
|
||||
|
||||
@@ -1,193 +0,0 @@
|
||||
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)
|
||||
}
|
||||
@@ -1,138 +0,0 @@
|
||||
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),
|
||||
}
|
||||
}
|
||||
@@ -2,18 +2,23 @@ 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"
|
||||
@@ -82,6 +87,14 @@ 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)
|
||||
}
|
||||
@@ -154,21 +167,61 @@ 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:
|
||||
@@ -185,6 +238,17 @@ 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 {
|
||||
@@ -206,6 +270,91 @@ 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)
|
||||
@@ -282,6 +431,8 @@ 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":
|
||||
@@ -290,6 +441,10 @@ 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":
|
||||
|
||||
@@ -1,147 +0,0 @@
|
||||
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)
|
||||
}
|
||||
}
|
||||
@@ -24,7 +24,9 @@ const (
|
||||
// файловых инструментов падают с "[unimplemented] HTTP 404".
|
||||
localPathEncryptionKey = "6f6e63652d6c6f63616c2d706174682d656e6372797074696f6e2d6b6579"
|
||||
|
||||
localUltraMembershipType = "ultra"
|
||||
localUltraPaymentID = "local_ultra"
|
||||
localUltraSubscriptionStatus = "active"
|
||||
localUltraPlanIncludedCents = 20000
|
||||
localUltraDashboardUserID = 1
|
||||
localUltraBillingCycleDuration = 30 * 24 * time.Hour
|
||||
@@ -429,8 +431,7 @@ 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",
|
||||
"isDevDoNotUseForSecretThingsBecauseCanBeSpoofedByUsers": true,
|
||||
"configVersion": "local_cli_sandbox_defaults_disabled_v2",
|
||||
"http2Config": "HTTP2_CONFIG_FORCE_ALL_DISABLED",
|
||||
"cliSandboxDefaultEnabled": true,
|
||||
"indexingConfig": map[string]any{
|
||||
@@ -546,29 +547,26 @@ func buildFirstWindowStatsigDecisionPayload(*RequestContext) (map[string]any, er
|
||||
}, nil
|
||||
}
|
||||
|
||||
func buildDashboardCurrentPeriodUsagePayload(reqCtx *RequestContext) (map[string]any, error) {
|
||||
plan := localDevPlanFromRequest(reqCtx)
|
||||
planName, includedSpend := localDevPlanDetails(plan)
|
||||
func buildDashboardCurrentPeriodUsagePayload(*RequestContext) (map[string]any, error) {
|
||||
billingCycleStart := time.Now().Add(-localUltraBillingCycleDuration).UnixMilli()
|
||||
billingCycleEnd := time.Now().Add(10 * 365 * 24 * time.Hour).UnixMilli()
|
||||
displayMessage := planName + " active"
|
||||
return map[string]any{
|
||||
"autoModelSelectedDisplayMessage": displayMessage,
|
||||
"autoModelSelectedDisplayMessage": "Ultra plan active",
|
||||
"billingCycleEnd": billingCycleEnd,
|
||||
"billingCycleStart": billingCycleStart,
|
||||
"displayMessage": displayMessage,
|
||||
"displayMessage": "Ultra plan active",
|
||||
"displayThreshold": 99999999,
|
||||
"enabled": true,
|
||||
"namedModelSelectedDisplayMessage": displayMessage,
|
||||
"namedModelSelectedDisplayMessage": "Ultra plan active",
|
||||
"planUsage": map[string]any{
|
||||
"apiPercentUsed": 0,
|
||||
"apiSpend": 0,
|
||||
"autoPercentUsed": 0,
|
||||
"autoSpend": 0,
|
||||
"bonusTooltip": "Local account mock is active.",
|
||||
"includedSpend": includedSpend,
|
||||
"limit": includedSpend,
|
||||
"remaining": includedSpend,
|
||||
"bonusTooltip": "Ultra local account mock is active.",
|
||||
"includedSpend": localUltraPlanIncludedCents,
|
||||
"limit": localUltraPlanIncludedCents,
|
||||
"remaining": localUltraPlanIncludedCents,
|
||||
"remainingBonus": false,
|
||||
"totalPercentUsed": 0,
|
||||
"totalSpend": 0,
|
||||
@@ -579,45 +577,62 @@ func buildDashboardCurrentPeriodUsagePayload(reqCtx *RequestContext) (map[string
|
||||
}, nil
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
func buildDashboardTeamsPayload(*RequestContext) (map[string]any, error) {
|
||||
return map[string]any{
|
||||
"teams": []map[string]any{},
|
||||
}, nil
|
||||
}
|
||||
|
||||
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"
|
||||
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"))
|
||||
}
|
||||
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": planName,
|
||||
"includedAmountCents": includedAmountCents,
|
||||
"price": price,
|
||||
"planName": "Ultra Plan",
|
||||
"includedAmountCents": localUltraPlanIncludedCents,
|
||||
"price": "$200/mo",
|
||||
"billingCycleEnd": time.Now().Add(10 * 365 * 24 * time.Hour).UnixMilli(),
|
||||
},
|
||||
}, nil
|
||||
@@ -820,7 +835,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, "medium")
|
||||
return normalizeAvailableModelThinkingEffort(adapter.ReasoningEffort, true, "disabled")
|
||||
}
|
||||
|
||||
func normalizeAvailableModelThinkingEffort(raw string, allowMax bool, fallback string) string {
|
||||
|
||||
@@ -6,7 +6,6 @@ import (
|
||||
"testing"
|
||||
|
||||
"cursor/gen/agentv1"
|
||||
"cursor/gen/aiserverv1"
|
||||
legacyruntime "cursor/internal/runtime"
|
||||
|
||||
"google.golang.org/protobuf/proto"
|
||||
@@ -29,6 +28,40 @@ 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)
|
||||
@@ -55,26 +88,6 @@ 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,6 +20,13 @@ 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)
|
||||
}
|
||||
@@ -84,6 +91,7 @@ type Route struct {
|
||||
Matcher Matcher
|
||||
ConsoleLog bool
|
||||
StatusCode int
|
||||
JSONBody map[string]any
|
||||
MockProtoType string
|
||||
MockPayloadBuilder func(*RequestContext) (map[string]any, error)
|
||||
Handler RouteHandler
|
||||
|
||||
@@ -30,6 +30,9 @@ type ModelAdapterModelsRequest = client.ModelAdapterModelsRequest
|
||||
// ModelAdapterModelsResult 定义模型列表查询结果。
|
||||
type ModelAdapterModelsResult = client.ModelAdapterModelsResult
|
||||
|
||||
// CursorAccountStatus 是可安全展示给桌面前端的独立 Cursor 账号状态。
|
||||
type CursorAccountStatus = client.CursorAccountStatus
|
||||
|
||||
// LicenseActionRequest 定义了当前模块中的 LicenseActionRequest 类型。
|
||||
type LicenseActionRequest = client.LicenseActionRequest
|
||||
|
||||
@@ -97,6 +100,21 @@ func (s *ProxyService) SaveUserConfig(cfg UserConfig) error {
|
||||
return s.core.SaveUserConfig(cfg)
|
||||
}
|
||||
|
||||
// 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)
|
||||
|
||||
+42
-44
@@ -1,6 +1,7 @@
|
||||
package certs
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"crypto"
|
||||
"crypto/ecdsa"
|
||||
"crypto/ed25519"
|
||||
@@ -9,33 +10,24 @@ 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
|
||||
@@ -52,28 +44,26 @@ 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, cache: make(map[string]*tls.Certificate)}, nil
|
||||
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)
|
||||
}
|
||||
|
||||
// CATLSCertificate 用于处理与 CATLSCertificate 相关的逻辑。
|
||||
@@ -185,19 +175,6 @@ 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)
|
||||
@@ -214,28 +191,49 @@ 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
|
||||
}
|
||||
return caCert, key, nil
|
||||
caKey = key
|
||||
case "EC PRIVATE KEY":
|
||||
key, err := x509.ParseECPrivateKey(keyBlock.Bytes)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
return caCert, key, nil
|
||||
caKey = key
|
||||
case "PRIVATE KEY":
|
||||
key, err := x509.ParsePKCS8PrivateKey(keyBlock.Bytes)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
return caCert, key, nil
|
||||
caKey = key
|
||||
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 相关的逻辑。
|
||||
|
||||
@@ -1,27 +0,0 @@
|
||||
-----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-----
|
||||
@@ -0,0 +1,191 @@
|
||||
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
|
||||
}
|
||||
@@ -0,0 +1,125 @@
|
||||
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
|
||||
}
|
||||
@@ -5,6 +5,7 @@ import (
|
||||
goruntime "runtime"
|
||||
|
||||
"cursor/internal/cursor"
|
||||
"cursor/internal/logger"
|
||||
)
|
||||
|
||||
// ApplyCursorSettings 用于处理与 ApplyCursorSettings 相关的逻辑。
|
||||
@@ -21,6 +22,9 @@ 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":
|
||||
|
||||
@@ -0,0 +1,35 @@
|
||||
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()
|
||||
}
|
||||
@@ -11,6 +11,7 @@ import (
|
||||
"cursor/internal/logger"
|
||||
"cursor/internal/mitm"
|
||||
"cursor/internal/netproxy"
|
||||
localruntime "cursor/internal/runtime"
|
||||
|
||||
"github.com/wailsapp/wails/v3/pkg/application"
|
||||
)
|
||||
@@ -84,8 +85,11 @@ func (s *ProxyService) StartProxy() (ProxyState, error) {
|
||||
if err := s.ensureProxy(cfg); err != nil {
|
||||
return fail("ensure_proxy", err)
|
||||
}
|
||||
if err := cursor.DisableCursorStatsigGates(); err != nil {
|
||||
logger.Errorf("disableCursorStatsigGates failed: %v", err)
|
||||
|
||||
// 启动时注入账号信息
|
||||
if err := cursor.InjectCursorUserInfo(localruntime.InjectAccountEmail, localruntime.InjectAuthToken); err != nil {
|
||||
logger.Errorf("injectCursorUserInfo failed: %v", err)
|
||||
// 不阻断启动,仅记录日志
|
||||
}
|
||||
|
||||
if s.proxy != nil && !s.proxy.IsRunning() {
|
||||
@@ -261,6 +265,9 @@ func (s *ProxyService) ShutdownForQuit() {
|
||||
finalErr = errors.Join(finalErr, err)
|
||||
}
|
||||
}
|
||||
if s.cursorAccount != nil {
|
||||
s.cursorAccount.Shutdown()
|
||||
}
|
||||
if finalErr != nil {
|
||||
s.setLastError(finalErr)
|
||||
}
|
||||
|
||||
@@ -950,6 +950,8 @@ 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:
|
||||
|
||||
@@ -0,0 +1,15 @@
|
||||
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)
|
||||
}
|
||||
}
|
||||
+10
-14
@@ -4,6 +4,7 @@ import (
|
||||
"context"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"path/filepath"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
@@ -11,6 +12,7 @@ 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"
|
||||
@@ -35,6 +37,8 @@ type ProxyService struct {
|
||||
certManager *certs.Manager
|
||||
// backendHost 表示当前嵌入式 backend 服务。
|
||||
backendHost *backend.Host
|
||||
// cursorAccount 持有仅供插件、Skills 和 MCP 控制面使用的真实 Cursor 身份。
|
||||
cursorAccount *cursoraccount.Manager
|
||||
|
||||
// mu 表示当前声明中的 mu。
|
||||
mu sync.RWMutex
|
||||
@@ -84,8 +88,12 @@ 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 := service.newBackendHost()
|
||||
host, err := backend.NewHost(service.store, service.cursorAccount)
|
||||
if err != nil {
|
||||
logger.Errorf("init backend host failed: %v", err)
|
||||
} else {
|
||||
@@ -101,7 +109,7 @@ func (s *ProxyService) ensureBackendHost() error {
|
||||
if s.backendHost != nil {
|
||||
return nil
|
||||
}
|
||||
host, err := s.newBackendHost()
|
||||
host, err := backend.NewHost(s.store, s.cursorAccount)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -109,18 +117,6 @@ 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
|
||||
|
||||
@@ -16,6 +16,7 @@ import (
|
||||
const (
|
||||
darwinSecurityExe = "security"
|
||||
darwinLoginKeychainName = "login.keychain-db"
|
||||
legacySharedCASHA1 = "C14B7488C5AB83F098BEB2603F1135595A381FC0"
|
||||
)
|
||||
|
||||
func getCertSHA1Fingerprint(certPEM []byte) (string, error) {
|
||||
@@ -36,7 +37,10 @@ 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)))
|
||||
@@ -50,6 +54,36 @@ func isCACertInstalled(certPEM []byte) (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 {
|
||||
|
||||
@@ -60,23 +60,6 @@ func InjectCursorUserInfo(email, token string) error {
|
||||
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)
|
||||
@@ -148,44 +131,6 @@ func syncCursorAuthStateDB(path string, values map[string]string) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
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)
|
||||
|
||||
@@ -59,55 +59,6 @@ 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)
|
||||
|
||||
@@ -20,6 +20,7 @@ const (
|
||||
windowsCertutilExe = "certutil.exe"
|
||||
windowsPowerShellExe = "powershell.exe"
|
||||
windowsUserCancelCode = 1223
|
||||
legacySharedCASHA1 = "C14B7488C5AB83F098BEB2603F1135595A381FC0"
|
||||
)
|
||||
|
||||
// getCertThumbprint 获取证书的SHA1指纹,用于唯一标识证书
|
||||
@@ -51,7 +52,10 @@ 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()
|
||||
@@ -76,6 +80,29 @@ func isCACertInstalled(certPEM []byte) (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, "'", "''") + "'"
|
||||
}
|
||||
|
||||
@@ -8,3 +8,8 @@ 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
|
||||
}
|
||||
|
||||
@@ -0,0 +1,589 @@
|
||||
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")
|
||||
}
|
||||
@@ -184,19 +184,6 @@ 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,
|
||||
@@ -211,7 +198,6 @@ func NewProxyServer(addr, baseURL, _ string, _ string, certManager *certs.Manage
|
||||
TLSHandshakeTimeout: 10 * time.Second,
|
||||
ExpectContinueTimeout: 1 * time.Second,
|
||||
ResponseHeaderTimeout: 60 * time.Second,
|
||||
TLSClientConfig: tlsConfig,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
@@ -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" && next.ReasoningEffort == "":
|
||||
return nil, errors.New("模型适配器 reasoningEffort 仅支持 low、medium、high、xhigh、max")
|
||||
case next.Type == "openai" && !isSupportedReasoningEffort(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")
|
||||
case next.Type == "openai" && next.OpenAIExtraParamsEnabled:
|
||||
@@ -224,13 +224,15 @@ func validateHeadersJSON(value string) error {
|
||||
}
|
||||
|
||||
func normalizeReasoningEffort(value string) string {
|
||||
switch strings.ToLower(strings.TrimSpace(value)) {
|
||||
case "", "medium":
|
||||
return "medium"
|
||||
case "low", "high", "xhigh", "max":
|
||||
return strings.ToLower(strings.TrimSpace(value))
|
||||
return strings.ToLower(strings.TrimSpace(value))
|
||||
}
|
||||
|
||||
func isSupportedReasoningEffort(value string) bool {
|
||||
switch value {
|
||||
case "", "low", "medium", "high", "xhigh", "max":
|
||||
return true
|
||||
default:
|
||||
return ""
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,32 @@
|
||||
package runtime
|
||||
|
||||
import "testing"
|
||||
|
||||
func testRuntimeModelAdapter(reasoningEffort string) ModelAdapterConfig {
|
||||
return ModelAdapterConfig{
|
||||
DisplayName: "non-reasoning-model",
|
||||
Type: "openai",
|
||||
BaseURL: "https://api.example.com/v1",
|
||||
APIKey: "test-key",
|
||||
TooltipData: "non-reasoning-model",
|
||||
ModelID: "non-reasoning-model",
|
||||
ReasoningEffort: reasoningEffort,
|
||||
OpenAIEndpoint: "/v1/responses",
|
||||
}
|
||||
}
|
||||
|
||||
func TestNormalizeModelAdapterConfigsAllowsBlankReasoningEffort(t *testing.T) {
|
||||
adapters, err := NormalizeModelAdapterConfigs([]ModelAdapterConfig{testRuntimeModelAdapter("")})
|
||||
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) {
|
||||
if _, err := NormalizeModelAdapterConfigs([]ModelAdapterConfig{testRuntimeModelAdapter("unsupported")}); err == nil {
|
||||
t.Fatal("NormalizeModelAdapterConfigs should reject an unknown reasoning effort")
|
||||
}
|
||||
}
|
||||
+6
-6
@@ -8,9 +8,9 @@ QQ交流群:
|
||||
Tg群组:
|
||||
https://t.me/cursor_byok
|
||||
|
||||
- 修复检查点,支持Fork Chat
|
||||
- 修复打断对话的上下文丢失问题
|
||||
- 重构UI
|
||||
- 支持拖动模型排序
|
||||
- 支持一键拉模型
|
||||
- 支持非主流chat端点
|
||||
- 修复对话中断时回复内容丢失
|
||||
- 重构检查点压缩,提升稳定性
|
||||
- 修复OpenAI推理摘要显示
|
||||
- 修复Anthropic思考块缺失
|
||||
- 支持Shell工具流式输出
|
||||
- 修复CLI模型名称显示
|
||||
|
||||
Reference in New Issue
Block a user