This commit is contained in:
leokun
2026-06-30 10:38:52 +08:00
commit c083be5ec2
312 changed files with 146628 additions and 0 deletions
+84
View File
@@ -0,0 +1,84 @@
package ads
import (
"net/http"
"strings"
)
const bridgeScript = `<script>
(function () {
var source = "cursor-ad";
function post(type, payload) {
if (!window.parent) {
return;
}
var message = Object.assign({ source: source, type: type }, payload || {});
window.parent.postMessage(message, "*");
}
window.AdBridge = Object.assign({}, window.AdBridge || {}, {
close: function () {
post("close");
},
openExternal: function (url) {
post("openExternal", { url: String(url || "") });
}
});
})();
</script>`
func NewHTTPHandler(storeRoot string) http.Handler {
service := NewService(Options{StoreRoot: storeRoot})
return http.HandlerFunc(service.ServeHTTP)
}
func (service *Service) ServeHTTP(writer http.ResponseWriter, request *http.Request) {
if request == nil {
http.NotFound(writer, request)
return
}
if request.Method != http.MethodGet && request.Method != http.MethodHead {
http.Error(writer, "method not allowed", http.StatusMethodNotAllowed)
return
}
asset, packageHash, ok, err := service.LoadAsset(request.Context(), request.URL.Path)
if err != nil {
http.Error(writer, err.Error(), http.StatusBadRequest)
return
}
if !ok {
http.NotFound(writer, request)
return
}
payload := asset.data
contentType := strings.TrimSpace(asset.contentType)
if asset.path == "index.html" {
payload = injectBridge(payload)
contentType = "text/html; charset=utf-8"
}
if contentType == "" {
contentType = "application/octet-stream"
}
writer.Header().Set("Content-Type", contentType)
writer.Header().Set("Cache-Control", "no-store")
writer.Header().Set("X-Content-Type-Options", "nosniff")
if packageHash != "" {
writer.Header().Set("ETag", `"`+packageHash+`"`)
}
writer.WriteHeader(http.StatusOK)
if request.Method == http.MethodHead {
return
}
_, _ = writer.Write(payload)
}
func injectBridge(data []byte) []byte {
html := string(data)
lower := strings.ToLower(html)
if index := strings.Index(lower, "</head>"); index >= 0 {
return []byte(html[:index] + bridgeScript + html[index:])
}
if index := strings.Index(lower, "<body"); index >= 0 {
return []byte(html[:index] + bridgeScript + html[index:])
}
return []byte(bridgeScript + html)
}
+816
View File
@@ -0,0 +1,816 @@
package ads
import (
"archive/zip"
"bytes"
"context"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"fmt"
"io"
"mime"
"net/http"
"os"
"path"
"path/filepath"
"runtime"
"strconv"
"strings"
"sync"
"time"
"cursor/internal/netproxy"
"gopkg.in/yaml.v3"
)
const (
defaultFetchTimeout = 45 * time.Second
maxZipBytes = 50 << 20
maxAssetBytes = 20 << 20
noAdStatusCode = http.StatusNotFound
)
type MetricsProvider func(context.Context) (MetricsSnapshot, error)
type ProviderCountProvider func(context.Context) (int, error)
type DeviceIDProvider func() (string, error)
type Options struct {
StoreRoot string
HTTPClient *http.Client
AppVersion string
AssetBaseURL string
DeviceID DeviceIDProvider
Metrics MetricsProvider
ProviderCount ProviderCountProvider
}
type Service struct {
storeRoot string
httpClient *http.Client
appVersion string
assetBaseURL string
deviceID DeviceIDProvider
metrics MetricsProvider
providerCount ProviderCountProvider
assetBaseURLMu sync.RWMutex
refreshMu sync.Mutex
}
type parsedPackage struct {
hash string
rawConfig []byte
configJSON string
assets []parsedAsset
}
type parsedAsset struct {
path string
contentType string
data []byte
}
type adPackageFile struct {
Hash string `json:"hash"`
ConfigJSON string `json:"config_json"`
Assets map[string]adAssetRef `json:"assets"`
FetchedAt time.Time `json:"fetched_at"`
ActivatedAt time.Time `json:"activated_at"`
}
type adAssetRef struct {
File string `json:"file"`
ContentType string `json:"content_type"`
}
type adConfigYAML struct {
Enabled *bool `yaml:"enabled"`
Window WindowConfig `yaml:"window"`
Home HomePlacementConfig `yaml:"home"`
}
type packageState int
const (
packageMissing packageState = iota
packageValid
packageCorrupt
)
type packageInspection struct {
state packageState
pkg adPackageFile
}
func NewService(options Options) *Service {
client := options.HTTPClient
if client == nil {
client = netproxy.NewHTTPClient(defaultFetchTimeout)
}
return &Service{
storeRoot: strings.TrimSpace(options.StoreRoot),
httpClient: client,
appVersion: strings.TrimSpace(options.AppVersion),
assetBaseURL: normalizeAssetBaseURL(options.AssetBaseURL),
deviceID: options.DeviceID,
metrics: options.Metrics,
providerCount: options.ProviderCount,
}
}
func (service *Service) SetAssetBaseURL(rawURL string) bool {
if service == nil {
return false
}
next := normalizeAssetBaseURL(rawURL)
service.assetBaseURLMu.Lock()
defer service.assetBaseURLMu.Unlock()
if service.assetBaseURL == next {
return false
}
service.assetBaseURL = next
return true
}
func (service *Service) currentAssetBaseURL() string {
if service == nil {
return ""
}
service.assetBaseURLMu.RLock()
defer service.assetBaseURLMu.RUnlock()
return service.assetBaseURL
}
func normalizeAssetBaseURL(rawURL string) string {
return strings.TrimRight(strings.TrimSpace(rawURL), "/")
}
func (service *Service) Refresh(parent context.Context) (Runtime, bool, error) {
if service == nil {
return Runtime{}, false, fmt.Errorf("ad service is nil")
}
if parent == nil {
parent = context.Background()
}
service.refreshMu.Lock()
defer service.refreshMu.Unlock()
ctx, cancel := context.WithTimeout(parent, defaultFetchTimeout)
defer cancel()
result, err := service.FetchOnce(ctx)
if err != nil {
return Runtime{}, false, err
}
runtimeState, err := service.GetRuntime(context.Background())
if err != nil {
return Runtime{}, result.Changed, err
}
return runtimeState, result.Changed, nil
}
func (service *Service) FetchOnce(ctx context.Context) (FetchResult, error) {
if service == nil {
return FetchResult{}, fmt.Errorf("ad service is nil")
}
var output FetchResult
var firstErr error
var succeeded bool
for _, slot := range Slots {
result, err := service.fetchSlotOnce(ctx, slot)
if err != nil {
if firstErr == nil {
firstErr = err
}
continue
}
succeeded = true
if result.Changed {
output.Changed = true
}
if output.Hash == "" {
output.Hash = result.Hash
}
}
if !succeeded && firstErr != nil {
return FetchResult{}, firstErr
}
return output, nil
}
func (service *Service) fetchSlotOnce(ctx context.Context, slot Slot) (FetchResult, error) {
slotID := normalizeSlotID(slot.ID)
fetchURL := strings.TrimSpace(slot.FetchURL)
if fetchURL == "" {
return FetchResult{}, nil
}
inspection := service.inspectSlotPackage(ctx, slotID)
currentHash := ""
if inspection.state == packageValid {
currentHash = strings.TrimSpace(inspection.pkg.Hash)
}
request, err := http.NewRequestWithContext(ctx, http.MethodGet, fetchURL, nil)
if err != nil {
return FetchResult{}, err
}
service.applyReportHeaders(ctx, request, currentHash)
response, err := service.httpClient.Do(request)
if err != nil {
return FetchResult{}, err
}
defer response.Body.Close()
if response.StatusCode == noAdStatusCode {
changed, err := service.clearPackage(ctx, slotID)
if err != nil {
return FetchResult{}, err
}
return FetchResult{Changed: changed}, nil
}
if response.StatusCode < http.StatusOK || response.StatusCode >= http.StatusMultipleChoices {
return FetchResult{}, fmt.Errorf("广告拉取返回状态码 %d", response.StatusCode)
}
body, err := readLimited(response.Body, maxZipBytes, "广告压缩包")
if err != nil {
return FetchResult{}, err
}
hashBytes := sha256.Sum256(body)
hash := slotStorageHash(slotID, hex.EncodeToString(hashBytes[:]))
if inspection.state == packageValid && strings.EqualFold(hash, currentHash) {
return FetchResult{Hash: hash, Changed: false}, nil
}
parsed, err := parsePackage(hash, body)
if err != nil {
return FetchResult{}, err
}
if err := service.replacePackage(ctx, slotID, parsed); err != nil {
return FetchResult{}, err
}
return FetchResult{Hash: hash, Changed: true}, nil
}
func (service *Service) CurrentHash(ctx context.Context) (string, error) {
return service.currentHash(ctx, normalizeSlotID(""))
}
func (service *Service) currentHash(ctx context.Context, slotID string) (string, error) {
inspection := service.inspectSlotPackage(ctx, slotID)
if inspection.state != packageValid {
return "", nil
}
return strings.TrimSpace(inspection.pkg.Hash), nil
}
func (service *Service) GetRuntime(ctx context.Context) (Runtime, error) {
runtimeState := Runtime{
AssetBaseURL: service.currentAssetBaseURL(),
}
for _, slot := range Slots {
slotRuntime, err := service.getSlotRuntime(ctx, normalizeSlotID(slot.ID))
if err != nil {
return runtimeState, err
}
runtimeState.Slots = append(runtimeState.Slots, slotRuntime)
}
if len(runtimeState.Slots) > 0 {
first := runtimeState.Slots[0]
runtimeState.Available = first.Available
runtimeState.Enabled = first.Enabled
runtimeState.PackageHash = first.PackageHash
runtimeState.AssetBaseURL = first.AssetBaseURL
runtimeState.IndexURL = first.IndexURL
runtimeState.Window = first.Window
runtimeState.Home = first.Home
}
return runtimeState, nil
}
func (service *Service) getSlotRuntime(ctx context.Context, slotID string) (SlotRuntime, error) {
slotID = normalizeSlotID(slotID)
runtimeState := SlotRuntime{
ID: slotID,
AssetBaseURL: service.assetURL(slotID, ""),
IndexURL: service.assetURL(slotID, "index.html"),
}
inspection := service.inspectSlotPackage(ctx, slotID)
if inspection.state != packageValid {
return runtimeState, nil
}
var cfg Config
if err := json.Unmarshal([]byte(inspection.pkg.ConfigJSON), &cfg); err != nil {
return runtimeState, nil
}
runtimeState.Available = true
runtimeState.Enabled = cfg.Enabled
runtimeState.PackageHash = strings.TrimSpace(inspection.pkg.Hash)
runtimeState.Window = cfg.Window
runtimeState.Home = cfg.Home
return runtimeState, nil
}
func (service *Service) LoadAsset(ctx context.Context, rawPath string) (parsedAsset, string, bool, error) {
slotID, assetPath, err := normalizeRequestAssetPath(rawPath)
if err != nil {
return parsedAsset{}, "", false, err
}
inspection := service.inspectSlotPackage(ctx, slotID)
if inspection.state != packageValid {
return parsedAsset{}, "", false, nil
}
ref, ok := inspection.pkg.Assets[assetPath]
if !ok {
return parsedAsset{}, inspection.pkg.Hash, false, nil
}
data, ok, err := service.readAssetFile(slotID, assetPath, ref)
if err != nil || !ok {
return parsedAsset{}, inspection.pkg.Hash, false, nil
}
return parsedAsset{
path: assetPath,
contentType: ref.ContentType,
data: data,
}, inspection.pkg.Hash, true, nil
}
func (service *Service) applyReportHeaders(ctx context.Context, request *http.Request, currentHash string) {
if request == nil {
return
}
version := firstNonEmpty(service.appVersion, "0.0.0")
request.Header.Set("Accept", "application/zip, application/octet-stream, */*")
request.Header.Set("User-Agent", "cursor-local-assistant/"+headerValue(version))
request.Header.Set("X-Cursor-Assistant-Version", headerValue(version))
request.Header.Set("X-Cursor-Assistant-OS", headerValue(displayOSName()))
request.Header.Set("X-Cursor-Assistant-OS-Version", headerValue(displayOSVersion()))
request.Header.Set("X-Cursor-Assistant-Arch", headerValue(runtime.GOARCH))
request.Header.Set("X-Cursor-Assistant-Current-Ad-Hash", headerValue(stripSlotStorageHash(currentHash)))
if service.deviceID != nil {
if value, err := service.deviceID(); err == nil {
request.Header.Set("X-Cursor-Assistant-Device-ID", headerValue(value))
}
}
if service.metrics != nil {
if metrics, err := service.metrics(ctx); err == nil {
request.Header.Set("X-Cursor-Assistant-Turns", strconv.Itoa(metrics.TurnsTotal))
request.Header.Set("X-Cursor-Assistant-Request-Tokens", strconv.FormatInt(metrics.RequestTokensTotal, 10))
request.Header.Set("X-Cursor-Assistant-Prompt-Tokens", strconv.FormatInt(metrics.PromptTokensTotal, 10))
request.Header.Set("X-Cursor-Assistant-Cache-Read-Tokens", strconv.FormatInt(metrics.CacheReadTokens, 10))
request.Header.Set("X-Cursor-Assistant-Cache-Write-Tokens", strconv.FormatInt(metrics.CacheWriteTokens, 10))
}
}
if service.providerCount != nil {
if count, err := service.providerCount(ctx); err == nil {
request.Header.Set("X-Cursor-Assistant-Provider-Count", strconv.Itoa(maxInt(count, 0)))
}
}
}
func (service *Service) currentPackage(ctx context.Context, slotID string) (adPackageFile, bool, error) {
_ = ctx
inspection := service.inspectSlotPackage(ctx, slotID)
switch inspection.state {
case packageMissing:
return adPackageFile{}, false, nil
case packageValid:
return inspection.pkg, true, nil
default:
return adPackageFile{}, false, fmt.Errorf("广告缓存已损坏")
}
}
func (service *Service) inspectSlotPackage(ctx context.Context, slotID string) packageInspection {
_ = ctx
body, err := os.ReadFile(service.slotPackagePath(slotID))
if err != nil {
if os.IsNotExist(err) {
return packageInspection{state: packageMissing}
}
return packageInspection{state: packageCorrupt}
}
var item adPackageFile
if err := json.Unmarshal(body, &item); err != nil {
return packageInspection{state: packageCorrupt}
}
if item.Assets == nil {
item.Assets = make(map[string]adAssetRef)
}
if !validSlotStorageHash(slotID, item.Hash) {
return packageInspection{state: packageCorrupt}
}
var cfg Config
if err := json.Unmarshal([]byte(item.ConfigJSON), &cfg); err != nil {
return packageInspection{state: packageCorrupt}
}
if cfg.Window.Width <= 0 || cfg.Window.Height <= 0 {
return packageInspection{state: packageCorrupt}
}
if _, ok := item.Assets["index.html"]; !ok {
return packageInspection{state: packageCorrupt}
}
for assetPath, ref := range item.Assets {
if _, err := normalizeArchivePath(assetPath); err != nil {
return packageInspection{state: packageCorrupt}
}
if _, ok, err := service.readAssetFile(slotID, assetPath, ref); err != nil || !ok {
return packageInspection{state: packageCorrupt}
}
}
return packageInspection{state: packageValid, pkg: item}
}
func (service *Service) readAssetFile(slotID string, assetPath string, ref adAssetRef) ([]byte, bool, error) {
fileName := strings.TrimSpace(ref.File)
if !isSafeAssetFileName(fileName) {
return nil, false, nil
}
filePath := filepath.Join(service.slotAssetsDir(slotID), fileName)
info, err := os.Lstat(filePath)
if err != nil {
if os.IsNotExist(err) {
return nil, false, nil
}
return nil, false, err
}
if !info.Mode().IsRegular() || info.Mode()&os.ModeSymlink != 0 {
return nil, false, nil
}
data, err := os.ReadFile(filePath)
if err != nil {
if os.IsNotExist(err) {
return nil, false, nil
}
return nil, false, err
}
if !assetFileMatchesData(assetPath, fileName, data) {
return nil, false, nil
}
return data, true, nil
}
func isSafeAssetFileName(fileName string) bool {
if fileName == "" || fileName == "." || fileName == ".." {
return false
}
if filepath.IsAbs(fileName) || path.IsAbs(fileName) {
return false
}
if strings.Contains(fileName, "/") || strings.Contains(fileName, "\\") || strings.Contains(fileName, "..") {
return false
}
return true
}
func assetFileMatchesData(assetPath string, fileName string, data []byte) bool {
sum := sha256.Sum256(data)
expectedPrefix := hex.EncodeToString(sum[:])
expectedExtension := strings.TrimPrefix(strings.ToLower(path.Ext(assetPath)), ".")
if expectedExtension == "" {
return strings.EqualFold(fileName, expectedPrefix)
}
actualPrefix, extension, ok := strings.Cut(fileName, ".")
if !ok || !strings.EqualFold(actualPrefix, expectedPrefix) {
return false
}
return strings.EqualFold(extension, expectedExtension)
}
func validSlotStorageHash(slotID string, hash string) bool {
hash = strings.TrimSpace(hash)
expectedPrefix := normalizeSlotID(slotID) + ":"
if !strings.HasPrefix(hash, expectedPrefix) {
return false
}
value := strings.TrimPrefix(hash, expectedPrefix)
if len(value) != sha256.Size*2 {
return false
}
_, err := hex.DecodeString(value)
return err == nil
}
func (service *Service) replacePackage(ctx context.Context, slotID string, parsed parsedPackage) error {
_ = ctx
now := time.Now().UTC()
slotDir := service.slotDir(slotID)
tempDir := slotDir + ".tmp-" + strconv.FormatInt(now.UnixNano(), 10)
if err := os.MkdirAll(filepath.Join(tempDir, "assets"), 0o755); err != nil {
return err
}
cleanup := true
defer func() {
if cleanup {
_ = os.RemoveAll(tempDir)
}
}()
assets := make(map[string]adAssetRef, len(parsed.assets))
for _, asset := range parsed.assets {
fileName := adAssetFileName(asset)
if err := os.WriteFile(filepath.Join(tempDir, "assets", fileName), asset.data, 0o644); err != nil {
return err
}
assets[asset.path] = adAssetRef{
File: fileName,
ContentType: asset.contentType,
}
}
pkg := adPackageFile{
Hash: parsed.hash,
ConfigJSON: parsed.configJSON,
Assets: assets,
FetchedAt: now,
ActivatedAt: now,
}
if err := writeJSONFile(filepath.Join(tempDir, "package.json"), pkg); err != nil {
return err
}
backupDir := slotDir + ".old-" + strconv.FormatInt(now.UnixNano(), 10)
hadExisting := false
if _, err := os.Stat(slotDir); err == nil {
hadExisting = true
if err := os.Rename(slotDir, backupDir); err != nil {
return err
}
} else if !os.IsNotExist(err) {
return err
}
if err := os.Rename(tempDir, slotDir); err != nil {
if hadExisting {
_ = os.Rename(backupDir, slotDir)
}
return err
}
cleanup = false
if hadExisting {
_ = os.RemoveAll(backupDir)
}
return nil
}
func (service *Service) clearPackage(ctx context.Context, slotID string) (bool, error) {
_ = ctx
slotDir := service.slotDir(slotID)
if _, err := os.Stat(slotDir); err != nil {
if os.IsNotExist(err) {
return false, nil
}
return false, err
}
return true, os.RemoveAll(slotDir)
}
func (service *Service) assetURL(slotID string, assetPath string) string {
base := strings.TrimRight(service.currentAssetBaseURL(), "/")
if base == "" {
return ""
}
slotID = normalizeSlotID(slotID)
cleaned := strings.TrimLeft(strings.TrimSpace(assetPath), "/")
if cleaned == "" {
return base + "/" + slotID
}
return base + "/" + slotID + "/" + cleaned
}
func (service *Service) slotDir(slotID string) string {
root := strings.TrimSpace(service.storeRoot)
if root == "" {
root = "ads"
}
return filepath.Join(root, normalizeSlotID(slotID))
}
func (service *Service) slotAssetsDir(slotID string) string {
return filepath.Join(service.slotDir(slotID), "assets")
}
func (service *Service) slotPackagePath(slotID string) string {
return filepath.Join(service.slotDir(slotID), "package.json")
}
func adAssetFileName(asset parsedAsset) string {
sum := sha256.Sum256(asset.data)
extension := strings.ToLower(path.Ext(asset.path))
return hex.EncodeToString(sum[:]) + extension
}
func writeJSONFile(filePath string, payload any) error {
if err := os.MkdirAll(filepath.Dir(filePath), 0o755); err != nil {
return err
}
data, err := json.MarshalIndent(payload, "", " ")
if err != nil {
return err
}
return os.WriteFile(filePath, append(data, '\n'), 0o644)
}
func parsePackage(hash string, data []byte) (parsedPackage, error) {
reader, err := zip.NewReader(bytes.NewReader(data), int64(len(data)))
if err != nil {
return parsedPackage{}, fmt.Errorf("解析广告压缩包失败: %w", err)
}
seen := map[string]struct{}{}
var rawConfig []byte
var assets []parsedAsset
for _, file := range reader.File {
if file == nil || file.FileInfo().IsDir() {
continue
}
cleanPath, err := normalizeArchivePath(file.Name)
if err != nil {
return parsedPackage{}, err
}
if _, exists := seen[cleanPath]; exists {
return parsedPackage{}, fmt.Errorf("广告压缩包存在重复文件: %s", cleanPath)
}
seen[cleanPath] = struct{}{}
payload, err := readZipFile(file)
if err != nil {
return parsedPackage{}, err
}
if cleanPath == "config.yaml" {
rawConfig = payload
}
assets = append(assets, parsedAsset{
path: cleanPath,
contentType: detectContentType(cleanPath, payload),
data: payload,
})
}
if len(rawConfig) == 0 {
return parsedPackage{}, fmt.Errorf("广告压缩包缺少根目录 config.yaml")
}
if _, ok := seen["index.html"]; !ok {
return parsedPackage{}, fmt.Errorf("广告压缩包缺少根目录 index.html")
}
cfg, err := parseConfig(rawConfig)
if err != nil {
return parsedPackage{}, err
}
configJSON, err := json.Marshal(cfg)
if err != nil {
return parsedPackage{}, err
}
return parsedPackage{
hash: strings.TrimSpace(hash),
rawConfig: rawConfig,
configJSON: string(configJSON),
assets: assets,
}, nil
}
func parseConfig(data []byte) (Config, error) {
var raw adConfigYAML
if err := yaml.Unmarshal(data, &raw); err != nil {
return Config{}, fmt.Errorf("解析广告 config.yaml 失败: %w", err)
}
if raw.Enabled == nil {
return Config{}, fmt.Errorf("广告 config.yaml 缺少 enabled")
}
cfg := Config{
Enabled: *raw.Enabled,
Window: WindowConfig{
Width: raw.Window.Width,
Height: raw.Window.Height,
},
Home: HomePlacementConfig{
Title: strings.TrimSpace(raw.Home.Title),
Subtitle: strings.TrimSpace(raw.Home.Subtitle),
},
}
if cfg.Window.Width <= 0 || cfg.Window.Height <= 0 {
return Config{}, fmt.Errorf("广告 window.width/window.height 必须为正整数")
}
return cfg, nil
}
func normalizeArchivePath(raw string) (string, error) {
name := strings.TrimSpace(raw)
if name == "" {
return "", fmt.Errorf("广告压缩包存在空文件名")
}
if strings.Contains(name, "\\") || strings.HasPrefix(name, "/") {
return "", fmt.Errorf("广告压缩包存在非法路径: %s", raw)
}
cleaned := path.Clean(name)
if cleaned == "." || cleaned == ".." || strings.HasPrefix(cleaned, "../") || strings.Contains(cleaned, "/../") {
return "", fmt.Errorf("广告压缩包存在路径穿越: %s", raw)
}
return cleaned, nil
}
func normalizeRequestAssetPath(raw string) (string, string, error) {
pathValue := strings.TrimSpace(raw)
pathValue = strings.TrimPrefix(pathValue, RoutePrefix)
pathValue = strings.TrimLeft(pathValue, "/")
slotID := normalizeSlotID("")
if head, tail, ok := strings.Cut(pathValue, "/"); ok && isKnownSlotID(head) {
slotID = normalizeSlotID(head)
pathValue = tail
} else if isKnownSlotID(pathValue) {
slotID = normalizeSlotID(pathValue)
pathValue = ""
}
if pathValue == "" {
pathValue = "index.html"
}
assetPath, err := normalizeArchivePath(pathValue)
return slotID, assetPath, err
}
func normalizeSlotID(value string) string {
value = strings.TrimSpace(value)
if value == "" {
return "1"
}
return value
}
func isKnownSlotID(value string) bool {
value = normalizeSlotID(value)
for _, slot := range Slots {
if normalizeSlotID(slot.ID) == value {
return true
}
}
return false
}
func slotStorageHash(slotID string, hash string) string {
return normalizeSlotID(slotID) + ":" + strings.TrimSpace(hash)
}
func stripSlotStorageHash(hash string) string {
_, value, ok := strings.Cut(strings.TrimSpace(hash), ":")
if !ok {
return strings.TrimSpace(hash)
}
return strings.TrimSpace(value)
}
func readZipFile(file *zip.File) ([]byte, error) {
if file.UncompressedSize64 > maxAssetBytes {
return nil, fmt.Errorf("广告资源过大: %s", file.Name)
}
reader, err := file.Open()
if err != nil {
return nil, err
}
defer reader.Close()
return readLimited(reader, maxAssetBytes, "广告资源 "+file.Name)
}
func readLimited(reader io.Reader, limit int64, label string) ([]byte, error) {
limited := io.LimitReader(reader, limit+1)
data, err := io.ReadAll(limited)
if err != nil {
return nil, err
}
if int64(len(data)) > limit {
return nil, fmt.Errorf("%s超过大小限制", strings.TrimSpace(label))
}
return data, nil
}
func detectContentType(assetPath string, data []byte) string {
extension := strings.ToLower(path.Ext(assetPath))
if extension != "" {
if value := mime.TypeByExtension(extension); strings.TrimSpace(value) != "" {
return value
}
}
switch extension {
case ".js", ".mjs":
return "application/javascript; charset=utf-8"
case ".css":
return "text/css; charset=utf-8"
case ".html", ".htm":
return "text/html; charset=utf-8"
}
if len(data) == 0 {
return "application/octet-stream"
}
return http.DetectContentType(data)
}
func headerValue(value string) string {
value = strings.TrimSpace(value)
value = strings.ReplaceAll(value, "\r", " ")
value = strings.ReplaceAll(value, "\n", " ")
return value
}
func firstNonEmpty(values ...string) string {
for _, value := range values {
if strings.TrimSpace(value) != "" {
return strings.TrimSpace(value)
}
}
return ""
}
func maxInt(left int, right int) int {
if left > right {
return left
}
return right
}
+64
View File
@@ -0,0 +1,64 @@
package ads
import (
"os"
"os/exec"
"runtime"
"strings"
"time"
)
func displayOSName() string {
switch runtime.GOOS {
case "darwin":
return "macOS"
case "linux":
return "Linux"
default:
return "Windows"
}
}
func displayOSVersion() string {
switch runtime.GOOS {
case "darwin":
return commandOutput(2*time.Second, "sw_vers", "-productVersion")
case "linux":
if version := linuxPrettyName(); version != "" {
return version
}
return commandOutput(2*time.Second, "uname", "-r")
default:
return ""
}
}
func commandOutput(timeout time.Duration, name string, args ...string) string {
cmd := exec.Command(name, args...)
timer := time.AfterFunc(timeout, func() {
if cmd.Process != nil {
_ = cmd.Process.Kill()
}
})
defer timer.Stop()
output, err := cmd.Output()
if err != nil {
return ""
}
return strings.TrimSpace(string(output))
}
func linuxPrettyName() string {
data, err := os.ReadFile("/etc/os-release")
if err != nil {
return ""
}
for _, line := range strings.Split(string(data), "\n") {
key, value, ok := strings.Cut(line, "=")
if !ok || strings.TrimSpace(key) != "PRETTY_NAME" {
continue
}
return strings.Trim(strings.TrimSpace(value), `"`)
}
return ""
}
+71
View File
@@ -0,0 +1,71 @@
package ads
const (
// FetchURL 保留旧名称表示默认广告位地址。
FetchURL = "https://ads.leokun.cn/1"
// FetchURL = "http://localhost:3000/ad.zip"
RoutePrefix = "/ad"
EventUpdated = "ad:updated"
)
type Slot struct {
ID string
FetchURL string
}
var Slots = []Slot{
{ID: "1", FetchURL: "https://ads.leokun.cn/1"},
{ID: "2", FetchURL: "https://ads.leokun.cn/2"},
{ID: "3", FetchURL: "https://ads.leokun.cn/3"},
}
type WindowConfig struct {
Width int `json:"width" yaml:"width"`
Height int `json:"height" yaml:"height"`
}
type HomePlacementConfig struct {
Title string `json:"title" yaml:"title"`
Subtitle string `json:"subtitle" yaml:"subtitle"`
}
type Config struct {
Enabled bool `json:"enabled" yaml:"enabled"`
Window WindowConfig `json:"window" yaml:"window"`
Home HomePlacementConfig `json:"home" yaml:"home"`
}
type Runtime struct {
Available bool `json:"available"`
Enabled bool `json:"enabled"`
PackageHash string `json:"packageHash"`
AssetBaseURL string `json:"assetBaseURL"`
IndexURL string `json:"indexURL"`
Window WindowConfig `json:"window"`
Home HomePlacementConfig `json:"home"`
Slots []SlotRuntime `json:"slots"`
}
type SlotRuntime struct {
ID string `json:"id"`
Available bool `json:"available"`
Enabled bool `json:"enabled"`
PackageHash string `json:"packageHash"`
AssetBaseURL string `json:"assetBaseURL"`
IndexURL string `json:"indexURL"`
Window WindowConfig `json:"window"`
Home HomePlacementConfig `json:"home"`
}
type MetricsSnapshot struct {
TurnsTotal int
RequestTokensTotal int64
PromptTokensTotal int64
CacheReadTokens int64
CacheWriteTokens int64
}
type FetchResult struct {
Hash string
Changed bool
}
+2
View File
@@ -0,0 +1,2 @@
// Package app 负责组装桌面应用入口、窗口、托盘和 Wails 服务。
package app
+394
View File
@@ -0,0 +1,394 @@
package app
import (
"context"
"crypto/sha256"
"crypto/x509"
"encoding/hex"
"encoding/pem"
"io/fs"
"net"
goruntime "runtime"
"strings"
"time"
"cursor/internal/ads"
"cursor/internal/appdata"
serverconfig "cursor/internal/backend/server/config"
"cursor/internal/buildinfo"
"cursor/internal/cursor"
"cursor/internal/historymetrics"
"github.com/leaanthony/u"
bridge "cursor/internal/bridge"
"cursor/internal/certs"
"cursor/internal/logger"
"cursor/internal/mitm"
"cursor/internal/netproxy"
"cursor/internal/updater"
"github.com/wailsapp/wails/v3/pkg/application"
"github.com/wailsapp/wails/v3/pkg/events"
)
const (
// appName 表示当前模块中的 appName 状态值。
appName = "Cursor助手"
// adRefreshInterval 表示后台广告拉取间隔。
adRefreshInterval = 3 * time.Minute
)
// EmbeddedResources 定义了当前模块中的 EmbeddedResources 类型。
type EmbeddedResources struct {
// Assets 表示当前声明中的 Assets。
Assets fs.FS
// AppIcon 表示当前声明中的 AppIcon。
AppIcon []byte
// TrayIcon 表示当前声明中的 TrayIcon。
TrayIcon []byte
}
// init 用于处理与 init 相关的逻辑。
func init() {
application.RegisterEvent[bridge.ProxyState]("proxy:state")
application.RegisterEvent[bridge.UserConfig]("user-config:changed")
application.RegisterEvent[bridge.ModelAdapterTestResultsPayload]("model-adapter-test:updated")
application.RegisterEvent[bridge.AdRuntime](ads.EventUpdated)
application.RegisterEvent[updater.StatePayload](updater.EventState)
application.RegisterEvent[updater.ProgressPayload](updater.EventProgress)
application.RegisterEvent[updater.ReadyPayload](updater.EventReady)
application.RegisterEvent[updater.ErrorPayload](updater.EventError)
}
// Run 用于处理与 Run 相关的逻辑。
func Run(resources EmbeddedResources) error {
logger.Init()
netproxy.InstallDefaultTransport()
embeddedCACertPEM := certs.EmbeddedCACertPEM()
logEmbeddedCAInfo(embeddedCACertPEM)
certManager, err := certs.NewEmbeddedManager()
if err != nil {
return err
}
defaultBackendBaseURL := "http://" + serverconfig.DefaultBackendListenAddr
proxyServer, err := mitm.NewProxyServer(serverconfig.DefaultProxyListenAddr, defaultBackendBaseURL, "", "", certManager)
if err != nil {
return err
}
proxyService := bridge.NewProxyService(proxyServer, certManager, embeddedCACertPEM)
adAssetBaseURL := defaultBackendBaseURL
if cfg, err := proxyService.LoadUserConfig(); err == nil {
adAssetBaseURL = browserReachableLoopbackBaseURL(cfg.BackendListenAddr)
}
metricsService := bridge.NewMetricsService()
windowService := bridge.NewWindowService()
adCore := ads.NewService(ads.Options{
StoreRoot: appdata.AdsRootPath(),
HTTPClient: netproxy.NewHTTPClient(30 * time.Second),
AppVersion: buildinfo.CurrentVersion(),
AssetBaseURL: adAssetBaseURL + ads.RoutePrefix,
DeviceID: cursor.GetDeviceID,
Metrics: func(context.Context) (ads.MetricsSnapshot, error) {
if err := appdata.EnsureAssistantHome(); err != nil {
return ads.MetricsSnapshot{}, err
}
summary, err := historymetrics.LoadUsageSummary(appdata.UsageFilePath())
if err != nil {
return ads.MetricsSnapshot{}, err
}
return ads.MetricsSnapshot{
TurnsTotal: summary.TurnsTotal,
RequestTokensTotal: summary.RequestTokensTotal,
PromptTokensTotal: summary.PromptTokensTotal,
CacheReadTokens: summary.CacheReadTokens,
CacheWriteTokens: summary.CacheWriteTokens,
}, nil
},
ProviderCount: func(context.Context) (int, error) {
cfg, err := proxyService.LoadUserConfig()
if err != nil {
return 0, err
}
return len(cfg.ModelAdapters), nil
},
})
adService := bridge.NewAdService(adCore)
var updateManager *updater.Manager
var mainWindow *application.WebviewWindow
adRefreshCtx, stopAdRefresh := context.WithCancel(context.Background())
app := application.New(application.Options{
Name: appName,
Description: appName,
Services: []application.Service{
application.NewService(proxyService),
application.NewService(metricsService),
application.NewService(windowService),
application.NewService(adService),
},
Assets: application.AssetOptions{
Handler: application.AssetFileServerFS(resources.Assets),
},
Mac: application.MacOptions{
ActivationPolicy: application.ActivationPolicyAccessory,
ApplicationShouldTerminateAfterLastWindowClosed: false,
},
OnShutdown: func() {
stopAdRefresh()
if updateManager != nil {
updateManager.Shutdown()
}
proxyService.ShutdownForQuit()
},
SingleInstance: &application.SingleInstanceOptions{
UniqueID: "com.cursor-assistant.single-instance",
OnSecondInstanceLaunch: func(data application.SecondInstanceData) {
logger.Infof("检测到实例请求,已忽略")
// 不激活窗口,避免干扰用户工作
},
},
})
refreshAdAssetBaseURL := func() bool {
state := proxyService.GetState()
backendListenAddr := strings.TrimSpace(state.BackendListenAddr)
if backendListenAddr == "" {
backendListenAddr = serverconfig.DefaultBackendListenAddr
}
return adCore.SetAssetBaseURL(browserReachableLoopbackBaseURL(backendListenAddr) + ads.RoutePrefix)
}
refreshAdRuntime := func() {
runtimeState, err := adCore.GetRuntime(context.Background())
if err != nil {
return
}
app.Event.Emit(ads.EventUpdated, runtimeState)
}
refreshAd := func(ctx context.Context) {
if ctx == nil {
ctx = context.Background()
}
runtimeState, changed, err := adCore.Refresh(ctx)
if err != nil || !changed {
return
}
app.Event.Emit(ads.EventUpdated, runtimeState)
}
refreshAdAsync := func() {
go func() {
refreshAd(context.Background())
}()
}
startAdRefreshLoop := func(ctx context.Context) {
go func() {
refreshAd(ctx)
ticker := time.NewTicker(adRefreshInterval)
defer ticker.Stop()
for {
select {
case <-ctx.Done():
return
case <-ticker.C:
refreshAd(ctx)
}
}
}()
}
updateManager = updater.NewManager(app)
windowService.SetApp(app)
windowService.SetUpdater(updateManager)
mainWindow = app.Window.NewWithOptions(application.WebviewWindowOptions{
Title: appName,
Width: 700,
Height: 520,
MinWidth: 640,
MinHeight: 480,
DisableResize: false,
Frameless: goruntime.GOOS == "windows",
URL: "/",
Hidden: false,
HideOnEscape: false,
MinimiseButtonState: application.ButtonEnabled,
MaximiseButtonState: application.ButtonHidden,
CloseButtonState: application.ButtonEnabled,
BackgroundColour: application.RGBA{Red: 25, Green: 25, Blue: 25, Alpha: 255},
Mac: application.MacWindow{
Backdrop: application.MacBackdropLiquidGlass,
DisableShadow: false,
TitleBar: application.MacTitleBar{
AppearsTransparent: true,
Hide: false,
HideTitle: true,
FullSizeContent: true,
UseToolbar: false,
HideToolbarSeparator: true,
},
WebviewPreferences: application.MacWebviewPreferences{
FullscreenEnabled: u.True,
TextInteractionEnabled: u.True,
AllowsBackForwardNavigationGestures: u.False,
},
},
Windows: application.WindowsWindow{
HiddenOnTaskbar: false,
},
})
window := mainWindow
window.RegisterHook(events.Common.WindowClosing, func(e *application.WindowEvent) {
window.Hide()
e.Cancel()
})
window.RegisterHook(events.Common.WindowFocus, func(e *application.WindowEvent) {
refreshAdAsync()
})
showMainWindow := func() {
window.Show().Focus()
}
toggleMainWindow := func() {
if window.IsVisible() {
window.Hide()
return
}
showMainWindow()
}
systray := app.SystemTray.New()
menu := app.Menu.New()
statusItem := menu.Add("状态:未启动").SetEnabled(false)
menu.AddSeparator()
startItem := menu.Add("启动服务")
stopItem := menu.Add("停止服务")
menu.Add("检查更新").OnClick(func(ctx *application.Context) {
updateManager.CheckNow(true)
})
menu.AddSeparator()
menu.Add("显示窗口").OnClick(func(ctx *application.Context) {
showMainWindow()
})
menu.Add("隐藏窗口").OnClick(func(ctx *application.Context) {
window.Hide()
})
menu.AddSeparator()
menu.Add("退出").OnClick(func(ctx *application.Context) {
proxyService.ShutdownForQuit()
app.Quit()
})
refreshTray := func() {
state := proxyService.GetState()
if state.Running {
statusItem.SetLabel("状态:运行中")
startItem.SetEnabled(false)
stopItem.SetEnabled(true)
} else {
statusItem.SetLabel("状态:未启动")
startItem.SetEnabled(true)
stopItem.SetEnabled(false)
}
if refreshAdAssetBaseURL() {
refreshAdRuntime()
}
}
app.Event.On("proxy:state", func(event *application.CustomEvent) {
refreshTray()
})
app.Event.OnApplicationEvent(events.Common.ApplicationStarted, func(event *application.ApplicationEvent) {
logger.Infof("应用版本:v%s", buildinfo.CurrentVersion())
updateManager.Start()
startAdRefreshLoop(adRefreshCtx)
go func() {
logger.Infof("application started, begin auto start service in background")
if _, err := proxyService.StartProxy(); err != nil {
logger.Errorf("自动启动服务失败: %v", err)
} else {
state := proxyService.GetState()
if refreshAdAssetBaseURL() {
refreshAdRuntime()
}
logger.Infof("代理已自动启动: %s", state.ProxyListenAddr)
}
}()
})
startItem.OnClick(func(ctx *application.Context) {
if _, err := proxyService.StartProxy(); err != nil {
logger.Errorf("启动服务失败: %v", err)
} else if refreshAdAssetBaseURL() {
refreshAdRuntime()
}
refreshTray()
})
stopItem.OnClick(func(ctx *application.Context) {
if _, err := proxyService.StopProxy(); err != nil {
logger.Errorf("停止服务失败: %v", err)
}
refreshTray()
})
if len(resources.AppIcon) > 0 {
switch goruntime.GOOS {
case "darwin":
systray.SetTemplateIcon(resources.TrayIcon)
case "windows":
systray.SetIcon(resources.AppIcon)
default:
systray.SetIcon(resources.TrayIcon)
}
}
systray.SetTooltip(appName)
systray.OnClick(toggleMainWindow).SetMenu(menu)
refreshTray()
return app.Run()
}
func browserReachableLoopbackBaseURL(listenAddr string) string {
host, port, err := net.SplitHostPort(strings.TrimSpace(listenAddr))
if err != nil || strings.TrimSpace(port) == "" {
return "http://" + serverconfig.DefaultBackendListenAddr
}
host = strings.TrimSpace(host)
if host == "" || host == "0.0.0.0" || host == "::" || host == "[::]" {
host = "127.0.0.1"
}
return "http://" + net.JoinHostPort(host, port)
}
// logEmbeddedCAInfo 用于处理与 logEmbeddedCAInfo 相关的逻辑。
func logEmbeddedCAInfo(certPEM []byte) {
if len(certPEM) == 0 {
logger.Errorf("embedded CA is empty")
return
}
cert, err := parseEmbeddedCert(certPEM)
if err != nil {
logger.Errorf("parse embedded CA failed: %v", err)
return
}
sum := sha256.Sum256(cert.Raw)
logger.Infof(
"embedded CA loaded: sha256=%s subject=%s valid=%s~%s",
strings.ToUpper(hex.EncodeToString(sum[:])),
cert.Subject.String(),
cert.NotBefore.Format(time.RFC3339),
cert.NotAfter.Format(time.RFC3339),
)
}
// parseEmbeddedCert 用于处理与 parseEmbeddedCert 相关的逻辑。
func parseEmbeddedCert(data []byte) (*x509.Certificate, error) {
if block, _ := pem.Decode(data); block != nil {
return x509.ParseCertificate(block.Bytes)
}
return x509.ParseCertificate(data)
}
+83
View File
@@ -0,0 +1,83 @@
package appdata
import (
"fmt"
"io"
"os"
"path/filepath"
)
func ensureAssistantHome() error {
migrateLegacyAssistantHome()
if err := os.MkdirAll(RootDir(), 0o755); err != nil {
return fmt.Errorf("create assistant home: %w", err)
}
if err := os.MkdirAll(DataRootPath(), 0o755); err != nil {
return fmt.Errorf("create data root: %w", err)
}
if err := os.MkdirAll(HistoryRootPath(), 0o755); err != nil {
return fmt.Errorf("create history root: %w", err)
}
if err := os.MkdirAll(RulesRootPath(), 0o755); err != nil {
return fmt.Errorf("create rules root: %w", err)
}
if err := os.MkdirAll(LogsRootPath(), 0o755); err != nil {
return fmt.Errorf("create logs root: %w", err)
}
return nil
}
func EnsureAssistantHome() error {
return ensureAssistantHome()
}
func migrateLegacyAssistantHome() {
legacyRoot := legacyRootDir()
copyLegacyFile(filepath.Join(legacyRoot, "config.yaml"), filepath.Join(RootDir(), "config.yaml"))
copyLegacyRules(filepath.Join(legacyRoot, "rules"), RulesRootPath())
_ = os.RemoveAll(legacyRoot)
}
func copyLegacyRules(sourceRoot string, targetRoot string) {
_ = filepath.Walk(sourceRoot, func(path string, info os.FileInfo, err error) error {
if err != nil || info == nil {
return nil
}
rel, err := filepath.Rel(sourceRoot, path)
if err != nil {
return nil
}
targetPath := filepath.Join(targetRoot, rel)
if info.IsDir() {
_ = os.MkdirAll(targetPath, info.Mode().Perm())
return nil
}
if !info.Mode().IsRegular() {
return nil
}
copyLegacyFile(path, targetPath)
return nil
})
}
func copyLegacyFile(sourcePath string, targetPath string) {
sourceFile, err := os.Open(sourcePath)
if err != nil {
return
}
defer sourceFile.Close()
info, err := sourceFile.Stat()
if err != nil || !info.Mode().IsRegular() {
return
}
if err := os.MkdirAll(filepath.Dir(targetPath), 0o755); err != nil {
return
}
targetFile, err := os.OpenFile(targetPath, os.O_CREATE|os.O_WRONLY|os.O_TRUNC, info.Mode().Perm())
if err != nil {
return
}
defer targetFile.Close()
_, _ = io.Copy(targetFile, sourceFile)
}
+72
View File
@@ -0,0 +1,72 @@
package appdata
import (
"os"
"path/filepath"
"strings"
)
const (
appDirName = ".cursor-local-assistant-v2"
legacyAppDirName = ".cursor-local-assistant"
)
// RootDir 返回应用配置根目录。
func RootDir() string {
return appRootDir(appDirName)
}
func legacyRootDir() string {
return appRootDir(legacyAppDirName)
}
func appRootDir(dirName string) string {
homeDir, err := os.UserHomeDir()
if err != nil || strings.TrimSpace(homeDir) == "" {
return dirName
}
return filepath.Join(homeDir, dirName)
}
// ConfigFilePath 返回统一用户配置文件路径。
func ConfigFilePath() string {
return filepath.Join(RootDir(), "config.yaml")
}
func DataRootPath() string {
return filepath.Join(RootDir(), "data")
}
func HistoryRootPath() string {
return filepath.Join(RootDir(), "history")
}
func UsageFilePath() string {
return filepath.Join(HistoryRootPath(), "usage.json")
}
func AdsRootPath() string {
return filepath.Join(DataRootPath(), "ads")
}
func CodebaseIndexRootPath() string {
return filepath.Join(DataRootPath(), "codebase-index")
}
func DocsIndexRootPath() string {
return filepath.Join(DataRootPath(), "docs-index")
}
func RulesRootPath() string {
return filepath.Join(RootDir(), "rules")
}
// LogsRootPath 返回统一日志根目录路径。
func LogsRootPath() string {
return filepath.Join(RootDir(), "logs")
}
// CACertFilePath 返回注入给宿主的 CA 文件路径。
func CACertFilePath() string {
return filepath.Join(DataRootPath(), "ca.crt")
}
+149
View File
@@ -0,0 +1,149 @@
# Backend 架构说明
`internal/backend` 当前支持本地助手模式与直连上游模式。
关于 backend agent「最小事实集合」的第一阶段研究文档,见 [`../../docs/backend-agent-minimum-facts-phase1.md`](../../docs/backend-agent-minimum-facts-phase1.md)。
核心边界:
- `server`
- 本地 HTTP/Connect 入口层
- 负责路由、中间件、错误编码和少量本地 mock
- `forwarder`
- 本地协议兼容与 LLM 转发内核
- 负责 `BidiAppend`、`RunSSE`、history JSON、prompt 编译、provider 流式调用和广播
- `host`
- 唯一组装点
- 负责把 `server/config.Manager`、`forwarder.Module` 和根路由装起来
当前实现不再支持:
- Pro / `cursor-byok`
- HTTP/protocol trace debug UI
- DB-backed store、会话索引和 searchable conversation memory
## 目录结构
```text
internal/backend/
README.md
host.go
server/
context.go
errors.go
local.go
middleware.go
policy.go
route.go
url.go
config/
manager.go
store.go
types.go
legacy_runtime.go
resolver.go
upstream/
action.go
mocks.go
types.go
forwarder/
artifacts.go
broker.go
compiler.go
events.go
file_store.go
legacy_stream.go
module.go
projector.go
provider.go
reminders.go
service.go
tool_catalog.go
types.go
agent/
bridge/
exec/
bridge.go
interaction/
bridge.go
core/
types.go
model/
router.go
openai.go
anthropic.go
artifacts.go
provider_limits.go
http_error.go
types.go
prompt/
engine.go
replay.go
protocol/
inbound.go
```
## 持久化布局
助手目录固定为:
- `~/.cursor-local-assistant-v2/config.yaml`
- `~/.cursor-local-assistant-v2/data/ca.crt`
- `~/.cursor-local-assistant-v2/data/ads/`
- `~/.cursor-local-assistant-v2/history/`
- `~/.cursor-local-assistant-v2/logs/`
约定:
- `config.yaml` 是用户配置
- `data/ca.crt` 是注入给宿主的 CA 证书
- `data/ads/` 是广告包与资源缓存目录
- `history/` 是会话事实与全局 usage JSON 目录,不属于日志
- `logs/` 只保留必要文本运行日志
当前 `history/` 目录布局:
```text
history/
usage.json
<conversation_id>/
state.json
context.json
conversation.lock
```
`state.json` 只表达当前 loop 状态与持久化内存,不保存可投射给 LLM 的历史内容。当前 loop status 语义为:
- `idle`:没有正在推进的 loop。
- `running`:已落入本轮输入或中间上下文,正在等待/发起模型推进。
- `waiting_tool`:已落完整 tool call,正在等待工具结果。
- `completed`:本轮已正常完成。
- `canceled`:本轮被取消,不制造 assistant 输出。
- `provider_error`:provider/LLM 调用失败,错误作为 context tag 记录。
- `failed`:本地内部失败,例如投影、持久化、usage JSON 写入或桥接收口失败;它不等同于 provider 错误。
## 请求流
1. 请求进入 backend 根路由。
2. `PolicyMiddleware` 根据 `routing.mode` 与 `X-Server-Upstream-URL` 选择本地或上游分支。
3. `BidiAppend` / `RunSSE` 进入 `forwarder`。
4. `forwarder` 先把当前 loop 状态写入 `state.json`,再把已发生语义事件追加到 `context.json`。
5. 发给 LLM 的 prompt 只由 `context.json` 投射生成;`state.json` 不保存可投射历史。
6. provider usage/cache 与聚合统计写入 `history/usage.json`,不从 conversation 文件现场扫描。
7. `checkpoint` 只表示同一 backend 进程内的 live state。
## 模型渠道
- 用户在配置里填写 `displayName`、`baseURL`、`apiKey`、`modelID`
- 运行时渠道唯一 ID 不再由 `modelID` 决定
- 当前唯一 ID 是 `url + modelID + key + name` 的短 `SHA-256` hash(前 16 个十六进制字符)
- `modelID` 仅表示 provider model
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,2 @@
// Package execbridge 负责把执行型工具调用映射为 ExecServerMessage,并归一化 ExecClientMessage 结果。
package execbridge
@@ -0,0 +1,929 @@
// bridge.go 实现 MVP 阶段的交互桥协议映射。
package interaction
import (
"bytes"
"encoding/json"
"fmt"
"html"
"io"
"mime"
"net"
"net/http"
neturl "net/url"
"regexp"
"strings"
"sync/atomic"
"time"
"unicode/utf8"
readability "codeberg.org/readeck/go-readability/v2"
htmlmarkdown "github.com/firecrawl/html-to-markdown"
mdplugin "github.com/firecrawl/html-to-markdown/plugin"
"cursor/gen/agentv1"
"cursor/internal/backend/agent/core"
"cursor/internal/netproxy"
)
// InteractionApplyResult 表示一次交互桥结果归一化后的最小产物。
type InteractionApplyResult struct {
// ToolCallID 表示该结果所属工具调用标识。
ToolCallID string
// InteractionID 表示该结果所属交互桥标识。
InteractionID string
// IsTerminal 表示交互桥是否已经收口。
IsTerminal bool
// ToolResultPayload 表示可继续喂给模型的结果摘要。
ToolResultPayload string
// ToolCall 保存可用于发 ToolCallCompletedUpdate 的工具调用对象。
ToolCall *agentv1.ToolCall
}
// InteractionBridge 定义交互桥接口。
type InteractionBridge interface {
// OpenQuery 打开一条交互型工具调用。
OpenQuery(toolCall runtimecore.ToolInvocation) (*agentv1.AgentServerMessage, runtimecore.PendingInteraction, error)
// ApplyInteractionResponse 处理交互响应。
ApplyInteractionResponse(msg *agentv1.InteractionResponse, pending runtimecore.PendingInteraction) (InteractionApplyResult, error)
}
// Bridge 实现 MVP 阶段的交互桥。
type Bridge struct {
// nextID 生成交互消息编号。
nextID atomic.Uint32
// httpClient 负责执行 web search / web fetch 等需要外网的操作。
httpClient *http.Client
}
// NewBridge 创建一个交互桥实例。
func NewBridge() *Bridge {
return &Bridge{
httpClient: netproxy.NewHTTPClient(15 * time.Second),
}
}
// OpenQuery 打开一条交互型工具调用。
func (bridge *Bridge) OpenQuery(toolCall runtimecore.ToolInvocation) (*agentv1.AgentServerMessage, runtimecore.PendingInteraction, error) {
switch toolCall.ToolName {
case "AskQuestion":
return bridge.openAskQuestion(toolCall)
case "CreatePlan":
return bridge.openCreatePlan(toolCall)
case "WebSearch":
return bridge.openWebSearch(toolCall)
case "WebFetch":
return bridge.openWebFetch(toolCall)
case "SwitchMode":
return bridge.openSwitchMode(toolCall)
default:
return nil, runtimecore.PendingInteraction{}, fmt.Errorf("unsupported interaction tool: %s", toolCall.ToolName)
}
}
// ApplyInteractionResponse 处理交互响应。
func (bridge *Bridge) ApplyInteractionResponse(msg *agentv1.InteractionResponse, pending runtimecore.PendingInteraction) (InteractionApplyResult, error) {
if msg == nil {
return InteractionApplyResult{}, fmt.Errorf("interaction response is required")
}
result := InteractionApplyResult{
ToolCallID: pending.ToolCallID,
InteractionID: pending.InteractionID,
IsTerminal: true,
}
switch pending.InteractionKind {
case "ask_question":
var args agentv1.AskQuestionArgs
_ = json.Unmarshal(pending.ArgsJSON, &args)
result.ToolResultPayload = summarizeAskQuestionResponse(msg.GetAskQuestionInteractionResponse())
result.ToolCall = &agentv1.ToolCall{
Tool: &agentv1.ToolCall_AskQuestionToolCall{
AskQuestionToolCall: &agentv1.AskQuestionToolCall{
Args: &args,
Result: msg.GetAskQuestionInteractionResponse().GetResult(),
},
},
}
return result, nil
case "create_plan":
args, err := runtimecore.DecodeCreatePlanArgsJSON(pending.ArgsJSON)
if err != nil {
args = &agentv1.CreatePlanArgs{}
}
createPlanResult := normalizeCreatePlanResult(msg.GetCreatePlanRequestResponse())
result.ToolResultPayload = summarizeCreatePlanResult(createPlanResult)
result.ToolCall = &agentv1.ToolCall{
Tool: &agentv1.ToolCall_CreatePlanToolCall{
CreatePlanToolCall: &agentv1.CreatePlanToolCall{
Args: args,
Result: createPlanResult,
},
},
}
return result, nil
case "web_search":
var args agentv1.WebSearchArgs
_ = json.Unmarshal(pending.ArgsJSON, &args)
webSearchResult, payload := bridge.applyWebSearchResponse(msg.GetWebSearchRequestResponse(), &args)
result.ToolResultPayload = payload
result.ToolCall = &agentv1.ToolCall{
Tool: &agentv1.ToolCall_WebSearchToolCall{
WebSearchToolCall: &agentv1.WebSearchToolCall{
Args: &args,
Result: webSearchResult,
},
},
}
return result, nil
case "web_fetch":
var args agentv1.WebFetchArgs
_ = json.Unmarshal(pending.ArgsJSON, &args)
webFetchResult, payload := bridge.applyWebFetchResponse(msg.GetWebFetchRequestResponse(), &args)
result.ToolResultPayload = payload
result.ToolCall = &agentv1.ToolCall{
Tool: &agentv1.ToolCall_WebFetchToolCall{
WebFetchToolCall: &agentv1.WebFetchToolCall{
Args: &args,
Result: webFetchResult,
},
},
}
return result, nil
case "switch_mode":
var args agentv1.SwitchModeArgs
_ = json.Unmarshal(pending.ArgsJSON, &args)
switchModeResult := buildSwitchModeResult(msg.GetSwitchModeRequestResponse(), &args)
result.ToolResultPayload = summarizeSwitchModeResponse(switchModeResult)
result.ToolCall = &agentv1.ToolCall{
Tool: &agentv1.ToolCall_SwitchModeToolCall{
SwitchModeToolCall: &agentv1.SwitchModeToolCall{
Args: &args,
Result: switchModeResult,
},
},
}
return result, nil
default:
return InteractionApplyResult{}, fmt.Errorf("unsupported pending interaction kind: %s", pending.InteractionKind)
}
}
// nextMessageID 返回下一个交互消息编号。
func (bridge *Bridge) nextMessageID() uint32 {
current := bridge.nextID.Add(1)
if current == 0 {
current = bridge.nextID.Add(1)
}
return current
}
// openAskQuestion 构造 AskQuestion 交互查询。
func (bridge *Bridge) openAskQuestion(toolCall runtimecore.ToolInvocation) (*agentv1.AgentServerMessage, runtimecore.PendingInteraction, error) {
var args agentv1.AskQuestionArgs
if err := json.Unmarshal(toolCall.ArgsJSON, &args); err != nil {
return nil, runtimecore.PendingInteraction{}, fmt.Errorf("decode AskQuestion args failed: %w", err)
}
messageID := bridge.nextMessageID()
serverMessage := &agentv1.AgentServerMessage{
Message: &agentv1.AgentServerMessage_InteractionQuery{
InteractionQuery: &agentv1.InteractionQuery{
Id: messageID,
Query: &agentv1.InteractionQuery_AskQuestionInteractionQuery{
AskQuestionInteractionQuery: &agentv1.AskQuestionInteractionQuery{
Args: &args,
ToolCallId: toolCall.CallID,
},
},
},
},
}
return serverMessage, runtimecore.PendingInteraction{
InteractionID: fmt.Sprintf("%d", messageID),
ArgsJSON: append([]byte(nil), toolCall.ArgsJSON...),
ToolCallID: toolCall.CallID,
InteractionKind: "ask_question",
}, nil
}
// openCreatePlan 构造 CreatePlan 交互查询。
func (bridge *Bridge) openCreatePlan(toolCall runtimecore.ToolInvocation) (*agentv1.AgentServerMessage, runtimecore.PendingInteraction, error) {
args, err := runtimecore.DecodeCreatePlanArgsJSON(toolCall.ArgsJSON)
if err != nil {
return nil, runtimecore.PendingInteraction{}, fmt.Errorf("decode CreatePlan args failed: %w", err)
}
messageID := bridge.nextMessageID()
serverMessage := &agentv1.AgentServerMessage{
Message: &agentv1.AgentServerMessage_InteractionQuery{
InteractionQuery: &agentv1.InteractionQuery{
Id: messageID,
Query: &agentv1.InteractionQuery_CreatePlanRequestQuery{
CreatePlanRequestQuery: &agentv1.CreatePlanRequestQuery{
Args: args,
ToolCallId: toolCall.CallID,
},
},
},
},
}
return serverMessage, runtimecore.PendingInteraction{
InteractionID: fmt.Sprintf("%d", messageID),
ArgsJSON: append([]byte(nil), toolCall.ArgsJSON...),
ToolCallID: toolCall.CallID,
InteractionKind: "create_plan",
}, nil
}
// openWebSearch 构造 WebSearch 交互查询。
func (bridge *Bridge) openWebSearch(toolCall runtimecore.ToolInvocation) (*agentv1.AgentServerMessage, runtimecore.PendingInteraction, error) {
var input struct {
SearchTerm string `json:"search_term"`
}
if err := json.Unmarshal(toolCall.ArgsJSON, &input); err != nil {
return nil, runtimecore.PendingInteraction{}, fmt.Errorf("decode WebSearch args failed: %w", err)
}
messageID := bridge.nextMessageID()
serverMessage := &agentv1.AgentServerMessage{
Message: &agentv1.AgentServerMessage_InteractionQuery{
InteractionQuery: &agentv1.InteractionQuery{
Id: messageID,
Query: &agentv1.InteractionQuery_WebSearchRequestQuery{
WebSearchRequestQuery: &agentv1.WebSearchRequestQuery{
Args: &agentv1.WebSearchArgs{
SearchTerm: input.SearchTerm,
ToolCallId: toolCall.CallID,
},
},
},
},
},
}
return serverMessage, runtimecore.PendingInteraction{
InteractionID: fmt.Sprintf("%d", messageID),
ArgsJSON: append([]byte(nil), toolCall.ArgsJSON...),
ToolCallID: toolCall.CallID,
InteractionKind: "web_search",
}, nil
}
// openWebFetch 构造 WebFetch 交互查询。
func (bridge *Bridge) openWebFetch(toolCall runtimecore.ToolInvocation) (*agentv1.AgentServerMessage, runtimecore.PendingInteraction, error) {
var input struct {
URL string `json:"url"`
}
if err := json.Unmarshal(toolCall.ArgsJSON, &input); err != nil {
return nil, runtimecore.PendingInteraction{}, fmt.Errorf("decode WebFetch args failed: %w", err)
}
messageID := bridge.nextMessageID()
serverMessage := &agentv1.AgentServerMessage{
Message: &agentv1.AgentServerMessage_InteractionQuery{
InteractionQuery: &agentv1.InteractionQuery{
Id: messageID,
Query: &agentv1.InteractionQuery_WebFetchRequestQuery{
WebFetchRequestQuery: &agentv1.WebFetchRequestQuery{
Args: &agentv1.WebFetchArgs{
Url: input.URL,
ToolCallId: toolCall.CallID,
},
},
},
},
},
}
argsPayload, _ := json.Marshal(agentv1.WebFetchArgs{
Url: input.URL,
ToolCallId: toolCall.CallID,
})
return serverMessage, runtimecore.PendingInteraction{
InteractionID: fmt.Sprintf("%d", messageID),
ArgsJSON: argsPayload,
ToolCallID: toolCall.CallID,
InteractionKind: "web_fetch",
}, nil
}
// openSwitchMode 构造 SwitchMode 交互查询。
func (bridge *Bridge) openSwitchMode(toolCall runtimecore.ToolInvocation) (*agentv1.AgentServerMessage, runtimecore.PendingInteraction, error) {
var args agentv1.SwitchModeArgs
if err := json.Unmarshal(toolCall.ArgsJSON, &args); err != nil {
return nil, runtimecore.PendingInteraction{}, fmt.Errorf("decode SwitchMode args failed: %w", err)
}
if err := validateSwitchModeTargetID(args.GetTargetModeId()); err != nil {
return nil, runtimecore.PendingInteraction{}, err
}
args.ToolCallId = toolCall.CallID
messageID := bridge.nextMessageID()
serverMessage := &agentv1.AgentServerMessage{
Message: &agentv1.AgentServerMessage_InteractionQuery{
InteractionQuery: &agentv1.InteractionQuery{
Id: messageID,
Query: &agentv1.InteractionQuery_SwitchModeRequestQuery{
SwitchModeRequestQuery: &agentv1.SwitchModeRequestQuery{
Args: &args,
},
},
},
},
}
argsPayload, _ := json.Marshal(args)
return serverMessage, runtimecore.PendingInteraction{
InteractionID: fmt.Sprintf("%d", messageID),
ArgsJSON: argsPayload,
ToolCallID: toolCall.CallID,
InteractionKind: "switch_mode",
}, nil
}
func validateSwitchModeTargetID(raw string) error {
switch strings.ToLower(strings.TrimSpace(raw)) {
case "agent", "ask", "plan":
return nil
default:
return fmt.Errorf("unsupported target mode id: %q", strings.TrimSpace(raw))
}
}
// summarizeAskQuestionResponse 生成 AskQuestion 响应摘要。
func summarizeAskQuestionResponse(response *agentv1.AskQuestionInteractionResponse) string {
if response == nil || response.GetResult() == nil {
return "ask question response missing"
}
switch item := response.GetResult().GetResult().(type) {
case *agentv1.AskQuestionResult_Success:
if len(item.Success.GetAnswers()) == 0 {
return "ask question success"
}
return fmt.Sprintf("ask question answers=%d", len(item.Success.GetAnswers()))
case *agentv1.AskQuestionResult_Error:
return item.Error.GetErrorMessage()
case *agentv1.AskQuestionResult_Rejected:
return item.Rejected.GetReason()
case *agentv1.AskQuestionResult_Async:
return "ask question async accepted"
default:
return "unknown ask question response"
}
}
const createPlanEmptyURIError = "create plan failed: Cursor returned success with empty planUri"
// normalizeCreatePlanResult 兜底客户端 success 但未返回 planUri 的异常形态。
func normalizeCreatePlanResult(response *agentv1.CreatePlanRequestResponse) *agentv1.CreatePlanResult {
if response == nil || response.GetResult() == nil {
return nil
}
result := response.GetResult()
if result.GetSuccess() != nil && strings.TrimSpace(result.GetPlanUri()) == "" {
return &agentv1.CreatePlanResult{
Result: &agentv1.CreatePlanResult_Error{
Error: &agentv1.CreatePlanError{Error: createPlanEmptyURIError},
},
}
}
return result
}
// summarizeCreatePlanResult 生成 CreatePlan 响应摘要。
func summarizeCreatePlanResult(result *agentv1.CreatePlanResult) string {
if result == nil {
return "create plan response missing"
}
switch item := result.GetResult().(type) {
case *agentv1.CreatePlanResult_Success:
return fmt.Sprintf("create plan success uri=%s", result.GetPlanUri())
case *agentv1.CreatePlanResult_Error:
return item.Error.GetError()
default:
return "unknown create plan response"
}
}
// summarizeWebSearchResponse 生成 WebSearch 响应摘要。
func summarizeWebSearchResponse(response *agentv1.WebSearchRequestResponse) string {
if response == nil {
return "web search response missing"
}
switch item := response.GetResult().(type) {
case *agentv1.WebSearchRequestResponse_Approved_:
_ = item
return "web search approved"
case *agentv1.WebSearchRequestResponse_Rejected_:
return item.Rejected.GetReason()
default:
return "unknown web search response"
}
}
// applyWebSearchResponse 把 WebSearch approval 响应转换成最终工具结果。
func (bridge *Bridge) applyWebSearchResponse(response *agentv1.WebSearchRequestResponse, args *agentv1.WebSearchArgs) (*agentv1.WebSearchResult, string) {
if response == nil {
return &agentv1.WebSearchResult{
Result: &agentv1.WebSearchResult_Error{
Error: &agentv1.WebSearchError{Error: "web search response missing"},
},
}, "web search response missing"
}
switch item := response.GetResult().(type) {
case *agentv1.WebSearchRequestResponse_Approved_:
_ = item
references, payload, err := bridge.executeWebSearch(strings.TrimSpace(args.GetSearchTerm()))
if err != nil {
return &agentv1.WebSearchResult{
Result: &agentv1.WebSearchResult_Error{
Error: &agentv1.WebSearchError{Error: err.Error()},
},
}, err.Error()
}
references, payload = truncateWebSearchReplay(strings.TrimSpace(args.GetSearchTerm()), references, payload)
return &agentv1.WebSearchResult{
Result: &agentv1.WebSearchResult_Success{
Success: &agentv1.WebSearchSuccess{References: references},
},
}, payload
case *agentv1.WebSearchRequestResponse_Rejected_:
return &agentv1.WebSearchResult{
Result: &agentv1.WebSearchResult_Rejected{
Rejected: &agentv1.WebSearchRejected{Reason: item.Rejected.GetReason()},
},
}, item.Rejected.GetReason()
default:
return &agentv1.WebSearchResult{
Result: &agentv1.WebSearchResult_Error{
Error: &agentv1.WebSearchError{Error: "unknown web search response"},
},
}, "unknown web search response"
}
}
// applyWebFetchResponse 把 WebFetch approval 响应转换成最终工具结果。
func (bridge *Bridge) applyWebFetchResponse(response *agentv1.WebFetchRequestResponse, args *agentv1.WebFetchArgs) (*agentv1.WebFetchResult, string) {
if response == nil {
return &agentv1.WebFetchResult{
Result: &agentv1.WebFetchResult_Error{
Error: &agentv1.WebFetchError{
Url: args.GetUrl(),
Error: "web fetch response missing",
},
},
}, "web fetch response missing"
}
switch item := response.GetResult().(type) {
case *agentv1.WebFetchRequestResponse_Approved_:
_ = item
markdown, err := bridge.executeWebFetch(strings.TrimSpace(args.GetUrl()))
if err != nil {
return &agentv1.WebFetchResult{
Result: &agentv1.WebFetchResult_Error{
Error: &agentv1.WebFetchError{
Url: args.GetUrl(),
Error: err.Error(),
},
},
}, err.Error()
}
return &agentv1.WebFetchResult{
Result: &agentv1.WebFetchResult_Success{
Success: &agentv1.WebFetchSuccess{
Url: args.GetUrl(),
Markdown: markdown,
},
},
}, markdown
case *agentv1.WebFetchRequestResponse_Rejected_:
return &agentv1.WebFetchResult{
Result: &agentv1.WebFetchResult_Rejected{
Rejected: &agentv1.WebFetchRejected{Reason: item.Rejected.GetReason()},
},
}, item.Rejected.GetReason()
default:
return &agentv1.WebFetchResult{
Result: &agentv1.WebFetchResult_Error{
Error: &agentv1.WebFetchError{
Url: args.GetUrl(),
Error: "unknown web fetch response",
},
},
}, "unknown web fetch response"
}
}
// buildSwitchModeResult 把 SwitchMode approval 响应转换成最终工具结果。
func buildSwitchModeResult(response *agentv1.SwitchModeRequestResponse, args *agentv1.SwitchModeArgs) *agentv1.SwitchModeResult {
if response == nil {
return &agentv1.SwitchModeResult{
Result: &agentv1.SwitchModeResult_Error{
Error: &agentv1.SwitchModeError{Error: "switch mode response missing"},
},
}
}
switch item := response.GetResult().(type) {
case *agentv1.SwitchModeRequestResponse_Approved_:
_ = item
targetModeID := strings.ToLower(strings.TrimSpace(args.GetTargetModeId()))
return &agentv1.SwitchModeResult{
Result: &agentv1.SwitchModeResult_Success{
Success: &agentv1.SwitchModeSuccess{
FromModeId: "unknown",
ToModeId: targetModeID,
},
},
}
case *agentv1.SwitchModeRequestResponse_Rejected_:
return &agentv1.SwitchModeResult{
Result: &agentv1.SwitchModeResult_Rejected{
Rejected: &agentv1.SwitchModeRejected{Reason: item.Rejected.GetReason()},
},
}
default:
return &agentv1.SwitchModeResult{
Result: &agentv1.SwitchModeResult_Error{
Error: &agentv1.SwitchModeError{Error: "unknown switch mode response"},
},
}
}
}
// summarizeSwitchModeResponse 生成 SwitchMode 响应摘要。
func summarizeSwitchModeResponse(result *agentv1.SwitchModeResult) string {
if result == nil {
return "switch mode result missing"
}
switch item := result.GetResult().(type) {
case *agentv1.SwitchModeResult_Success:
return fmt.Sprintf("switch mode success to=%s", item.Success.GetToModeId())
case *agentv1.SwitchModeResult_Rejected:
return item.Rejected.GetReason()
case *agentv1.SwitchModeResult_Error:
return item.Error.GetError()
default:
return "unknown switch mode result"
}
}
var (
webSearchAnchorPattern = regexp.MustCompile(`(?is)<a[^>]*class="[^"]*result__a[^"]*"[^>]*href="([^"]+)"[^>]*>(.*?)</a>`)
webSearchSnippetPattern = regexp.MustCompile(`(?is)<(?:a|div)[^>]*class="[^"]*result__snippet[^"]*"[^>]*>(.*?)</(?:a|div)>`)
htmlTitlePattern = regexp.MustCompile(`(?is)<title[^>]*>(.*?)</title>`)
htmlTagPattern = regexp.MustCompile(`(?is)<[^>]+>`)
webSearchURLOverride = "https://html.duckduckgo.com/html/?q="
)
const (
webFetchBodyLimit = 2 * 1024 * 1024
webFetchMarkdownLimit = 32 * 1024
webSearchPayloadLimit = 16 * 1024
webSearchTitleLimit = 512
webSearchChunkLimit = 2 * 1024
)
func (bridge *Bridge) executeWebSearch(searchTerm string) ([]*agentv1.WebSearchReference, string, error) {
if strings.TrimSpace(searchTerm) == "" {
return nil, "", fmt.Errorf("web search search_term is required")
}
client := bridge.httpClient
if client == nil {
client = netproxy.NewHTTPClient(15 * time.Second)
}
requestURL := webSearchURLOverride + neturl.QueryEscape(searchTerm)
request, err := http.NewRequest(http.MethodGet, requestURL, nil)
if err != nil {
return nil, "", err
}
request.Header.Set("User-Agent", "cursor-local-agent/1.0")
response, err := client.Do(request)
if err != nil {
return nil, "", err
}
defer response.Body.Close()
if response.StatusCode < 200 || response.StatusCode >= 300 {
return nil, "", fmt.Errorf("web search http status %d", response.StatusCode)
}
body, err := io.ReadAll(io.LimitReader(response.Body, 2*1024*1024))
if err != nil {
return nil, "", err
}
references := extractWebSearchReferences(string(body))
if len(references) == 0 {
return nil, "", fmt.Errorf("web search returned no parseable results")
}
if len(references) > 5 {
references = references[:5]
}
return references, formatWebSearchPayload(searchTerm, references), nil
}
func extractWebSearchReferences(body string) []*agentv1.WebSearchReference {
anchorMatches := webSearchAnchorPattern.FindAllStringSubmatch(body, 8)
snippetMatches := webSearchSnippetPattern.FindAllStringSubmatch(body, 8)
references := make([]*agentv1.WebSearchReference, 0, len(anchorMatches))
for index, match := range anchorMatches {
if len(match) < 3 {
continue
}
title := cleanupWebSearchHTML(match[2])
url := strings.TrimSpace(html.UnescapeString(match[1]))
snippet := ""
if index < len(snippetMatches) && len(snippetMatches[index]) >= 2 {
snippet = cleanupWebSearchHTML(snippetMatches[index][1])
}
if title == "" || url == "" {
continue
}
references = append(references, &agentv1.WebSearchReference{
Title: title,
Url: url,
Chunk: snippet,
})
}
return references
}
func cleanupWebSearchHTML(value string) string {
withoutTags := htmlTagPattern.ReplaceAllString(value, " ")
unescaped := html.UnescapeString(withoutTags)
return strings.Join(strings.Fields(unescaped), " ")
}
func formatWebSearchPayload(searchTerm string, references []*agentv1.WebSearchReference) string {
lines := []string{
fmt.Sprintf("Title: Web search results for query: %s", strings.TrimSpace(searchTerm)),
"Content: Links:",
}
for index, reference := range references {
if reference == nil {
continue
}
lines = append(lines, fmt.Sprintf("%d. [%s](%s)", index+1, strings.TrimSpace(reference.GetTitle()), strings.TrimSpace(reference.GetUrl())))
}
snippets := make([]string, 0, len(references))
for _, reference := range references {
if reference == nil {
continue
}
chunk := strings.TrimSpace(reference.GetChunk())
if chunk == "" {
continue
}
snippets = append(snippets, fmt.Sprintf("- %s", chunk))
}
if len(snippets) > 0 {
lines = append(lines, "", strings.Join(snippets, "\n"))
}
return strings.Join(lines, "\n")
}
func truncateWebSearchReplay(searchTerm string, references []*agentv1.WebSearchReference, payload string) ([]*agentv1.WebSearchReference, string) {
truncated := false
nextReferences := make([]*agentv1.WebSearchReference, 0, len(references))
for _, reference := range references {
if reference == nil {
continue
}
next := *reference
title := truncateInteractionText("WebSearch title", next.GetTitle(), webSearchTitleLimit)
chunk := truncateInteractionText("WebSearch snippet", next.GetChunk(), webSearchChunkLimit)
if title != next.GetTitle() || chunk != next.GetChunk() {
truncated = true
}
next.Title = title
next.Chunk = chunk
nextReferences = append(nextReferences, &next)
}
nextPayload := formatWebSearchPayload(searchTerm, nextReferences)
if strings.TrimSpace(payload) != "" && len(nextPayload) == 0 {
nextPayload = payload
}
if len(nextPayload) > webSearchPayloadLimit {
truncated = true
nextPayload = truncateInteractionText("WebSearch", nextPayload, webSearchPayloadLimit)
}
if truncated && len(nextReferences) > 0 {
last := nextReferences[len(nextReferences)-1]
last.Chunk = strings.TrimSpace(last.GetChunk() + "\n\n" + interactionTruncationNotice("WebSearch", webSearchPayloadLimit, len(nextPayload), len(payload)))
nextPayload = formatWebSearchPayload(searchTerm, nextReferences)
nextPayload = truncateInteractionText("WebSearch", nextPayload, webSearchPayloadLimit)
}
return nextReferences, nextPayload
}
func (bridge *Bridge) executeWebFetch(rawURL string) (string, error) {
parsedURL, err := validateWebFetchURL(rawURL)
if err != nil {
return "", err
}
client := bridge.httpClient
if client == nil {
client = netproxy.NewHTTPClient(15 * time.Second)
}
client = webFetchHTTPClient(client)
request, err := http.NewRequest(http.MethodGet, parsedURL.String(), nil)
if err != nil {
return "", err
}
request.Header.Set("User-Agent", "cursor-local-agent/1.0")
request.Header.Set("Accept", "text/html,application/xhtml+xml,text/plain,application/xml,application/json;q=0.9,*/*;q=0.1")
response, err := client.Do(request)
if err != nil {
return "", err
}
defer response.Body.Close()
if response.StatusCode < 200 || response.StatusCode >= 300 {
return "", fmt.Errorf("web fetch http status %d", response.StatusCode)
}
body, err := io.ReadAll(io.LimitReader(response.Body, webFetchBodyLimit+1))
if err != nil {
return "", err
}
if len(body) == 0 {
return "", fmt.Errorf("web fetch returned empty body")
}
if len(body) > webFetchBodyLimit {
body = body[:webFetchBodyLimit]
}
contentType := response.Header.Get("Content-Type")
if contentType == "" {
contentType = http.DetectContentType(body)
}
if !isWebFetchTextContentType(contentType) {
return "", fmt.Errorf("web fetch unsupported content type %q", contentType)
}
markdown, title, err := renderWebFetchMarkdown(parsedURL, body, contentType)
if err != nil {
return "", err
}
markdown = strings.TrimSpace(markdown)
if markdown == "" {
return "", fmt.Errorf("web fetch returned empty markdown")
}
title = strings.TrimSpace(title)
if title == "" {
title = parsedURL.String()
}
payload := fmt.Sprintf("Title: %s\nURL: %s\n\nContent:\n%s", title, parsedURL.String(), markdown)
return truncateWebFetchMarkdown(payload), nil
}
func validateWebFetchURL(rawURL string) (*neturl.URL, error) {
rawURL = strings.TrimSpace(rawURL)
if rawURL == "" {
return nil, fmt.Errorf("web fetch url is required")
}
parsedURL, err := neturl.Parse(rawURL)
if err != nil {
return nil, fmt.Errorf("web fetch invalid url: %w", err)
}
switch strings.ToLower(parsedURL.Scheme) {
case "http", "https":
default:
return nil, fmt.Errorf("web fetch only supports http and https urls")
}
host := strings.TrimSpace(parsedURL.Hostname())
if host == "" {
return nil, fmt.Errorf("web fetch url host is required")
}
if isBlockedWebFetchHost(host) {
return nil, fmt.Errorf("web fetch host is not public-web accessible")
}
return parsedURL, nil
}
func isBlockedWebFetchHost(host string) bool {
host = strings.Trim(strings.ToLower(host), "[]")
if host == "localhost" || strings.HasSuffix(host, ".localhost") {
return true
}
ip := net.ParseIP(host)
if ip == nil {
return false
}
return ip.IsLoopback() ||
ip.IsPrivate() ||
ip.IsLinkLocalUnicast() ||
ip.IsLinkLocalMulticast() ||
ip.IsUnspecified()
}
func isWebFetchTextContentType(contentType string) bool {
mediaType, _, err := mime.ParseMediaType(contentType)
if err != nil {
mediaType = strings.ToLower(strings.TrimSpace(strings.Split(contentType, ";")[0]))
}
if strings.HasPrefix(mediaType, "text/") {
return true
}
switch mediaType {
case "application/xhtml+xml", "application/xml", "application/json", "application/ld+json", "application/rss+xml", "application/atom+xml":
return true
default:
return strings.HasSuffix(mediaType, "+xml") || strings.HasSuffix(mediaType, "+json")
}
}
func renderWebFetchMarkdown(pageURL *neturl.URL, body []byte, contentType string) (string, string, error) {
if !isHTMLLikeContentType(contentType) {
return string(body), "", nil
}
article, err := readability.FromReader(bytes.NewReader(body), pageURL)
if err == nil {
var articleHTML bytes.Buffer
if renderErr := article.RenderHTML(&articleHTML); renderErr == nil && strings.TrimSpace(articleHTML.String()) != "" {
if markdown, convertErr := convertHTMLToMarkdown(pageURL, articleHTML.String()); convertErr == nil && strings.TrimSpace(markdown) != "" {
return markdown, article.Title(), nil
}
}
}
markdown, err := convertHTMLToMarkdown(pageURL, string(body))
if err != nil {
return "", "", fmt.Errorf("web fetch markdown conversion failed: %w", err)
}
return markdown, extractWebFetchHTMLTitle(string(body)), nil
}
func isHTMLLikeContentType(contentType string) bool {
mediaType, _, err := mime.ParseMediaType(contentType)
if err != nil {
mediaType = strings.ToLower(strings.TrimSpace(strings.Split(contentType, ";")[0]))
}
return mediaType == "text/html" || mediaType == "application/xhtml+xml" || mediaType == ""
}
func convertHTMLToMarkdown(pageURL *neturl.URL, htmlBody string) (string, error) {
converter := htmlmarkdown.NewConverter(htmlmarkdown.DomainFromURL(pageURL.String()), true, nil)
converter.Use(mdplugin.GitHubFlavored())
return converter.ConvertString(htmlBody)
}
func extractWebFetchHTMLTitle(htmlBody string) string {
matches := htmlTitlePattern.FindStringSubmatch(htmlBody)
if len(matches) < 2 {
return ""
}
return cleanupWebSearchHTML(matches[1])
}
func truncateWebFetchMarkdown(markdown string) string {
return truncateInteractionText("WebFetch", markdown, webFetchMarkdownLimit)
}
func truncateInteractionText(toolName string, text string, limit int) string {
if limit <= 0 || len(text) <= limit {
return text
}
original := len(text)
notice := fmt.Sprintf("\n\n%s", interactionTruncationNotice(toolName, limit, limit, original))
for {
keep := limit - len(notice)
if keep <= 0 {
return truncateInteractionUTF8(text, limit)
}
kept := truncateInteractionUTF8(text, keep)
nextNotice := fmt.Sprintf("\n\n%s", interactionTruncationNotice(toolName, limit, len(kept), original))
output := strings.TrimRight(kept, "\n") + nextNotice
if len(output) <= limit || nextNotice == notice {
return output
}
notice = nextNotice
}
}
func interactionTruncationNotice(toolName string, limit int, kept int, original int) string {
return fmt.Sprintf("[truncated: %s result exceeded %d bytes; showing %d of %d bytes]", toolName, limit, kept, original)
}
func truncateInteractionUTF8(text string, limit int) string {
if limit <= 0 {
return ""
}
if len(text) <= limit {
return text
}
if limit > len(text) {
limit = len(text)
}
truncated := text[:limit]
for !utf8.ValidString(truncated) && len(truncated) > 0 {
truncated = truncated[:len(truncated)-1]
}
return truncated
}
func webFetchHTTPClient(base *http.Client) *http.Client {
if base == nil {
base = netproxy.NewHTTPClient(15 * time.Second)
}
client := *base
previousCheckRedirect := client.CheckRedirect
client.CheckRedirect = func(request *http.Request, via []*http.Request) error {
if len(via) >= 10 {
return fmt.Errorf("web fetch stopped after 10 redirects")
}
if _, err := validateWebFetchURL(request.URL.String()); err != nil {
return err
}
if previousCheckRedirect != nil {
return previousCheckRedirect(request, via)
}
return nil
}
return &client
}
@@ -0,0 +1,2 @@
// Package interaction 负责把交互型工具调用映射为 InteractionQuery,并归一化 InteractionResponse。
package interaction
+237
View File
@@ -0,0 +1,237 @@
package runtimecore
import (
"encoding/json"
"fmt"
"strconv"
"strings"
"cursor/gen/agentv1"
)
// DecodeCreatePlanArgsJSON 解析 CreatePlan 参数,并兼容字符串形式的 todo status。
func DecodeCreatePlanArgsJSON(raw []byte) (*agentv1.CreatePlanArgs, error) {
if len(strings.TrimSpace(string(raw))) == 0 {
return &agentv1.CreatePlanArgs{}, nil
}
var direct agentv1.CreatePlanArgs
if err := json.Unmarshal(raw, &direct); err == nil {
return &direct, nil
}
var payload map[string]any
if err := json.Unmarshal(raw, &payload); err != nil {
return nil, err
}
if payload == nil {
return &agentv1.CreatePlanArgs{}, nil
}
todos, err := decodeCreatePlanTodoItems(createPlanValueByAlias(payload, "todos"))
if err != nil {
return nil, fmt.Errorf("decode todos: %w", err)
}
phases, err := decodeCreatePlanPhases(createPlanValueByAlias(payload, "phases"))
if err != nil {
return nil, fmt.Errorf("decode phases: %w", err)
}
return &agentv1.CreatePlanArgs{
Plan: createPlanStringValue(createPlanValueByAlias(payload, "plan")),
Overview: createPlanStringValue(createPlanValueByAlias(payload, "overview")),
Name: strings.TrimSpace(createPlanStringValue(createPlanValueByAlias(payload, "name"))),
IsProject: createPlanBoolValue(createPlanValueByAlias(payload, "is_project", "isProject")),
Todos: todos,
Phases: phases,
}, nil
}
func decodeCreatePlanPhases(value any) ([]*agentv1.Phase, error) {
if value == nil {
return nil, nil
}
items, ok := value.([]any)
if !ok {
return nil, fmt.Errorf("phases must be an array")
}
if len(items) == 0 {
return nil, nil
}
phases := make([]*agentv1.Phase, 0, len(items))
for index, item := range items {
object, ok := item.(map[string]any)
if !ok {
return nil, fmt.Errorf("phase %d must be an object", index)
}
todos, err := decodeCreatePlanTodoItems(createPlanValueByAlias(object, "todos"))
if err != nil {
return nil, fmt.Errorf("phase %d todos: %w", index, err)
}
phases = append(phases, &agentv1.Phase{
Name: strings.TrimSpace(createPlanStringValue(createPlanValueByAlias(object, "name"))),
Todos: todos,
})
}
return phases, nil
}
func decodeCreatePlanTodoItems(value any) ([]*agentv1.TodoItem, error) {
if value == nil {
return nil, nil
}
items, ok := value.([]any)
if !ok {
return nil, fmt.Errorf("todos must be an array")
}
if len(items) == 0 {
return nil, nil
}
todos := make([]*agentv1.TodoItem, 0, len(items))
for index, item := range items {
object, ok := item.(map[string]any)
if !ok {
return nil, fmt.Errorf("todo %d must be an object", index)
}
status, err := decodeCreatePlanTodoStatus(createPlanValueByAlias(object, "status"))
if err != nil {
return nil, fmt.Errorf("todo %d status: %w", index, err)
}
todos = append(todos, &agentv1.TodoItem{
Id: strings.TrimSpace(createPlanStringValue(createPlanValueByAlias(object, "id"))),
Content: strings.TrimSpace(createPlanStringValue(createPlanValueByAlias(object, "content"))),
Status: status,
CreatedAt: createPlanInt64Value(createPlanValueByAlias(object, "created_at", "createdAt")),
UpdatedAt: createPlanInt64Value(createPlanValueByAlias(object, "updated_at", "updatedAt")),
Dependencies: createPlanStringSliceValue(createPlanValueByAlias(object, "dependencies")),
})
}
return todos, nil
}
func decodeCreatePlanTodoStatus(value any) (agentv1.TodoStatus, error) {
switch item := value.(type) {
case nil:
return agentv1.TodoStatus_TODO_STATUS_UNSPECIFIED, nil
case float64:
return agentv1.TodoStatus(int32(item)), nil
case float32:
return agentv1.TodoStatus(int32(item)), nil
case int:
return agentv1.TodoStatus(item), nil
case int32:
return agentv1.TodoStatus(item), nil
case int64:
return agentv1.TodoStatus(item), nil
case string:
return decodeCreatePlanTodoStatusString(item)
default:
return agentv1.TodoStatus_TODO_STATUS_UNSPECIFIED, fmt.Errorf("unsupported todo status type %T", value)
}
}
func decodeCreatePlanTodoStatusString(raw string) (agentv1.TodoStatus, error) {
normalized := strings.ToLower(strings.TrimSpace(raw))
if normalized == "" || normalized == "unspecified" || normalized == "todo_status_unspecified" {
return agentv1.TodoStatus_TODO_STATUS_UNSPECIFIED, nil
}
if numeric, err := strconv.ParseInt(normalized, 10, 32); err == nil {
return agentv1.TodoStatus(numeric), nil
}
switch normalized {
case "pending", "todo_status_pending":
return agentv1.TodoStatus_TODO_STATUS_PENDING, nil
case "in_progress", "in-progress", "inprogress", "todo_status_in_progress":
return agentv1.TodoStatus_TODO_STATUS_IN_PROGRESS, nil
case "completed", "complete", "todo_status_completed":
return agentv1.TodoStatus_TODO_STATUS_COMPLETED, nil
case "cancelled", "canceled", "todo_status_cancelled":
return agentv1.TodoStatus_TODO_STATUS_CANCELLED, nil
default:
return agentv1.TodoStatus_TODO_STATUS_UNSPECIFIED, fmt.Errorf("unsupported todo status %q", raw)
}
}
func createPlanValueByAlias(payload map[string]any, aliases ...string) any {
for _, alias := range aliases {
if value, ok := payload[alias]; ok {
return value
}
}
return nil
}
func createPlanStringValue(value any) string {
text, ok := value.(string)
if !ok {
return ""
}
return text
}
func createPlanBoolValue(value any) bool {
switch item := value.(type) {
case bool:
return item
case string:
return strings.EqualFold(strings.TrimSpace(item), "true")
default:
return false
}
}
func createPlanInt64Value(value any) int64 {
switch item := value.(type) {
case float64:
return int64(item)
case float32:
return int64(item)
case int:
return int64(item)
case int32:
return int64(item)
case int64:
return item
case string:
parsed, err := strconv.ParseInt(strings.TrimSpace(item), 10, 64)
if err != nil {
return 0
}
return parsed
default:
return 0
}
}
func createPlanStringSliceValue(value any) []string {
switch item := value.(type) {
case []string:
if len(item) == 0 {
return nil
}
result := make([]string, 0, len(item))
for _, text := range item {
trimmed := strings.TrimSpace(text)
if trimmed == "" {
continue
}
result = append(result, trimmed)
}
return result
case []any:
if len(item) == 0 {
return nil
}
result := make([]string, 0, len(item))
for _, entry := range item {
text := strings.TrimSpace(createPlanStringValue(entry))
if text == "" {
continue
}
result = append(result, text)
}
return result
default:
return nil
}
}
+2
View File
@@ -0,0 +1,2 @@
// Package runtimecore 定义 runtime/loop、checkpoint、session 之间共享的状态与事件模型。
package runtimecore
+110
View File
@@ -0,0 +1,110 @@
package runtimecore
import (
"encoding/json"
"strings"
)
// MCPToolPayload 表示 CallMcpTool 的宽容解码结果。
type MCPToolPayload struct {
Server string
ProviderIdentifier string
ToolName string
Name string
Arguments map[string]any
}
// DecodeMCPToolPayload 解析 CallMcpTool 参数,并兼容字符串化的 arguments 对象。
func DecodeMCPToolPayload(raw []byte) (MCPToolPayload, error) {
payload := MCPToolPayload{
Arguments: make(map[string]any),
}
if len(raw) == 0 {
return payload, nil
}
var decoded map[string]any
if err := json.Unmarshal(raw, &decoded); err != nil {
return payload, err
}
if decoded == nil {
return payload, nil
}
payload.Server = decodeJSONStringValue(decoded["server"])
payload.ProviderIdentifier = decodeJSONStringValue(decoded["providerIdentifier"])
payload.ToolName = decodeJSONStringValue(decoded["toolName"])
payload.Name = decodeJSONStringValue(decoded["name"])
payload.Arguments = decodeJSONObjectLike(decoded["arguments"])
if len(payload.Arguments) == 0 {
payload.Arguments = decodeJSONObjectLike(decoded["args"])
}
return payload, nil
}
// InferMCPServerIdentifier 从 canonical lookup name 中反推出 server identifier。
func InferMCPServerIdentifier(name string) string {
trimmed := strings.TrimSpace(name)
if trimmed == "" {
return ""
}
if index := strings.Index(trimmed, "-"); index > 0 {
return strings.TrimSpace(trimmed[:index])
}
return ""
}
// InferMCPToolName 从 canonical lookup name 中反推出 tool name。
func InferMCPToolName(serverIdentifier string, name string) string {
trimmedName := strings.TrimSpace(name)
if trimmedName == "" {
return ""
}
trimmedServer := strings.TrimSpace(serverIdentifier)
if trimmedServer != "" && strings.HasPrefix(trimmedName, trimmedServer+"-") {
return strings.TrimSpace(strings.TrimPrefix(trimmedName, trimmedServer+"-"))
}
return trimmedName
}
func decodeJSONStringValue(value any) string {
text, ok := value.(string)
if !ok {
return ""
}
return strings.TrimSpace(text)
}
func decodeJSONObjectLike(value any) map[string]any {
switch item := value.(type) {
case map[string]any:
if item == nil {
return make(map[string]any)
}
return item
case string:
return decodeJSONObjectBytes([]byte(item))
case []byte:
return decodeJSONObjectBytes(item)
case json.RawMessage:
return decodeJSONObjectBytes([]byte(item))
default:
return make(map[string]any)
}
}
func decodeJSONObjectBytes(raw []byte) map[string]any {
trimmed := strings.TrimSpace(string(raw))
if trimmed == "" {
return make(map[string]any)
}
var decoded any
if err := json.Unmarshal([]byte(trimmed), &decoded); err != nil {
return make(map[string]any)
}
object, ok := decoded.(map[string]any)
if !ok || object == nil {
return make(map[string]any)
}
return object
}
+370
View File
@@ -0,0 +1,370 @@
package runtimecore
import (
"bytes"
"encoding/json"
"fmt"
"io"
"math"
"reflect"
"strconv"
"strings"
)
// DecodeArgsMap decodes model-produced built-in tool arguments while preserving
// JSON number spellings for lossless numeric coercion by the typed readers below.
func DecodeArgsMap(raw []byte) (map[string]any, error) {
if len(bytes.TrimSpace(raw)) == 0 {
return map[string]any{}, nil
}
decoder := json.NewDecoder(bytes.NewReader(raw))
decoder.UseNumber()
var result map[string]any
if err := decoder.Decode(&result); err != nil {
return nil, err
}
if result == nil {
return map[string]any{}, nil
}
var trailing any
if err := decoder.Decode(&trailing); err != io.EOF {
if err == nil {
return nil, fmt.Errorf("invalid JSON arguments: multiple top-level values")
}
return nil, err
}
return result, nil
}
// ReadStringArg reads the first string value matching one of the provided keys.
func ReadStringArg(args map[string]any, keys ...string) string {
for _, key := range keys {
value, ok := args[key]
if !ok || value == nil {
continue
}
if text, ok := value.(string); ok {
return text
}
}
return ""
}
// ReadBoolArg reads the first bool value matching one of the provided keys.
func ReadBoolArg(args map[string]any, keys ...string) bool {
for _, key := range keys {
value, ok := args[key]
if !ok || value == nil {
continue
}
if item, ok := value.(bool); ok {
return item
}
}
return false
}
// HasArgKey reports whether any candidate key is present with a non-null value.
func HasArgKey(args map[string]any, keys ...string) bool {
for _, key := range keys {
value, ok := args[key]
if ok && value != nil {
return true
}
}
return false
}
// BoolPtrIfPresent returns a bool pointer only when a matching key is present.
func BoolPtrIfPresent(args map[string]any, keys ...string) *bool {
if !HasArgKey(args, keys...) {
return nil
}
value := ReadBoolArg(args, keys...)
return &value
}
// ReadStringSliceArg reads a string array value matching one of the provided keys.
func ReadStringSliceArg(args map[string]any, keys ...string) []string {
for _, key := range keys {
value, ok := args[key]
if !ok || value == nil {
continue
}
if direct, ok := value.([]string); ok {
return append([]string(nil), direct...)
}
items, ok := value.([]any)
if !ok {
continue
}
result := make([]string, 0, len(items))
for _, item := range items {
text, ok := item.(string)
if !ok {
continue
}
trimmed := strings.TrimSpace(text)
if trimmed != "" {
result = append(result, trimmed)
}
}
return result
}
return nil
}
func readArgValue(args map[string]any, keys ...string) (any, string, bool) {
for _, key := range keys {
value, ok := args[key]
if !ok || value == nil {
continue
}
return value, key, true
}
return nil, "", false
}
// ReadIntArg reads an int field from a JSON number or lossless numeric string.
func ReadIntArg(args map[string]any, keys ...string) (int, bool, error) {
value, key, found := readArgValue(args, keys...)
if !found {
return 0, false, nil
}
parsed, err := parseIntegerValue(value, key, int64MinForBits(strconv.IntSize), int64MaxForBits(strconv.IntSize))
if err != nil {
return 0, true, err
}
return int(parsed), true, nil
}
// ReadInt32Arg reads an int32 field from a JSON number or lossless numeric string.
func ReadInt32Arg(args map[string]any, keys ...string) (int32, bool, error) {
value, key, found := readArgValue(args, keys...)
if !found {
return 0, false, nil
}
parsed, err := parseIntegerValue(value, key, math.MinInt32, math.MaxInt32)
if err != nil {
return 0, true, err
}
return int32(parsed), true, nil
}
// ReadInt64Arg reads an int64 field from a JSON number or lossless numeric string.
func ReadInt64Arg(args map[string]any, keys ...string) (int64, bool, error) {
value, key, found := readArgValue(args, keys...)
if !found {
return 0, false, nil
}
parsed, err := parseIntegerValue(value, key, math.MinInt64, math.MaxInt64)
if err != nil {
return 0, true, err
}
return parsed, true, nil
}
// ReadUint32Arg reads a uint32 field from a JSON number or lossless numeric string.
func ReadUint32Arg(args map[string]any, keys ...string) (uint32, bool, error) {
value, key, found := readArgValue(args, keys...)
if !found {
return 0, false, nil
}
parsed, err := parseUnsignedIntegerValue(value, key, math.MaxUint32)
if err != nil {
return 0, true, err
}
return uint32(parsed), true, nil
}
// ReadFloat64Arg reads a float64 field from a JSON number or numeric string.
func ReadFloat64Arg(args map[string]any, keys ...string) (float64, bool, error) {
value, key, found := readArgValue(args, keys...)
if !found {
return 0, false, nil
}
parsed, err := parseFloatValue(value, key)
if err != nil {
return 0, true, err
}
return parsed, true, nil
}
func parseIntegerValue(value any, key string, minValue int64, maxValue int64) (int64, error) {
switch item := value.(type) {
case json.Number:
return parseIntegerLiteral(item.String(), key, minValue, maxValue)
case string:
return parseIntegerLiteral(strings.TrimSpace(item), key, minValue, maxValue)
case float64:
return parseIntegerFloat(item, key, minValue, maxValue)
case float32:
return parseIntegerFloat(float64(item), key, minValue, maxValue)
default:
reflected := reflect.ValueOf(value)
switch reflected.Kind() {
case reflect.Int, reflect.Int8, reflect.Int16, reflect.Int32, reflect.Int64:
parsed := reflected.Int()
if parsed < minValue || parsed > maxValue {
return 0, fmt.Errorf("%s is outside supported integer range", key)
}
return parsed, nil
case reflect.Uint, reflect.Uint8, reflect.Uint16, reflect.Uint32, reflect.Uint64:
parsed := reflected.Uint()
if parsed > uint64(maxValue) {
return 0, fmt.Errorf("%s is outside supported integer range", key)
}
return int64(parsed), nil
default:
return 0, fmt.Errorf("%s must be an integer", key)
}
}
}
func parseUnsignedIntegerValue(value any, key string, maxValue uint64) (uint64, error) {
switch item := value.(type) {
case json.Number:
return parseUnsignedIntegerLiteral(item.String(), key, maxValue)
case string:
return parseUnsignedIntegerLiteral(strings.TrimSpace(item), key, maxValue)
case float64:
return parseUnsignedIntegerFloat(item, key, maxValue)
case float32:
return parseUnsignedIntegerFloat(float64(item), key, maxValue)
default:
reflected := reflect.ValueOf(value)
switch reflected.Kind() {
case reflect.Int, reflect.Int8, reflect.Int16, reflect.Int32, reflect.Int64:
parsed := reflected.Int()
if parsed < 0 {
return 0, fmt.Errorf("%s must be a non-negative integer", key)
}
if uint64(parsed) > maxValue {
return 0, fmt.Errorf("%s is outside supported unsigned integer range", key)
}
return uint64(parsed), nil
case reflect.Uint, reflect.Uint8, reflect.Uint16, reflect.Uint32, reflect.Uint64:
parsed := reflected.Uint()
if parsed > maxValue {
return 0, fmt.Errorf("%s is outside supported unsigned integer range", key)
}
return parsed, nil
default:
return 0, fmt.Errorf("%s must be a non-negative integer", key)
}
}
}
func parseIntegerLiteral(raw string, key string, minValue int64, maxValue int64) (int64, error) {
if raw == "" {
return 0, fmt.Errorf("%s must be an integer", key)
}
parsed, err := strconv.ParseInt(raw, 10, 64)
if err != nil {
return 0, fmt.Errorf("%s must be an integer", key)
}
if parsed < minValue || parsed > maxValue {
return 0, fmt.Errorf("%s is outside supported integer range", key)
}
return parsed, nil
}
func parseUnsignedIntegerLiteral(raw string, key string, maxValue uint64) (uint64, error) {
if raw == "" {
return 0, fmt.Errorf("%s must be a non-negative integer", key)
}
parsed, err := strconv.ParseUint(raw, 10, 64)
if err != nil {
if strings.HasPrefix(raw, "-") {
return 0, fmt.Errorf("%s must be a non-negative integer", key)
}
return 0, fmt.Errorf("%s must be a non-negative integer", key)
}
if parsed > maxValue {
return 0, fmt.Errorf("%s is outside supported unsigned integer range", key)
}
return parsed, nil
}
func parseIntegerFloat(value float64, key string, minValue int64, maxValue int64) (int64, error) {
if !isFiniteFloat(value) || math.Trunc(value) != value {
return 0, fmt.Errorf("%s must be an integer", key)
}
if value < float64(minValue) || value > float64(maxValue) {
return 0, fmt.Errorf("%s is outside supported integer range", key)
}
return int64(value), nil
}
func parseUnsignedIntegerFloat(value float64, key string, maxValue uint64) (uint64, error) {
if !isFiniteFloat(value) || math.Trunc(value) != value {
return 0, fmt.Errorf("%s must be a non-negative integer", key)
}
if value < 0 {
return 0, fmt.Errorf("%s must be a non-negative integer", key)
}
if value > float64(maxValue) {
return 0, fmt.Errorf("%s is outside supported unsigned integer range", key)
}
return uint64(value), nil
}
func parseFloatValue(value any, key string) (float64, error) {
switch item := value.(type) {
case json.Number:
parsed, err := item.Float64()
if err != nil || !isFiniteFloat(parsed) {
return 0, fmt.Errorf("%s must be a finite number", key)
}
return parsed, nil
case string:
trimmed := strings.TrimSpace(item)
if trimmed == "" {
return 0, fmt.Errorf("%s must be a finite number", key)
}
parsed, err := strconv.ParseFloat(trimmed, 64)
if err != nil || !isFiniteFloat(parsed) {
return 0, fmt.Errorf("%s must be a finite number", key)
}
return parsed, nil
case float64:
if !isFiniteFloat(item) {
return 0, fmt.Errorf("%s must be a finite number", key)
}
return item, nil
case float32:
parsed := float64(item)
if !isFiniteFloat(parsed) {
return 0, fmt.Errorf("%s must be a finite number", key)
}
return parsed, nil
default:
reflected := reflect.ValueOf(value)
switch reflected.Kind() {
case reflect.Int, reflect.Int8, reflect.Int16, reflect.Int32, reflect.Int64:
return float64(reflected.Int()), nil
case reflect.Uint, reflect.Uint8, reflect.Uint16, reflect.Uint32, reflect.Uint64:
return float64(reflected.Uint()), nil
default:
return 0, fmt.Errorf("%s must be a finite number", key)
}
}
}
func int64MinForBits(bits int) int64 {
if bits >= 64 {
return math.MinInt64
}
return -(int64(1) << (bits - 1))
}
func int64MaxForBits(bits int) int64 {
if bits >= 64 {
return math.MaxInt64
}
return (int64(1) << (bits - 1)) - 1
}
func isFiniteFloat(value float64) bool {
return !math.IsNaN(value) && !math.IsInf(value, 0)
}
+399
View File
@@ -0,0 +1,399 @@
// types.go 定义运行时、公用命令、事件、状态与 pending 结构。
package runtimecore
import (
"encoding/json"
"fmt"
"strings"
"time"
"google.golang.org/protobuf/proto"
"cursor/gen/agentv1"
)
// SubagentModelOverrideSelection 表示父 run 对某类 subagent 的模型选择覆盖。
type SubagentModelOverrideSelection struct {
SubagentType string `json:"subagent_type"`
Selection string `json:"selection"`
ModelID string `json:"model_id,omitempty"`
MaxMode bool `json:"max_mode,omitempty"`
ParameterCount int `json:"parameter_count,omitempty"`
BuiltInModel bool `json:"built_in_model,omitempty"`
IsVariantStringRepresentation bool `json:"is_variant_string_representation,omitempty"`
}
// LookupSubagentModelOverride 按 Task subagent_type 查找运行期模型覆盖。
func LookupSubagentModelOverride(overrides map[string]SubagentModelOverrideSelection, subagentType string) (SubagentModelOverrideSelection, string, bool) {
if len(overrides) == 0 {
return SubagentModelOverrideSelection{}, "", false
}
for _, key := range subagentModelOverrideLookupKeys(subagentType) {
if selection, ok := overrides[key]; ok {
return selection, key, true
}
}
return SubagentModelOverrideSelection{}, "", false
}
func subagentModelOverrideLookupKeys(subagentType string) []string {
trimmed := strings.TrimSpace(subagentType)
if trimmed == "" {
return nil
}
keys := []string{trimmed}
switch trimmed {
case "generalPurpose":
keys = append(keys, "explore")
case "explore":
keys = append(keys, "generalPurpose")
case "browserUse":
keys = append(keys, "browser-use")
case "browser-use":
keys = append(keys, "browserUse")
}
return keys
}
type RunState string
const (
// RunStateIdle 表示空闲态,此时 session 已存在但没有活跃 run。
RunStateIdle RunState = "IDLE"
// RunStateRestoring 表示恢复态,此时正在装载会话状态与最小恢复信息。
RunStateRestoring RunState = "RESTORING"
// RunStatePreparingModelInput 表示模型输入准备态。
RunStatePreparingModelInput RunState = "PREPARING_MODEL_INPUT"
// RunStateStreamingModel 表示模型流消费态。
RunStateStreamingModel RunState = "STREAMING_MODEL"
// RunStateWaitingExec 表示执行桥等待态。
RunStateWaitingExec RunState = "WAITING_EXEC"
// RunStateWaitingInteraction 表示交互桥等待态。
RunStateWaitingInteraction RunState = "WAITING_INTERACTION"
// RunStateApplyingExternalResult 表示外部结果回写态。
RunStateApplyingExternalResult RunState = "APPLYING_EXTERNAL_RESULT"
// RunStateCheckpointing 表示检查点写入态。
RunStateCheckpointing RunState = "CHECKPOINTING"
// RunStateCompleted 表示正常完成态。
RunStateCompleted RunState = "COMPLETED"
// RunStateCanceled 表示取消终态。
RunStateCanceled RunState = "CANCELED"
// RunStateFailed 表示失败终态。
RunStateFailed RunState = "FAILED"
)
// CommandKind 表示运行时接收的上行命令类型。
type CommandKind string
const (
// CommandKindRunRequested 表示收到 `run_request`。
CommandKindRunRequested CommandKind = "run_requested"
// CommandKindPrewarmRequested 表示收到 `prewarm_request`。
CommandKindPrewarmRequested CommandKind = "prewarm_requested"
// CommandKindCancelRequested 表示收到 `conversation_action.cancel_action`。
CommandKindCancelRequested CommandKind = "cancel_requested"
// CommandKindConversationActionRecordOnly 表示收到非取消型的 `conversation_action`,当前阶段只记录不推进状态。
CommandKindConversationActionRecordOnly CommandKind = "conversation_action_record_only"
// CommandKindExecClientMessage 表示收到 `exec_client_message`。
CommandKindExecClientMessage CommandKind = "exec_client_message"
// CommandKindInteractionResponse 表示收到 `interaction_response`。
CommandKindInteractionResponse CommandKind = "interaction_response"
// CommandKindExecClientControlMessage 表示收到 `exec_client_control_message`,当前阶段只记录不推进状态。
CommandKindExecClientControlMessage CommandKind = "exec_client_control_message"
// CommandKindClientHeartbeat 表示收到客户端心跳,当前阶段只记录不推进状态。
CommandKindClientHeartbeat CommandKind = "client_heartbeat"
// CommandKindKVClientMessage 表示收到 `kv_client_message`,当前阶段只记录不推进状态。
CommandKindKVClientMessage CommandKind = "kv_client_message"
)
// Command 描述一次投递到运行时协调层的上行命令。
type Command struct {
// Kind 指定该命令的运行时语义。
Kind CommandKind
// IsResume 标记当前命令是否为恢复型启动。
IsResume bool
// ClientKind 保留协议层顶级消息种类,便于观测与调试。
ClientKind string
// HistoryEntry 保存协议摘要文本,供当前 MVP 的合成回复使用。
HistoryEntry string
// ClientMessage 保存解码后的完整上行协议消息。
ClientMessage *agentv1.AgentClientMessage
}
// EventKind 表示一次可回放下行事件的业务类型。
type EventKind string
const (
// EventKindRunStarted 表示新 run 已创建并开始进入恢复路径。
EventKindRunStarted EventKind = "run_started"
// EventKindStepStarted 表示步骤开始事件。
EventKindStepStarted EventKind = "step_started"
// EventKindTextDelta 表示文本增量事件。
EventKindTextDelta EventKind = "text_delta"
// EventKindStepCompleted 表示步骤完成事件。
EventKindStepCompleted EventKind = "step_completed"
// EventKindTurnEnded 表示回合结束事件。
EventKindTurnEnded EventKind = "turn_ended"
// EventKindCheckpoint 表示会话检查点事件。
EventKindCheckpoint EventKind = "checkpoint"
// EventKindCanceled 表示取消事件。
EventKindCanceled EventKind = "canceled"
// EventKindHeartbeat 表示服务端心跳事件。
EventKindHeartbeat EventKind = "heartbeat"
)
// Event 表示一条可广播、可回放的下行事件记录。
type Event struct {
// Seq 是请求维度内递增的事件序号。
Seq int64
// RequestID 是事件所属请求标识。
RequestID string
// RunID 是事件所属运行标识。
RunID string
// Kind 标识该事件的业务类型。
Kind EventKind
// Message 是要透传到 RunSSE 的协议消息体。
Message *agentv1.AgentServerMessage
// End 表示该事件会结束当前 SSE 读取。
End bool
// TerminalErrorCode 表示当前终态 SSE 需要返回的 connect error code,例如 canceled。
TerminalErrorCode string
// TerminalErrorMessage 表示当前终态 SSE 需要返回的错误消息。
TerminalErrorMessage string
// CreatedAt 是事件入库时间。
CreatedAt time.Time
}
// RunSnapshot 表示一次 run 的最小快照信息。
type RunSnapshot struct {
// RunID 是运行唯一标识。
RunID string
// RequestID 是当前 run 绑定的请求标识。
RequestID string
// ConversationID 是当前 run 绑定的会话标识。
ConversationID string
// ModelID 表示当前运行使用的模型标识。
ModelID string
// State 表示该 run 当前所处状态。
State RunState
// Mode 表示该 run 当前使用的会话模式。
Mode agentv1.AgentMode
// Version 是运行时版本号,便于后续扩展乐观更新。
Version int64
// StartedAt 记录 run 启动时间。
StartedAt time.Time
// UpdatedAt 记录 run 最近一次状态更新时间。
UpdatedAt time.Time
// CurrentUserMessageText 保存当前 turn 的用户输入文本,直到本 turn 提交进 `turns`。
CurrentUserMessageText string
// CustomSystemPrompt 保存当前 run 附带的自定义系统提示词。
CustomSystemPrompt string
// RequestContextPayload 保存当前 run 的 request_context proto 序列化结果。
RequestContextPayload []byte
// IsPrewarm 标记当前 run 是否由 `prewarm_request` 触发。
IsPrewarm bool
}
// PendingAssistantOutput 表示尚未收口的一条 assistant 输出记录。
type PendingAssistantOutput struct {
// RawMessage 保存原始序列化 assistant message。
RawMessage string
// Role 表示该记录的 role,当前常见值为 assistant。
Role string
// ContentKinds 记录内容块类型顺序,例如 text 或 tool-call。
ContentKinds []string
// ToolCallIDs 记录该输出中出现的全部 tool_call_id。
ToolCallIDs []string
// ToolNames 记录该输出中出现的全部工具名称。
ToolNames []string
// TextPreview 保存文本块的简要摘要。
TextPreview string
}
// PendingExec 表示一条尚未收口的执行桥记录。
type PendingExec struct {
// MessageID 是打开该执行桥时下发给客户端的桥消息编号。
MessageID uint32
// ExecID 是执行桥唯一标识。
ExecID string
// ProviderPass 表示创建该执行桥时所属的 provider pass。
ProviderPass int
// ModelCallID 是触发该执行桥的模型调用标识。
ModelCallID string
// ToolCallID 是与该执行桥关联的工具调用标识。
ToolCallID string
// ArgsJSON 保存打开该执行桥时的原始参数 JSON,便于恢复 completed ToolCall。
ArgsJSON []byte
// ReasoningContent 保存触发该工具调用时的 thinking 文本,供 checkpoint/replay 续跑复用。
ReasoningContent string
// ReasoningSignature 保存 provider 对当前 thinking 文本签发的签名。
ReasoningSignature string
// ReasoningSignatureSource 保存 reasoning signature 的 provider 语义来源。
ReasoningSignatureSource string
// ExecKind 描述执行桥类型,例如 read、write、shellStream。
ExecKind string
// StreamState 描述当前流式执行桥的阶段。
StreamState string
// OpenedAt 表示执行桥请求发出的时间。
OpenedAt time.Time
// FirstChunkAt 表示 shellStream 首个输出块时间。
FirstChunkAt time.Time
// ChunkCount 表示 shellStream 已接收的输出块数量。
ChunkCount int64
// LastShellActivityAt 记录最近一次 shell 相关上行事件时间,包括输出、start、heartbeat 和 close。
LastShellActivityAt time.Time
// LastShellHeartbeatAt 记录最近一次 shell heartbeat 到达时间。
LastShellHeartbeatAt time.Time
// ShellForegroundDeadline 表示前台 shell 预计最晚应收到终态的时间点。
ShellForegroundDeadline time.Time
// ShellRecoveryScheduled 标记是否已经为该 shell 安排了异常收口协程。
ShellRecoveryScheduled bool
// StdoutBuffer 保存当前 shell 已累计的 stdout 文本。
StdoutBuffer string
// StderrBuffer 保存当前 shell 已累计的 stderr 文本。
StderrBuffer string
// ArtifactPath 保存该 exec 对应的原始桥接工件路径。
ArtifactPath string
}
// PendingInteraction 表示一条尚未收口的交互桥记录。
type PendingInteraction struct {
// InteractionID 是交互桥唯一标识。
InteractionID string
// ProviderPass 表示创建该交互桥时所属的 provider pass。
ProviderPass int
// ModelCallID 是触发该交互桥的模型调用标识。
ModelCallID string
// ToolCallID 是与该交互桥关联的工具调用标识。
ToolCallID string
// ArgsJSON 保存打开该交互桥时的原始参数 JSON,便于结果回写时恢复结构化状态。
ArgsJSON []byte
// ReasoningContent 保存触发该工具调用时的 thinking 文本,供 checkpoint/replay 续跑复用。
ReasoningContent string
// ReasoningSignature 保存 provider 对当前 thinking 文本签发的签名。
ReasoningSignature string
// ReasoningSignatureSource 保存 reasoning signature 的 provider 语义来源。
ReasoningSignatureSource string
// InteractionKind 描述交互类型,例如 ask_question、create_plan。
InteractionKind string
// OpenedAt 表示交互请求发出的时间。
OpenedAt time.Time
// ArtifactPath 保存该 interaction 对应的原始桥接工件路径。
ArtifactPath string
}
// ActiveStep 表示当前正在推进、尚未收口的 step 元数据。
type ActiveStep struct {
// StepID 是当前 step 唯一标识。
StepID uint64
// ModelCallID 是当前 step 绑定的模型调用标识。
ModelCallID string
// StartedAt 是当前 step 的开始时间。
StartedAt time.Time
// InputTokens 保存当前 step 已知的输入 token 数。
InputTokens int64
// OutputTokens 保存当前 step 已知的输出 token 数。
OutputTokens int64
}
// ExternalResultSummary 表示 APPLYING_EXTERNAL_RESULT 后继续下一轮编译所需的最小上下文。
type ExternalResultSummary struct {
// Source 表示结果来源,例如 exec 或 interaction。
Source string
// ToolName 表示对应工具名或交互名。
ToolName string
// Payload 表示可直接注入 prompt 的结果摘要。
Payload string
}
// ToolInvocation 表示一次模型产出的工具调用意图。
type ToolInvocation struct {
// CallID 是模型层工具调用标识。
CallID string
// ToolName 表示工具名称,例如 Read、Write、AskQuestion。
ToolName string
// ArgsJSON 保存工具参数原始 JSON。
ArgsJSON []byte
// ReasoningContent 保存当前工具调用前伴随的 thinking 文本。
ReasoningContent string
// ReasoningSignature 保存 provider 对当前 thinking 文本签发的签名。
ReasoningSignature string
// ReasoningSignatureSource 保存 reasoning signature 的 provider 语义来源。
ReasoningSignatureSource string
// ReasoningProviderItemID 保存 provider 原始 reasoning output item id。
ReasoningProviderItemID string
// ReasoningProviderStatus 保存 provider 原始 reasoning output item status。
ReasoningProviderStatus string
// ReasoningProviderSummary 保存 provider 原始 reasoning output item summary。
ReasoningProviderSummary json.RawMessage
// ProviderItemID 保存 provider 原始 tool/function output item id。
ProviderItemID string
// ProviderCallID 保存 provider 原始 tool/function call id。
ProviderCallID string
// ProviderStatus 保存 provider 原始 tool/function output item status。
ProviderStatus string
// ModelCallID 表示本轮模型调用标识。
ModelCallID string
}
// NormalizeSupportedMode 规范化并校验当前支持的会话 mode。
//
// 当前默认口径:
// 1. 未显式携带 mode 或值为 `AGENT_MODE_UNSPECIFIED` 时,按 `AGENT_MODE_AGENT` 处理;
// 2. 仅允许 `AGENT_MODE_AGENT`、`AGENT_MODE_ASK`、`AGENT_MODE_PLAN`、`AGENT_MODE_DEBUG`、`AGENT_MODE_MULTITASK`;
// 3. 其他 mode 一律报错,不允许静默回退。
func NormalizeSupportedMode(mode agentv1.AgentMode) (agentv1.AgentMode, error) {
switch mode {
case agentv1.AgentMode_AGENT_MODE_UNSPECIFIED:
return agentv1.AgentMode_AGENT_MODE_AGENT, nil
case agentv1.AgentMode_AGENT_MODE_AGENT,
agentv1.AgentMode_AGENT_MODE_ASK,
agentv1.AgentMode_AGENT_MODE_PLAN,
agentv1.AgentMode_AGENT_MODE_DEBUG,
agentv1.AgentMode_AGENT_MODE_MULTITASK:
return mode, nil
default:
return agentv1.AgentMode_AGENT_MODE_UNSPECIFIED, fmt.Errorf("unsupported mode: %s", mode.String())
}
}
// CloneToolCallMap 深拷贝 tool_call 结果映射,避免共享 proto 指针。
func CloneToolCallMap(items map[string]*agentv1.ToolCall) map[string]*agentv1.ToolCall {
if len(items) == 0 {
return make(map[string]*agentv1.ToolCall)
}
cloned := make(map[string]*agentv1.ToolCall, len(items))
for key, value := range items {
if value == nil {
cloned[key] = nil
continue
}
typed, ok := proto.Clone(value).(*agentv1.ToolCall)
if !ok {
cloned[key] = nil
continue
}
cloned[key] = typed
}
return cloned
}
// IsCurrentlySupportedTool 判断当前 Phase 5 稳定化版本是否真正支持该工具。
//
// 当前规则:
// 1. 只返回 runtime/loop 当前已经具备完整推进链路的能力;
// 2. 结果用于限制实际对模型暴露的工具集合,避免模型调用未实现能力后把整轮 run 直接打失败;
// 3. 必须保持最小闭环优先,而不是优先暴露抓包里存在但服务端尚未支持的能力。
func IsCurrentlySupportedTool(name string) bool {
switch strings.TrimSpace(name) {
case "Read", "Write", "PatchEdit", "Delete", "Shell", "AwaitShell", "WriteShellStdin", "ForceBackgroundShell",
"Glob", "Grep", "ReadLints",
"AskQuestion", "CreatePlan", "SwitchMode", "WebSearch", "WebFetch",
"TodoWrite", "Task",
"CallMcpTool", "FetchMcpResource":
return true
default:
return false
}
}
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,48 @@
package modeladapter
import (
"encoding/json"
"fmt"
"testing"
)
func TestAnthropicMessageCacheBreakpointsPreserveAppendOnlyHistory(t *testing.T) {
for size := 1; size <= 32; size++ {
t.Run(fmt.Sprintf("size_%02d", size), func(t *testing.T) {
previous := anthropicMessagesForAppendOnlyTest(size)
current := anthropicMessagesForAppendOnlyTest(size + 1)
applyAnthropicMessageCacheBreakpoints(previous)
applyAnthropicMessageCacheBreakpoints(current)
want := mustMarshalAnthropicMessagesForTest(t, previous)
got := mustMarshalAnthropicMessagesForTest(t, current[:len(previous)])
if got != want {
t.Fatalf("historical message prefix changed after append\nwant: %s\ngot: %s", want, got)
}
})
}
}
func anthropicMessagesForAppendOnlyTest(count int) []anthropicMessage {
messages := make([]anthropicMessage, 0, count)
for index := 0; index < count; index++ {
messages = append(messages, anthropicMessage{
Role: "user",
Content: []map[string]any{{
"type": "text",
"text": fmt.Sprintf("message-%02d", index),
}},
})
}
return messages
}
func mustMarshalAnthropicMessagesForTest(t *testing.T, messages []anthropicMessage) string {
t.Helper()
payload, err := json.Marshal(messages)
if err != nil {
t.Fatalf("marshal anthropic messages: %v", err)
}
return string(payload)
}
+205
View File
@@ -0,0 +1,205 @@
package modeladapter
import (
"encoding/json"
"strings"
"time"
)
// recordLLMRequestArtifact 记录一次模型调用的原始请求工件。
func recordLLMRequestArtifact(req StreamRequest, provider string, model string, method string, url string, body any) {
if req.Observer == nil {
return
}
path, err := req.Observer.RecordLLMRequest(req.RequestID, req.RunID, req.ModelCallID, map[string]any{
"request_id": req.RequestID,
"run_id": req.RunID,
"model_call_id": req.ModelCallID,
"provider": strings.TrimSpace(provider),
"openai_endpoint": strings.TrimSpace(req.OpenAIEndpoint),
"model": firstNonEmptyString(model, req.ModelID),
"runtime_model_id": strings.TrimSpace(req.ModelID),
"resolved_channel_id": strings.TrimSpace(req.ResolvedChannelID),
"resolved_channel_name": strings.TrimSpace(req.ResolvedChannelName),
"resolved_context_window_tokens": req.ResolvedContextWindowTokens,
"url": url,
"method": method,
"body": body,
"request_knobs": req.RequestKnobs,
"compile_summary": req.CompileSummary,
"stable_message_count": req.StableMessageCount,
"tools_summary": summarizeTools(req.Tools),
"messages_summary": summarizeMessages(req.Messages),
})
if err == nil && req.ArtifactPaths != nil {
req.ArtifactPaths.RequestPath = path
}
}
// buildLLMSummaryPayload 生成 LLM 调用摘要工件内容。
func buildLLMSummaryPayload(
req StreamRequest,
provider string,
model string,
startedAt time.Time,
firstEventAt time.Time,
finishedAt time.Time,
finishReason string,
inputTokens int64,
outputTokens int64,
cacheReadTokens int64,
cacheWriteTokens int64,
err error,
) map[string]any {
finished := finishedAt
if finished.IsZero() {
finished = time.Now().UTC()
}
promptTokensTotal := inputTokens + cacheReadTokens + cacheWriteTokens
requestTokensTotal := promptTokensTotal + outputTokens
return map[string]any{
"provider": strings.TrimSpace(provider),
"model": strings.TrimSpace(model),
"started_at": normalizeModelArtifactTime(startedAt),
"first_event_at": normalizeModelArtifactTime(firstEventAt),
"finished_at": normalizeModelArtifactTime(finished),
"finish_reason": strings.TrimSpace(finishReason),
"input_tokens": inputTokens,
"output_tokens": outputTokens,
"cache_read_tokens": cacheReadTokens,
"cache_write_tokens": cacheWriteTokens,
"prompt_tokens_total": promptTokensTotal,
"request_tokens_total": requestTokensTotal,
"error": summarizeModelArtifactError(err),
"ttft_ms": computeTTFTMS(startedAt, firstEventAt),
"duration_ms": computeDurationMS(startedAt, finished),
}
}
// appendLLMResponseArtifact 追加模型调用的原始响应文本。
func appendLLMResponseArtifact(req StreamRequest, chunk string) (string, error) {
if req.Observer == nil {
return "", nil
}
path, err := req.Observer.AppendLLMResponseChunk(req.RequestID, req.RunID, req.ModelCallID, chunk)
if err == nil && req.ArtifactPaths != nil {
req.ArtifactPaths.ResponsePath = path
}
return path, err
}
// recordLLMSummaryArtifact 记录模型调用摘要。
func recordLLMSummaryArtifact(req StreamRequest, payload map[string]any) {
if req.Observer == nil {
return
}
path, err := req.Observer.RecordLLMSummary(req.RequestID, req.RunID, req.ModelCallID, payload)
if err == nil && req.ArtifactPaths != nil {
req.ArtifactPaths.SummaryPath = path
}
}
// summarizeTools 生成工具列表摘要。
func summarizeTools(items []json.RawMessage) []string {
result := make([]string, 0, len(items))
for _, item := range items {
var wrapper struct {
Function struct {
Name string `json:"name"`
} `json:"function"`
}
if err := json.Unmarshal(item, &wrapper); err != nil {
continue
}
if strings.TrimSpace(wrapper.Function.Name) != "" {
result = append(result, strings.TrimSpace(wrapper.Function.Name))
}
}
return result
}
// summarizeMessages 生成消息列表摘要。
func summarizeMessages(items []Message) []map[string]any {
result := make([]map[string]any, 0, len(items))
for _, item := range items {
content := truncateArtifactText(item.Content, 120)
imageCount := 0
for _, part := range item.ContentParts {
if normalizeContentPartType(part.Type) == contentPartTypeImage {
imageCount++
}
}
result = append(result, map[string]any{
"role": item.Role,
"content_preview": content,
"content_length": len([]rune(content)),
"image_count": imageCount,
"tool_call_count": len(item.ToolCalls),
"tool_call_id": item.ToolCallID,
"name": item.Name,
})
}
return result
}
func truncateArtifactText(text string, maxRunes int) string {
trimmed := strings.TrimSpace(text)
if trimmed == "" || maxRunes <= 0 {
return ""
}
runes := []rune(trimmed)
if len(runes) <= maxRunes {
return trimmed
}
return string(runes[:maxRunes]) + "..."
}
// normalizeModelArtifactTime 把时间格式化为 RFC3339Nano。
func normalizeModelArtifactTime(value time.Time) string {
if value.IsZero() {
return ""
}
return value.UTC().Format(time.RFC3339Nano)
}
// computeTTFTMS 计算首事件耗时。
func computeTTFTMS(startedAt time.Time, firstEventAt time.Time) int64 {
if startedAt.IsZero() || firstEventAt.IsZero() {
return 0
}
value := firstEventAt.Sub(startedAt).Milliseconds()
if value < 0 {
return 0
}
return value
}
// computeDurationMS 计算调用总耗时。
func computeDurationMS(startedAt time.Time, finishedAt time.Time) int64 {
if startedAt.IsZero() || finishedAt.IsZero() {
return 0
}
value := finishedAt.Sub(startedAt).Milliseconds()
if value < 0 {
return 0
}
return value
}
// summarizeModelArtifactError 返回可安全落盘的错误文本。
func summarizeModelArtifactError(err error) string {
if err == nil {
return ""
}
return strings.TrimSpace(err.Error())
}
// firstNonEmptyString 返回第一个非空字符串。
func firstNonEmptyString(values ...string) string {
for _, value := range values {
if strings.TrimSpace(value) != "" {
return strings.TrimSpace(value)
}
}
return ""
}
@@ -0,0 +1,213 @@
package modeladapter
import (
"encoding/base64"
"fmt"
"net/http"
"os"
"path/filepath"
"strings"
)
const (
contentPartTypeText = "text"
contentPartTypeImage = "image"
)
// ContentPart 表示一条消息中的结构化内容块。
type ContentPart struct {
Type string `json:"type"`
Text string `json:"text,omitempty"`
Image *ImageContent `json:"image,omitempty"`
}
// ImageContent 表示消息中携带的一张图片。
type ImageContent struct {
MIMEType string `json:"mime_type,omitempty"`
Path string `json:"path,omitempty"`
Data []byte `json:"data,omitempty"`
}
func hasImageContentParts(parts []ContentPart) bool {
for _, part := range parts {
if normalizeContentPartType(part.Type) == contentPartTypeImage {
return true
}
}
return false
}
func collapseTextContentParts(parts []ContentPart) string {
if len(parts) == 0 {
return ""
}
texts := make([]string, 0, len(parts))
for _, part := range parts {
if normalizeContentPartType(part.Type) != contentPartTypeText {
continue
}
if strings.TrimSpace(part.Text) == "" {
continue
}
texts = append(texts, part.Text)
}
return strings.Join(texts, "")
}
func normalizeContentPartType(value string) string {
trimmed := strings.TrimSpace(strings.ToLower(value))
switch trimmed {
case "", contentPartTypeText:
return contentPartTypeText
case contentPartTypeImage:
return contentPartTypeImage
default:
return trimmed
}
}
func openAIContentValue(message Message) (any, error) {
if !hasImageContentParts(message.ContentParts) {
content := message.Content
if strings.TrimSpace(content) == "" && len(message.ContentParts) > 0 {
content = collapseTextContentParts(message.ContentParts)
}
return content, nil
}
parts := make([]map[string]any, 0, len(message.ContentParts)+1)
if len(message.ContentParts) == 0 && strings.TrimSpace(message.Content) != "" {
parts = append(parts, map[string]any{
"type": contentPartTypeText,
"text": message.Content,
})
}
for _, part := range message.ContentParts {
switch normalizeContentPartType(part.Type) {
case contentPartTypeText:
if part.Text == "" {
continue
}
parts = append(parts, map[string]any{
"type": contentPartTypeText,
"text": part.Text,
})
case contentPartTypeImage:
dataURL, err := imageContentDataURL(part.Image)
if err != nil {
return nil, err
}
parts = append(parts, map[string]any{
"type": "image_url",
"image_url": map[string]any{
"url": dataURL,
},
})
default:
return nil, fmt.Errorf("unsupported openai content part type: %s", strings.TrimSpace(part.Type))
}
}
if len(parts) == 0 {
return message.Content, nil
}
return parts, nil
}
func anthropicContentBlocks(message Message) ([]map[string]any, error) {
if len(message.ContentParts) == 0 {
if strings.TrimSpace(message.Content) == "" {
return nil, nil
}
return []map[string]any{{
"type": contentPartTypeText,
"text": message.Content,
}}, nil
}
blocks := make([]map[string]any, 0, len(message.ContentParts))
for _, part := range message.ContentParts {
switch normalizeContentPartType(part.Type) {
case contentPartTypeText:
if part.Text == "" {
continue
}
blocks = append(blocks, map[string]any{
"type": contentPartTypeText,
"text": part.Text,
})
case contentPartTypeImage:
payload, mediaType, err := resolveImageContent(part.Image)
if err != nil {
return nil, err
}
blocks = append(blocks, map[string]any{
"type": contentPartTypeImage,
"source": map[string]any{
"type": "base64",
"media_type": mediaType,
"data": base64.StdEncoding.EncodeToString(payload),
},
})
default:
return nil, fmt.Errorf("unsupported anthropic content part type: %s", strings.TrimSpace(part.Type))
}
}
if len(blocks) == 0 && strings.TrimSpace(message.Content) != "" {
blocks = append(blocks, map[string]any{
"type": contentPartTypeText,
"text": message.Content,
})
}
return blocks, nil
}
func imageContentDataURL(image *ImageContent) (string, error) {
payload, mediaType, err := resolveImageContent(image)
if err != nil {
return "", err
}
return "data:" + mediaType + ";base64," + base64.StdEncoding.EncodeToString(payload), nil
}
func resolveImageContent(image *ImageContent) ([]byte, string, error) {
if image == nil {
return nil, "", fmt.Errorf("image content is required")
}
if len(image.Data) > 0 {
return image.Data, normalizeImageMIMEType(image.MIMEType, image.Path, image.Data), nil
}
path := strings.TrimSpace(image.Path)
if path == "" {
return nil, "", fmt.Errorf("image content is missing data and path")
}
payload, err := os.ReadFile(path)
if err != nil {
return nil, "", fmt.Errorf("read image content failed: %w", err)
}
return payload, normalizeImageMIMEType(image.MIMEType, path, payload), nil
}
func normalizeImageMIMEType(mimeType string, path string, payload []byte) string {
trimmed := strings.TrimSpace(strings.ToLower(mimeType))
if trimmed != "" {
return trimmed
}
if len(payload) > 0 {
detected := strings.TrimSpace(strings.ToLower(http.DetectContentType(payload)))
if strings.HasPrefix(detected, "image/") {
return detected
}
}
switch strings.ToLower(filepath.Ext(strings.TrimSpace(path))) {
case ".jpg", ".jpeg":
return "image/jpeg"
case ".png":
return "image/png"
case ".gif":
return "image/gif"
case ".webp":
return "image/webp"
default:
return "image/png"
}
}
@@ -0,0 +1,97 @@
package modeladapter
import (
"strings"
"time"
"google.golang.org/protobuf/encoding/protojson"
"cursor/gen/agentv1"
runtimecore "cursor/internal/backend/agent/core"
)
func emitCreatePlanToolProgress(
sink func(ModelEvent) error,
provider string,
model string,
callID string,
rawArgs string,
argsTextDelta string,
lastSnapshot *string,
) error {
if sink == nil || lastSnapshot == nil {
return nil
}
trimmedCallID := strings.TrimSpace(callID)
if trimmedCallID == "" {
return nil
}
args, ok := createPlanArgsProgressSnapshot(rawArgs)
if !ok {
return nil
}
signatureBytes, err := protojson.MarshalOptions{UseProtoNames: true}.Marshal(args)
if err != nil {
return err
}
signature := string(signatureBytes)
if signature == "" || signature == *lastSnapshot {
return nil
}
*lastSnapshot = signature
if err := sink(ModelEvent{
Kind: ModelEventKindPartialToolCall,
OccurredAt: time.Now().UTC(),
Provider: provider,
Model: model,
ToolCallID: trimmedCallID,
ArgsTextDelta: argsTextDelta,
ToolCall: &agentv1.ToolCall{
Tool: &agentv1.ToolCall_CreatePlanToolCall{
CreatePlanToolCall: &agentv1.CreatePlanToolCall{
Args: args,
},
},
},
}); err != nil {
return err
}
return nil
}
func createPlanArgsProgressSnapshot(rawArgs string) (*agentv1.CreatePlanArgs, bool) {
trimmed := strings.TrimSpace(rawArgs)
if trimmed == "" {
return nil, false
}
if args, err := runtimecore.DecodeCreatePlanArgsJSON([]byte(trimmed)); err == nil && hasCreatePlanArgsProgress(args) {
return args, true
}
args := &agentv1.CreatePlanArgs{}
if value, found, _ := extractJSONStringFieldPrefix(trimmed, "plan"); found {
args.Plan = value
}
if value, found, _ := extractJSONStringFieldPrefix(trimmed, "overview"); found {
args.Overview = value
}
if value, found, complete := extractJSONStringFieldPrefix(trimmed, "name"); found && complete {
args.Name = strings.TrimSpace(value)
}
if !hasCreatePlanArgsProgress(args) {
return nil, false
}
return args, true
}
func hasCreatePlanArgsProgress(args *agentv1.CreatePlanArgs) bool {
if args == nil {
return false
}
return args.GetPlan() != "" ||
args.GetOverview() != "" ||
args.GetName() != "" ||
args.GetIsProject() ||
len(args.GetTodos()) > 0 ||
len(args.GetPhases()) > 0
}
+2
View File
@@ -0,0 +1,2 @@
// Package modeladapter 提供 OpenAI / Anthropic 兼容流式模型适配器与路由器。
package modeladapter
@@ -0,0 +1,41 @@
// http_error.go 负责把非 2xx HTTP 响应整理成带响应体摘要的错误。
package modeladapter
import (
"fmt"
"io"
"net/http"
"strings"
)
const (
// maxErrorBodyBytes 表示错误响应体最多读取的字节数。
maxErrorBodyBytes = 8192
)
// buildHTTPStatusError 读取响应体摘要并生成带状态码的错误。
func buildHTTPStatusError(prefix string, resp *http.Response) error {
if resp == nil {
return fmt.Errorf("%s response is nil", strings.TrimSpace(prefix))
}
limitedBody, err := io.ReadAll(io.LimitReader(resp.Body, maxErrorBodyBytes))
if err != nil {
if retrySummary := ProviderRetryAttemptSummary(resp); retrySummary != "" {
return fmt.Errorf("%s status=%d %s body_read_error=%v", strings.TrimSpace(prefix), resp.StatusCode, retrySummary, err)
}
return fmt.Errorf("%s status=%d body_read_error=%v", strings.TrimSpace(prefix), resp.StatusCode, err)
}
retrySummary := ProviderRetryAttemptSummary(resp)
bodyText := strings.TrimSpace(string(limitedBody))
if bodyText == "" {
if retrySummary != "" {
return fmt.Errorf("%s status=%d %s", strings.TrimSpace(prefix), resp.StatusCode, retrySummary)
}
return fmt.Errorf("%s status=%d", strings.TrimSpace(prefix), resp.StatusCode)
}
if retrySummary != "" {
return fmt.Errorf("%s status=%d %s body=%s", strings.TrimSpace(prefix), resp.StatusCode, retrySummary, bodyText)
}
return fmt.Errorf("%s status=%d body=%s", strings.TrimSpace(prefix), resp.StatusCode, bodyText)
}
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,18 @@
package modeladapter
func maxAnthropicTokens(req StreamRequest) int {
if req.AnthropicMaxTokens > 0 {
return req.AnthropicMaxTokens
}
if req.MaxTokens > 0 {
return req.MaxTokens
}
return 65536
}
func maxThinkingBudget(req StreamRequest) int {
if req.ThinkingBudgetTokens > 0 {
return req.ThinkingBudgetTokens
}
return anthropicThinkingBudget(maxAnthropicTokens(req))
}
@@ -0,0 +1,121 @@
package modeladapter
import (
"encoding/json"
"fmt"
"net/http"
"strings"
)
func cloneRequestBodyOverride(input map[string]any) map[string]any {
if len(input) == 0 {
return nil
}
payload, err := json.Marshal(input)
if err != nil {
return nil
}
var cloned map[string]any
if err := json.Unmarshal(payload, &cloned); err != nil {
return nil
}
return cloned
}
func requestBodyToMap(input any) (map[string]any, error) {
if body, ok := input.(map[string]any); ok {
return body, nil
}
payload, err := json.Marshal(input)
if err != nil {
return nil, err
}
var body map[string]any
if err := json.Unmarshal(payload, &body); err != nil {
return nil, err
}
if body == nil {
body = map[string]any{}
}
return body, nil
}
func ApplyOpenAIExtraParams(body map[string]any, enabled bool, paramsJSON string) error {
return applyExtraParams(body, enabled, paramsJSON, "openai extra params json")
}
func ApplyAnthropicExtraParams(body map[string]any, enabled bool, paramsJSON string) error {
return applyExtraParams(body, enabled, paramsJSON, "anthropic extra params json")
}
func applyExtraParams(body map[string]any, enabled bool, paramsJSON string, label string) error {
if !enabled {
return nil
}
if body == nil {
return fmt.Errorf("%s target body is nil", label)
}
extraParams, err := parseJSONMap(paramsJSON, label)
if err != nil {
return err
}
for key, value := range extraParams {
name := strings.TrimSpace(key)
if name == "" {
continue
}
body[name] = value
}
return nil
}
func ApplyCustomHeaders(httpReq *http.Request, enabled bool, headersJSON string) error {
if !enabled {
return nil
}
if httpReq == nil {
return fmt.Errorf("custom headers target request is nil")
}
headers, err := parseStringJSONMap(headersJSON, "custom headers json")
if err != nil {
return err
}
for key, value := range headers {
name := strings.TrimSpace(key)
if name == "" {
continue
}
httpReq.Header.Set(name, value)
}
return nil
}
func parseJSONMap(value string, label string) (map[string]any, error) {
text := strings.TrimSpace(value)
if text == "" {
return nil, fmt.Errorf("%s is empty", label)
}
var parsed map[string]any
if err := json.Unmarshal([]byte(text), &parsed); err != nil {
return nil, fmt.Errorf("%s must be an object: %w", label, err)
}
if parsed == nil {
return nil, fmt.Errorf("%s must be an object", label)
}
return parsed, nil
}
func parseStringJSONMap(value string, label string) (map[string]string, error) {
text := strings.TrimSpace(value)
if text == "" {
return nil, fmt.Errorf("%s is empty", label)
}
var parsed map[string]string
if err := json.Unmarshal([]byte(text), &parsed); err != nil {
return nil, fmt.Errorf("%s must be an object with string values: %w", label, err)
}
if parsed == nil {
return nil, fmt.Errorf("%s must be an object", label)
}
return parsed, nil
}
+49
View File
@@ -0,0 +1,49 @@
// retry.go 保留 provider HTTP 请求入口的历史命名;provider 错误交给客户端重连链路处理。
package modeladapter
import (
"context"
"net/http"
)
// DoProviderRequestWithRetry 保留旧入口名;本地模式不在服务端重试 provider 请求。
func DoProviderRequestWithRetry(
ctx context.Context,
client *http.Client,
provider string,
requestID string,
modelCallID string,
buildRequest func(context.Context) (*http.Request, error),
) (*http.Response, error) {
return doProviderRequestWithRetry(ctx, client, provider, requestID, modelCallID, buildRequest)
}
func doProviderRequestWithRetry(
ctx context.Context,
client *http.Client,
provider string,
requestID string,
modelCallID string,
buildRequest func(context.Context) (*http.Request, error),
) (*http.Response, error) {
if client == nil {
client = http.DefaultClient
}
httpReq, err := buildRequest(ctx)
if err != nil {
return nil, err
}
resp, err := client.Do(httpReq)
if err != nil {
if resp != nil && resp.Body != nil {
_ = resp.Body.Close()
}
return nil, err
}
return resp, nil
}
// ProviderRetryAttemptSummary 返回空值;provider 请求不再有服务端内部重试摘要。
func ProviderRetryAttemptSummary(resp *http.Response) string {
return ""
}
+419
View File
@@ -0,0 +1,419 @@
// router.go 按模型标识选择 OpenAI 或 Anthropic 兼容适配器。
package modeladapter
import (
"context"
"fmt"
"strings"
"time"
legacyruntime "cursor/internal/runtime"
)
// Router 是 MVP 阶段的模型适配路由器。
type Router struct {
// openai 负责 OpenAI 兼容流式请求。
openai ModelAdapter
// anthropic 负责 Anthropic 兼容流式请求。
anthropic ModelAdapter
// resolver 负责从本地配置中解析实际模型通道。
resolver ChannelResolver
}
type ChannelResolver interface {
SelectChannelForModel(context.Context, string) (*legacyruntime.ResolvedChannel, error)
ProviderStreamIdleTimeout(context.Context) time.Duration
}
// NewRouter 创建模型适配路由器。
func NewRouter(resolver ChannelResolver) *Router {
return &Router{
openai: NewOpenAIAdapter(),
anthropic: NewAnthropicAdapter(),
resolver: resolver,
}
}
// Stream 根据模型标识选择具体 provider 并转发请求。
func (router *Router) Stream(ctx context.Context, req StreamRequest, sink func(ModelEvent) error) error {
if router == nil || router.resolver == nil {
return fmt.Errorf("model adapter resolver is unavailable")
}
channel, err := router.resolver.SelectChannelForModel(ctx, req.ModelID)
if err != nil {
return err
}
if channel == nil {
return fmt.Errorf("no available channel for model %q", req.ModelID)
}
resolved := req
resolved.Provider = strings.TrimSpace(channel.Provider)
resolved.BaseURL = strings.TrimSpace(channel.BaseURL)
resolved.APIKey = strings.TrimSpace(channel.APIKey)
resolved.ProviderModelID = strings.TrimSpace(channel.Model)
resolved.ResolvedChannelID = strings.TrimSpace(channel.ID)
resolved.ResolvedChannelName = strings.TrimSpace(channel.Name)
resolved.ResolvedContextWindowTokens = channel.ContextWindowTokens
resolved.ReasoningEffort = openAIReasoningEffortFromRuntime(channel.ReasoningEffort)
resolved.OpenAIEndpoint = strings.TrimSpace(channel.OpenAIEndpoint)
resolved.OpenAIExtraParamsEnabled = channel.OpenAIExtraParamsEnabled
resolved.OpenAIExtraParamsJSON = strings.TrimSpace(channel.OpenAIExtraParamsJSON)
resolved.CustomHeadersEnabled = channel.CustomHeadersEnabled
resolved.CustomHeadersJSON = strings.TrimSpace(channel.CustomHeadersJSON)
resolved.AnthropicExtraParamsEnabled = channel.AnthropicExtraParamsEnabled
resolved.AnthropicExtraParamsJSON = strings.TrimSpace(channel.AnthropicExtraParamsJSON)
resolved.AnthropicMaxTokens = channel.AnthropicMaxTokens
resolved.AnthropicThinkingEffort = strings.TrimSpace(channel.AnthropicThinkingEffort)
resolved.ThinkingBudgetTokens = channel.ThinkingBudgetTokens
resolved.ProviderStreamIdleTimeout = router.resolver.ProviderStreamIdleTimeout(ctx)
runtimeThinkingEffort := normalizeRuntimeThinkingEffort(req.ThinkingEffort)
if runtimeThinkingEffort != "" {
resolved.ThinkingEffort = runtimeThinkingEffort
if runtimeThinkingEffort == "disabled" {
resolved.ReasoningEffort = ""
resolved.AnthropicThinkingEffort = ""
} else {
resolved.ReasoningEffort = openAIReasoningEffortFromRuntime(runtimeThinkingEffort)
resolved.AnthropicThinkingEffort = runtimeThinkingEffort
}
} else {
resolved.ThinkingEffort = ""
}
if resolved.MaxTokens <= 0 && channel.MaxTokens > 0 {
resolved.MaxTokens = channel.MaxTokens
}
if req.MaxTokens > 0 && (resolved.AnthropicMaxTokens <= 0 || req.MaxTokens < resolved.AnthropicMaxTokens) {
resolved.AnthropicMaxTokens = req.MaxTokens
}
if resolved.AnthropicMaxTokens <= 0 && resolved.MaxTokens > 0 {
resolved.AnthropicMaxTokens = resolved.MaxTokens
}
if resolved.ProviderModelID == "" {
resolved.ProviderModelID = strings.TrimSpace(req.ModelID)
}
resolved.Messages = sanitizeProviderMessages(req.Messages)
if resolved.RequestKnobs != nil {
resolved.RequestKnobs["max_tokens"] = resolved.MaxTokens
if runtimeThinkingEffort != "" {
resolved.RequestKnobs["runtime_thinking_effort"] = runtimeThinkingEffort
} else {
delete(resolved.RequestKnobs, "runtime_thinking_effort")
}
if resolved.Provider == "openai" {
if strings.TrimSpace(resolved.ReasoningEffort) != "" {
resolved.RequestKnobs["reasoning_effort"] = strings.TrimSpace(resolved.ReasoningEffort)
} else {
delete(resolved.RequestKnobs, "reasoning_effort")
}
resolved.RequestKnobs["openai_endpoint"] = resolved.OpenAIEndpoint
resolved.RequestKnobs["openai_extra_params_enabled"] = resolved.OpenAIExtraParamsEnabled
resolved.RequestKnobs["custom_headers_enabled"] = resolved.CustomHeadersEnabled
} else if resolved.Provider == "anthropic" {
delete(resolved.RequestKnobs, "reasoning_effort")
resolved.RequestKnobs["custom_headers_enabled"] = resolved.CustomHeadersEnabled
resolved.RequestKnobs["anthropic_extra_params_enabled"] = resolved.AnthropicExtraParamsEnabled
anthropicMaxTokens := maxAnthropicTokens(resolved)
resolved.RequestKnobs["max_tokens"] = anthropicMaxTokens
resolved.RequestKnobs["anthropic_max_tokens"] = anthropicMaxTokens
if strings.TrimSpace(resolved.AnthropicThinkingEffort) != "" {
resolved.RequestKnobs["anthropic_thinking_effort"] = anthropicThinkingEffort(resolved)
} else {
delete(resolved.RequestKnobs, "anthropic_thinking_effort")
}
}
}
switch resolved.Provider {
case "anthropic":
return router.anthropic.Stream(ctx, resolved, sink)
case "openai":
return router.openai.Stream(ctx, resolved, sink)
default:
return fmt.Errorf("unsupported provider %q", resolved.Provider)
}
}
// sanitizeProviderMessages removes replay-only placeholders and trims trailing
// assistant prefill so providers that require a user/tool terminal message do
// not reject the request.
func sanitizeProviderMessages(input []Message) []Message {
if len(input) == 0 {
return nil
}
filtered := make([]Message, 0, len(input))
for _, message := range input {
if isAssistantPlaceholderMessage(message) {
continue
}
filtered = append(filtered, message)
}
filtered = mergeAdjacentAssistantToolCallMessages(filtered)
filtered = trimDanglingAssistantToolCalls(filtered)
for len(filtered) > 0 && isAssistantPrefillMessage(filtered[len(filtered)-1]) {
filtered = filtered[:len(filtered)-1]
}
if len(filtered) == 0 {
return nil
}
return filtered
}
func isAssistantPlaceholderMessage(message Message) bool {
if strings.TrimSpace(message.Role) != "assistant" {
return false
}
if len(message.ToolCalls) > 0 || len(message.ContentParts) > 0 {
return false
}
if strings.TrimSpace(message.ToolCallID) != "" || strings.TrimSpace(message.Name) != "" {
return false
}
if strings.TrimSpace(message.ReasoningContent) != "" {
return false
}
if strings.TrimSpace(message.ReasoningSignature) != "" {
return false
}
switch strings.TrimSpace(message.Content) {
case "":
return true
default:
return false
}
}
func isAssistantPrefillMessage(message Message) bool {
if strings.TrimSpace(message.Role) != "assistant" {
return false
}
if len(message.ToolCalls) > 0 {
return false
}
if strings.TrimSpace(message.ToolCallID) != "" || strings.TrimSpace(message.Name) != "" {
return false
}
return strings.TrimSpace(message.Content) != "" || strings.TrimSpace(message.ReasoningContent) != ""
}
func mergeAdjacentAssistantToolCallMessages(input []Message) []Message {
if len(input) == 0 {
return nil
}
merged := make([]Message, 0, len(input))
for _, raw := range input {
message := cloneProviderMessage(raw)
if mergeProviderAssistantToolCalls(&merged, message) {
continue
}
merged = append(merged, message)
}
return merged
}
func cloneProviderMessage(message Message) Message {
cloned := message
if len(message.ContentParts) > 0 {
cloned.ContentParts = append([]ContentPart(nil), message.ContentParts...)
}
if len(message.ToolCalls) > 0 {
cloned.ToolCalls = append([]ToolCallDescriptor(nil), message.ToolCalls...)
}
if len(message.OpenAIResponsesReasoningSummary) > 0 {
cloned.OpenAIResponsesReasoningSummary = append([]byte(nil), message.OpenAIResponsesReasoningSummary...)
}
return cloned
}
func mergeProviderAssistantToolCalls(messages *[]Message, message Message) bool {
if len(*messages) == 0 {
return false
}
last := &(*messages)[len(*messages)-1]
if !canMergeProviderAssistantToolCalls(*last, message) {
return false
}
startIndex := len(last.ToolCalls)
for index, toolCall := range message.ToolCalls {
item := toolCall
item.Index = startIndex + index
last.ToolCalls = append(last.ToolCalls, item)
}
last.ReasoningContent = mergeProviderReasoning(last.ReasoningContent, message.ReasoningContent)
mergeProviderReasoningMetadata(last, message)
return true
}
func canMergeProviderAssistantToolCalls(last Message, current Message) bool {
if strings.TrimSpace(last.Role) != "assistant" || strings.TrimSpace(current.Role) != "assistant" {
return false
}
if len(last.ToolCalls) == 0 || len(current.ToolCalls) == 0 {
return false
}
if strings.TrimSpace(last.ToolCallID) != "" || strings.TrimSpace(last.Name) != "" {
return false
}
if strings.TrimSpace(current.ToolCallID) != "" || strings.TrimSpace(current.Name) != "" {
return false
}
if strings.TrimSpace(current.Content) != "" || len(current.ContentParts) > 0 {
return false
}
return true
}
func mergeProviderReasoning(left string, right string) string {
left = strings.TrimSpace(left)
right = strings.TrimSpace(right)
switch {
case left == "":
return right
case right == "", right == left:
return left
default:
return left + "\n\n" + right
}
}
func mergeProviderReasoningSignature(left string, right string) string {
left = strings.TrimSpace(left)
right = strings.TrimSpace(right)
switch {
case left == "":
return right
case right == "", right == left:
return left
default:
return ""
}
}
func mergeProviderReasoningSignatureSource(left string, right string) string {
left = strings.TrimSpace(left)
right = strings.TrimSpace(right)
switch {
case left == "":
return right
case right == "", right == left:
return left
default:
return ""
}
}
func mergeProviderReasoningMetadata(last *Message, current Message) {
if last == nil {
return
}
leftSignature := strings.TrimSpace(last.ReasoningSignature)
rightSignature := strings.TrimSpace(current.ReasoningSignature)
mergedSignature := mergeProviderReasoningSignature(leftSignature, rightSignature)
last.ReasoningSignature = mergedSignature
if mergedSignature == "" {
last.ReasoningSignatureSource = ""
last.OpenAIResponsesReasoningID = ""
last.OpenAIResponsesReasoningStatus = ""
last.OpenAIResponsesReasoningSummary = nil
return
}
if leftSignature == "" && rightSignature != "" {
last.ReasoningSignatureSource = strings.TrimSpace(current.ReasoningSignatureSource)
last.OpenAIResponsesReasoningID = current.OpenAIResponsesReasoningID
last.OpenAIResponsesReasoningStatus = current.OpenAIResponsesReasoningStatus
last.OpenAIResponsesReasoningSummary = append([]byte(nil), current.OpenAIResponsesReasoningSummary...)
return
}
if leftSignature == rightSignature {
last.ReasoningSignatureSource = mergeProviderReasoningSignatureSource(last.ReasoningSignatureSource, current.ReasoningSignatureSource)
if strings.TrimSpace(last.OpenAIResponsesReasoningID) == "" {
last.OpenAIResponsesReasoningID = current.OpenAIResponsesReasoningID
}
if strings.TrimSpace(last.OpenAIResponsesReasoningStatus) == "" {
last.OpenAIResponsesReasoningStatus = current.OpenAIResponsesReasoningStatus
}
if len(last.OpenAIResponsesReasoningSummary) == 0 {
last.OpenAIResponsesReasoningSummary = append([]byte(nil), current.OpenAIResponsesReasoningSummary...)
}
}
}
func trimDanglingAssistantToolCalls(input []Message) []Message {
if len(input) == 0 {
return nil
}
trimmed := make([]Message, 0, len(input))
for index := 0; index < len(input); index++ {
message := cloneProviderMessage(input[index])
if strings.TrimSpace(message.Role) != "assistant" || len(message.ToolCalls) == 0 {
trimmed = append(trimmed, message)
continue
}
end := index + 1
responded := make(map[string]struct{}, len(message.ToolCalls))
for end < len(input) && strings.TrimSpace(input[end].Role) == "tool" {
toolCallID := strings.TrimSpace(input[end].ToolCallID)
if toolCallID != "" {
responded[toolCallID] = struct{}{}
}
end++
}
nextToolCalls := make([]ToolCallDescriptor, 0, len(message.ToolCalls))
allowedToolCallIDs := make(map[string]struct{}, len(message.ToolCalls))
for _, toolCall := range message.ToolCalls {
toolCallID := strings.TrimSpace(toolCall.ID)
if _, ok := responded[toolCallID]; !ok {
continue
}
item := toolCall
item.Index = len(nextToolCalls)
nextToolCalls = append(nextToolCalls, item)
allowedToolCallIDs[toolCallID] = struct{}{}
}
if len(nextToolCalls) > 0 {
message.ToolCalls = nextToolCalls
trimmed = append(trimmed, message)
for toolIndex := index + 1; toolIndex < end; toolIndex++ {
toolMessage := cloneProviderMessage(input[toolIndex])
if _, ok := allowedToolCallIDs[strings.TrimSpace(toolMessage.ToolCallID)]; !ok {
continue
}
trimmed = append(trimmed, toolMessage)
}
} else if strings.TrimSpace(message.Content) != "" || len(message.ContentParts) > 0 || strings.TrimSpace(message.ReasoningContent) != "" {
message.ToolCalls = nil
trimmed = append(trimmed, message)
}
index = end - 1
}
return trimmed
}
func normalizeRuntimeThinkingEffort(raw string) string {
switch strings.ToLower(strings.TrimSpace(raw)) {
case "disabled", "low", "medium", "high", "xhigh", "max":
return strings.ToLower(strings.TrimSpace(raw))
case "disable", "off", "none", "false", "no", "0":
return "disabled"
case "very_high", "very-high", "veryhigh", "x-high", "extra_high", "extra-high", "extrahigh":
return "xhigh"
case "maximum":
return "max"
default:
return ""
}
}
func openAIReasoningEffortFromRuntime(runtimeThinkingEffort string) string {
switch normalizeRuntimeThinkingEffort(runtimeThinkingEffort) {
case "low", "medium", "high", "xhigh", "max":
return normalizeRuntimeThinkingEffort(runtimeThinkingEffort)
default:
return ""
}
}
+133
View File
@@ -0,0 +1,133 @@
package modeladapter
import (
"context"
"fmt"
"io"
"sync"
"time"
)
const (
defaultProviderStreamIdleTimeout = 4 * time.Minute
minProviderStreamIdleTimeout = 30 * time.Second
)
type providerStreamIdleWatchdog struct {
ctx context.Context
cancel context.CancelCauseFunc
timeout time.Duration
timer *time.Timer
mu sync.Mutex
body io.Closer
stopped bool
timedOut bool
err error
}
func newProviderStreamIdleWatchdog(parent context.Context, timeout time.Duration) (context.Context, *providerStreamIdleWatchdog) {
if parent == nil {
parent = context.Background()
}
timeout = normalizeProviderStreamIdleTimeoutDuration(timeout)
ctx, cancel := context.WithCancelCause(parent)
watchdog := &providerStreamIdleWatchdog{
ctx: ctx,
cancel: cancel,
timeout: timeout,
err: providerStreamIdleTimeoutError(timeout),
}
watchdog.timer = time.AfterFunc(watchdog.timeout, watchdog.expire)
return ctx, watchdog
}
func (watchdog *providerStreamIdleWatchdog) AttachBody(body io.Closer) {
if watchdog == nil || body == nil {
return
}
watchdog.mu.Lock()
watchdog.body = body
shouldClose := watchdog.timedOut || watchdog.stopped
watchdog.mu.Unlock()
if shouldClose {
_ = body.Close()
}
}
func (watchdog *providerStreamIdleWatchdog) MarkEffectiveContent() {
if watchdog == nil {
return
}
watchdog.mu.Lock()
defer watchdog.mu.Unlock()
if watchdog.stopped || watchdog.timedOut || watchdog.timer == nil {
return
}
watchdog.timer.Reset(watchdog.timeout)
}
func (watchdog *providerStreamIdleWatchdog) Stop() {
if watchdog == nil {
return
}
watchdog.mu.Lock()
if watchdog.stopped {
watchdog.mu.Unlock()
return
}
watchdog.stopped = true
watchdog.body = nil
if watchdog.timer != nil {
watchdog.timer.Stop()
}
watchdog.mu.Unlock()
watchdog.cancel(nil)
}
func (watchdog *providerStreamIdleWatchdog) Err() error {
if watchdog == nil {
return nil
}
watchdog.mu.Lock()
defer watchdog.mu.Unlock()
if watchdog.timedOut {
return watchdog.err
}
return nil
}
func (watchdog *providerStreamIdleWatchdog) expire() {
watchdog.mu.Lock()
if watchdog.stopped || watchdog.timedOut {
watchdog.mu.Unlock()
return
}
watchdog.timedOut = true
body := watchdog.body
err := watchdog.err
watchdog.mu.Unlock()
watchdog.cancel(err)
if body != nil {
_ = body.Close()
}
}
func normalizeProviderStreamIdleTimeoutDuration(timeout time.Duration) time.Duration {
if timeout <= 0 {
return defaultProviderStreamIdleTimeout
}
if timeout < minProviderStreamIdleTimeout {
return minProviderStreamIdleTimeout
}
return timeout
}
func providerStreamIdleTimeoutError(timeout time.Duration) error {
seconds := int(timeout / time.Second)
if seconds > 0 && timeout == time.Duration(seconds)*time.Second {
return fmt.Errorf("provider stream idle timeout after %ds without effective content", seconds)
}
return fmt.Errorf("provider stream idle timeout after %s without effective content", timeout)
}
@@ -0,0 +1,126 @@
package modeladapter
import (
"crypto/sha256"
"encoding/hex"
"strings"
)
const (
maxProviderToolCallIDLen = 64
toolCallNamespaceHashLen = 12
toolCallValueHashLen = 12
)
// namespaceToolCallID 为 provider 原始 tool call id 增加 model-call 级别命名空间,
// 避免像 functions.Shell:0 这类跨轮复用的 id 在客户端被误判为同一个 bubble。
//
// OpenAI 等 provider 对 tool_call_id 长度有限制,因此这里使用 model_call_id 的短哈希
// 而不是完整 UUID,保证内部存储的 tool_call_id 既稳定又能安全回放给 provider。
func namespaceToolCallID(modelCallID string, rawToolCallID string) string {
raw := strings.TrimSpace(rawToolCallID)
if raw == "" {
return ""
}
if strings.Contains(raw, "::") {
return providerToolCallID(raw)
}
model := strings.TrimSpace(modelCallID)
if model == "" {
return providerToolCallID(raw)
}
return buildProviderSafeToolCallID(shortToolCallHash(model, toolCallNamespaceHashLen), raw)
}
// providerToolCallID 把内部持久化的 tool_call_id 规整成 provider 可接受的安全长度。
// 这样旧会话里已经落盘的 legacy "<modelCallID>::<rawID>" 也能继续回放。
func providerToolCallID(toolCallID string) string {
trimmed := strings.TrimSpace(toolCallID)
if trimmed == "" {
return ""
}
namespace, raw, ok := splitLegacyToolCallID(trimmed)
if ok {
return buildProviderSafeToolCallID(shortToolCallHash(namespace, toolCallNamespaceHashLen), raw)
}
if len(trimmed) <= maxProviderToolCallIDLen {
return trimmed
}
return buildProviderSafeToolCallID("", trimmed)
}
type providerToolCallDescriptor struct {
ID string `json:"id"`
Index int `json:"index,omitempty"`
Type string `json:"type"`
Function ToolCallFunctionShape `json:"function"`
}
func normalizeToolCallDescriptors(toolCalls []ToolCallDescriptor) []providerToolCallDescriptor {
if len(toolCalls) == 0 {
return nil
}
normalized := make([]providerToolCallDescriptor, 0, len(toolCalls))
for _, toolCall := range toolCalls {
item := providerToolCallDescriptor{
ID: providerToolCallID(toolCall.ID),
Index: toolCall.Index,
Type: toolCall.Type,
Function: toolCall.Function,
}
normalized = append(normalized, item)
}
return normalized
}
func buildProviderSafeToolCallID(namespace string, raw string) string {
trimmedRaw := strings.TrimSpace(raw)
if trimmedRaw == "" {
return ""
}
if namespace == "" && len(trimmedRaw) <= maxProviderToolCallIDLen && !strings.Contains(trimmedRaw, "::") {
return trimmedRaw
}
prefix := "tc"
if namespace != "" {
prefix += "_" + namespace
}
candidate := prefix + "_" + trimmedRaw
if len(candidate) <= maxProviderToolCallIDLen {
return candidate
}
rawHash := shortToolCallHash(trimmedRaw, toolCallValueHashLen)
remaining := maxProviderToolCallIDLen - len(prefix) - len(rawHash) - 2
if remaining <= 0 {
return prefix + "_" + rawHash
}
suffix := trimmedRaw
if len(suffix) > remaining {
suffix = suffix[len(suffix)-remaining:]
}
return prefix + "_" + rawHash + "_" + suffix
}
func splitLegacyToolCallID(value string) (namespace string, raw string, ok bool) {
namespace, raw, ok = strings.Cut(strings.TrimSpace(value), "::")
if !ok {
return "", "", false
}
namespace = strings.TrimSpace(namespace)
raw = strings.TrimSpace(raw)
if namespace == "" || raw == "" {
return "", "", false
}
return namespace, raw, true
}
func shortToolCallHash(value string, size int) string {
sum := sha256.Sum256([]byte(strings.TrimSpace(value)))
encoded := hex.EncodeToString(sum[:])
if size <= 0 || size > len(encoded) {
return encoded
}
return encoded[:size]
}
+243
View File
@@ -0,0 +1,243 @@
// types.go 定义模型适配层的统一请求、事件与路由接口。
package modeladapter
import (
"context"
"encoding/json"
"time"
"cursor/gen/agentv1"
runtimecore "cursor/internal/backend/agent/core"
)
const (
// ReasoningSignatureSourceAnthropic 表示 signature 来自 Anthropic thinking signature。
ReasoningSignatureSourceAnthropic = "anthropic"
// ReasoningSignatureSourceOpenAIResponses 表示 signature 来自 OpenAI Responses encrypted reasoning content。
ReasoningSignatureSourceOpenAIResponses = "openai_responses"
)
// Message 表示模型适配层统一使用的消息结构。
type Message struct {
// Role 表示消息角色。
Role string `json:"role"`
// Content 表示消息文本内容。
Content string `json:"content"`
// ContentParts 表示消息中的结构化内容块,例如文本或图片。
ContentParts []ContentPart `json:"content_parts,omitempty"`
// ReasoningContent 表示推理内容(用于支持 reasoning 的模型)。
ReasoningContent string `json:"reasoning_content,omitempty"`
// ReasoningSignature 表示 provider 对推理内容签发的签名(如 Anthropic thinking signature)。
ReasoningSignature string `json:"reasoning_signature,omitempty"`
// ReasoningSignatureSource 表示 reasoning signature 的 provider 语义来源。
ReasoningSignatureSource string `json:"reasoning_signature_source,omitempty"`
// OpenAIResponsesReasoningID 保存 Responses reasoning output item 的原始 id。
OpenAIResponsesReasoningID string `json:"openai_responses_reasoning_id,omitempty"`
// OpenAIResponsesReasoningStatus 保存 Responses reasoning output item 的原始 status。
OpenAIResponsesReasoningStatus string `json:"openai_responses_reasoning_status,omitempty"`
// OpenAIResponsesReasoningSummary 保存 Responses reasoning output item 的原始 summary。
OpenAIResponsesReasoningSummary json.RawMessage `json:"openai_responses_reasoning_summary,omitempty"`
// ToolCalls 表示 assistant 发起的函数调用。
ToolCalls []ToolCallDescriptor `json:"tool_calls,omitempty"`
// ToolCallID 表示 tool role 关联的调用 id。
ToolCallID string `json:"tool_call_id,omitempty"`
// Name 表示 tool role 的工具名。
Name string `json:"name,omitempty"`
}
type ToolCallDescriptor struct {
ID string `json:"id"`
Index int `json:"index,omitempty"`
Type string `json:"type"`
Function ToolCallFunctionShape `json:"function"`
OpenAIResponsesID string `json:"openai_responses_id,omitempty"`
OpenAIResponsesCallID string `json:"openai_responses_call_id,omitempty"`
OpenAIResponsesStatus string `json:"openai_responses_status,omitempty"`
}
type ToolCallFunctionShape struct {
Name string `json:"name"`
Arguments string `json:"arguments"`
}
// StreamRequest 表示一次统一的模型流请求。
type StreamRequest struct {
// RequestID 表示当前模型调用所属 request。
RequestID string
// RunID 表示当前模型调用所属 run。
RunID string
// ModelCallID 表示当前模型调用标识。
ModelCallID string
// ConversationID 表示当前模型调用所属会话,用于稳定 provider 侧 prompt cache 路由。
ConversationID string
// Mode 表示当前运行模式。
Mode agentv1.AgentMode
// ModelID 表示当前模型标识。
ModelID string
// ThinkingEffort 表示客户端在本轮运行时选择的思考强度覆盖。
ThinkingEffort string
// Provider 表示目标 provider 类型,例如 openai 或 anthropic。
Provider string
// BaseURL 表示请求应发送到的 provider 基础地址。
BaseURL string
// APIKey 表示 provider 鉴权凭据。
APIKey string
// ProviderModelID 表示 provider 侧真实模型标识。
ProviderModelID string
// ResolvedChannelID 表示本次请求实际命中的 adapter 渠道 ID。
ResolvedChannelID string
// ResolvedChannelName 表示本次请求实际命中的 adapter 展示名。
ResolvedChannelName string
// ResolvedContextWindowTokens 表示本次请求实际命中的 adapter 上下文窗口。
ResolvedContextWindowTokens int
// ReasoningEffort 表示 OpenAI 兼容 provider 的推理强度。
ReasoningEffort string
// OpenAIEndpoint 表示 OpenAI 兼容 provider 使用的 API 端点。
OpenAIEndpoint string
// OpenAIExtraParamsEnabled 表示是否启用 OpenAI 额外请求参数。
OpenAIExtraParamsEnabled bool
// OpenAIExtraParamsJSON 表示 OpenAI 额外请求参数 JSON 对象。
OpenAIExtraParamsJSON string
// CustomHeadersEnabled 表示是否启用自定义请求头。
CustomHeadersEnabled bool
// CustomHeadersJSON 表示自定义请求头 JSON 对象。
CustomHeadersJSON string
// AnthropicExtraParamsEnabled 表示是否启用 Anthropic 额外请求参数。
AnthropicExtraParamsEnabled bool
// AnthropicExtraParamsJSON 表示 Anthropic 额外请求参数 JSON 对象。
AnthropicExtraParamsJSON string
// AnthropicMaxTokens 表示 Anthropic 兼容 provider 的 max_tokens。
AnthropicMaxTokens int
// AnthropicThinkingEffort 表示 Anthropic adaptive thinking 的 output_config.effort。
AnthropicThinkingEffort string
// ThinkingBudgetTokens 表示 Anthropic thinking 预算。
ThinkingBudgetTokens int
// Messages 表示按顺序排列的消息列表。
Messages []Message
// StableMessageCount 表示 messages 中可作为稳定缓存前缀的 provider-visible 消息数量。
StableMessageCount int
// Tools 表示原始工具描述 JSON 列表。
Tools []json.RawMessage
// MaxTokens 表示本轮最大输出 token 数。
MaxTokens int
// Stream 表示当前请求必须使用流式。
Stream bool
// RequestKnobs 保存本轮请求的附加参数摘要。
RequestKnobs map[string]any
// CompileSummary 保存当前 prompt 编译摘要。
CompileSummary string
// Observer 负责写入 request-scoped LLM 工件。
Observer LLMArtifactObserver
// ArtifactPaths 用于由 adapter 回填工件路径。
ArtifactPaths *LLMArtifactPaths
// RequestBodyOverride 表示直接复用的 provider 原始请求体;设置后由 adapter 原样发送。
RequestBodyOverride map[string]any
// ProviderStreamIdleTimeout 表示 provider 流式响应无有效内容时的空闲超时。
ProviderStreamIdleTimeout time.Duration
}
// LLMArtifactPaths 表示一次模型调用相关工件路径。
type LLMArtifactPaths struct {
RequestPath string
ResponsePath string
SummaryPath string
}
// LLMArtifactObserver 定义模型调用原始工件写入接口。
type LLMArtifactObserver interface {
RecordLLMRequest(requestID string, runID string, modelCallID string, payload map[string]any) (string, error)
AppendLLMResponseChunk(requestID string, runID string, modelCallID string, chunk string) (string, error)
RecordLLMSummary(requestID string, runID string, modelCallID string, payload map[string]any) (string, error)
}
// ModelEventKind 表示统一模型事件类型。
type ModelEventKind string
const (
// ModelEventKindTextDelta 表示文本增量事件。
ModelEventKindTextDelta ModelEventKind = "text_delta"
// ModelEventKindThinkingDelta 表示思考增量事件。
ModelEventKindThinkingDelta ModelEventKind = "thinking_delta"
// ModelEventKindThinkingCompleted 表示思考结束事件。
ModelEventKindThinkingCompleted ModelEventKind = "thinking_completed"
// ModelEventKindPartialToolCall 表示工具调用已开始,但参数仍在流式生成中。
ModelEventKindPartialToolCall ModelEventKind = "partial_tool_call"
// ModelEventKindToolCallDelta 表示工具调用参数或输出的流式增量。
ModelEventKindToolCallDelta ModelEventKind = "tool_call_delta"
// ModelEventKindToolLikeCompleted 表示工具意图已完整收口。
ModelEventKindToolLikeCompleted ModelEventKind = "tool_like_completed"
// ModelEventKindTurnFinished 表示当前模型回合结束。
ModelEventKindTurnFinished ModelEventKind = "turn_finished"
// ModelEventKindProviderError 表示 provider 侧返回错误。
ModelEventKindProviderError ModelEventKind = "provider_error"
)
// ModelEvent 表示一条统一模型事件。
type ModelEvent struct {
// Kind 表示事件类型。
Kind ModelEventKind
// OccurredAt 表示当前 provider 事件发生时间。
OccurredAt time.Time
// Provider 表示当前事件所属 provider。
Provider string
// Model 表示当前事件所属模型标识。
Model string
// Text 表示文本增量。
Text string
// ThinkingStyle 表示思考样式。
ThinkingStyle agentv1.ThinkingStyle
// ThinkingDurationMS 表示思考持续时长。
ThinkingDurationMS int32
// ThinkingSignature 表示 provider 返回的思考签名(如 Anthropic signature_delta)。
ThinkingSignature string
// ThinkingSignatureSource 表示思考签名的 provider 语义来源。
ThinkingSignatureSource string
// ProviderItemID 保存 provider 原始 output item id,用于 stateless Responses replay。
ProviderItemID string
// ProviderStatus 保存 provider 原始 output item status,用于 stateless Responses replay。
ProviderStatus string
// ProviderSummary 保存 provider 原始 output item summary,用于 stateless Responses replay。
ProviderSummary json.RawMessage
// ProviderCallID 保存 provider 原始 tool/function call id,用于 stateless Responses replay。
ProviderCallID string
// ToolCallID 表示当前 partial/delta 对应的工具调用标识。
ToolCallID string
// ToolCall 保存 partial tool call 当前可公开的结构化快照。
ToolCall *agentv1.ToolCall
// ToolCallDelta 保存与当前工具调用相关的流式增量。
ToolCallDelta *agentv1.ToolCallDelta
// ArgsTextDelta 保存原始工具参数文本增量,供兼容层透传。
ArgsTextDelta string
// InputTokens 表示当前已知的输入 token 数。
InputTokens int64
// OutputTokens 表示当前已知的输出 token 数。
OutputTokens int64
// CacheReadTokens 表示当前已知的 cache read token 数。
CacheReadTokens int64
// CacheWriteTokens 表示当前已知的 cache write token 数。
CacheWriteTokens int64
// UsagePresent 表示 provider 本次流里实际返回过 usage 对象。
UsagePresent bool
// CacheReadPresent 表示 provider 明确返回了 cache read token 字段。
CacheReadPresent bool
// CacheWritePresent 表示 provider 明确返回了 cache write token 字段。
CacheWritePresent bool
// ToolInvocation 表示完成收口的工具调用意图。
ToolInvocation *runtimecore.ToolInvocation
// FinishReason 表示回合结束原因。
FinishReason string
// Err 表示 provider 错误。
Err error
}
// ModelAdapter 定义具体 provider 适配器接口。
type ModelAdapter interface {
// Stream 按流式方式发送请求,并持续产出统一模型事件。
Stream(ctx context.Context, req StreamRequest, sink func(ModelEvent) error) error
}
// ModelAdapterRouter 定义 provider 路由接口。
type ModelAdapterRouter interface {
// Stream 根据模型标识选择底层 provider 适配器。
Stream(ctx context.Context, req StreamRequest, sink func(ModelEvent) error) error
}
@@ -0,0 +1,8 @@
package modeladapter
const (
// ClaudeCodeUserAgent 用于将渠道模型请求伪装为 Claude Code 客户端。
ClaudeCodeUserAgent = "claude-cli/2.1.19 (external, sdk-cli)"
// AnthropicClaudeCodeUserAgent 用于 Anthropic provider 的 Claude Code UA 兼容。
AnthropicClaudeCodeUserAgent = "claude-cli/1.0.25"
)
@@ -0,0 +1,15 @@
package promptengine
// ContentPart 表示一条消息中的结构化内容块。
type ContentPart struct {
Type string `json:"type"`
Text string `json:"text,omitempty"`
Image *ImageContent `json:"image,omitempty"`
}
// ImageContent 表示消息中携带的一张图片。
type ImageContent struct {
MIMEType string `json:"mime_type,omitempty"`
Path string `json:"path,omitempty"`
Data []byte `json:"data,omitempty"`
}
+2
View File
@@ -0,0 +1,2 @@
// Package promptengine 负责把静态 prompt 资产、会话状态与外部结果编译成模型请求输入。
package promptengine
File diff suppressed because it is too large Load Diff
+518
View File
@@ -0,0 +1,518 @@
package promptengine
import (
"encoding/json"
"fmt"
"strings"
"google.golang.org/protobuf/proto"
"google.golang.org/protobuf/reflect/protoreflect"
"cursor/gen/agentv1"
)
// BuildUserQueryReplayMessage 构造一条可直接回放给模型的用户消息。
func BuildUserQueryReplayMessage(text string) (Message, bool) {
return buildUserReplayMessage(strings.TrimSpace(text), nil)
}
// BuildUserMessageReplayMessage 把包含 selected_context 的用户消息还原为 replay message。
func BuildUserMessageReplayMessage(userMessage *agentv1.UserMessage) (Message, bool) {
if userMessage == nil {
return Message{}, false
}
return buildUserReplayMessage(strings.TrimSpace(userMessage.GetText()), userMessage.GetSelectedContext())
}
func buildUserReplayMessage(text string, selectedContext *agentv1.SelectedContext) (Message, bool) {
images := buildSelectedImageContentParts(selectedContext)
sections := make([]string, 0, 4)
if text != "" {
sections = append(sections, formatMessageText(fmt.Sprintf("<user_query>\n%s\n</user_query>", text)))
}
if ideState := buildSelectedIDEStatePromptSection(selectedContext); ideState != "" {
sections = append(sections, ideState)
}
if selectedFiles := buildSelectedFilesPromptSection(selectedContext); selectedFiles != "" {
sections = append(sections, selectedFiles)
}
content := strings.TrimSpace(strings.Join(sections, "\n\n"))
if content == "" && len(images) == 0 {
return Message{}, false
}
if len(images) == 0 {
return Message{
Role: "user",
Content: content,
}, true
}
parts := make([]ContentPart, 0, len(images)+1)
if content != "" {
parts = append(parts, ContentPart{
Type: "text",
Text: content,
})
}
parts = append(parts, images...)
return Message{
Role: "user",
Content: content,
ContentParts: parts,
}, true
}
func buildSelectedIDEStatePromptSection(selectedContext *agentv1.SelectedContext) string {
if selectedContext == nil || selectedContext.GetInvocationContext() == nil {
return ""
}
ideState := selectedContext.GetInvocationContext().GetIdeState()
if ideState == nil {
return ""
}
sections := make([]string, 0, 2)
if visible := buildIDEStateFilesPromptSection("visible_files", ideState.GetVisibleFiles()); visible != "" {
sections = append(sections, visible)
}
if recent := buildIDEStateFilesPromptSection("recently_viewed_files", ideState.GetRecentlyViewedFiles()); recent != "" {
sections = append(sections, recent)
}
return strings.TrimSpace(strings.Join(sections, "\n\n"))
}
func buildIDEStateFilesPromptSection(tag string, files []*agentv1.InvocationContext_IdeState_File) string {
if len(files) == 0 {
return ""
}
entries := make([]string, 0, len(files))
for _, file := range files {
if file == nil {
continue
}
attrs := make([]string, 0, 4)
if path := strings.TrimSpace(file.GetPath()); path != "" {
attrs = append(attrs, fmt.Sprintf(`path="%s"`, escapePromptXML(path)))
}
if relativePath := strings.TrimSpace(file.GetRelativePath()); relativePath != "" {
attrs = append(attrs, fmt.Sprintf(`relative_path="%s"`, escapePromptXML(relativePath)))
}
if cursor := file.GetCursorPosition(); cursor != nil && cursor.GetLine() > 0 {
attrs = append(attrs, fmt.Sprintf(`cursor_line="%d"`, cursor.GetLine()))
}
if totalLines := file.GetTotalLines(); totalLines > 0 {
attrs = append(attrs, fmt.Sprintf(`total_lines="%d"`, totalLines))
}
if len(attrs) == 0 {
continue
}
entries = append(entries, "<file "+strings.Join(attrs, " ")+" />")
}
if len(entries) == 0 {
return ""
}
return fmt.Sprintf("<%s>\n%s\n</%s>", tag, strings.Join(entries, "\n"), tag)
}
func buildSelectedFilesPromptSection(selectedContext *agentv1.SelectedContext) string {
if selectedContext == nil || len(selectedContext.GetFiles()) == 0 {
return ""
}
entries := make([]string, 0, len(selectedContext.GetFiles()))
for _, file := range selectedContext.GetFiles() {
if file == nil || strings.TrimSpace(file.GetContent()) == "" {
continue
}
attrs := make([]string, 0, 2)
if path := strings.TrimSpace(file.GetPath()); path != "" {
attrs = append(attrs, fmt.Sprintf(`path="%s"`, escapePromptXML(path)))
}
if relativePath := strings.TrimSpace(file.GetRelativePath()); relativePath != "" {
attrs = append(attrs, fmt.Sprintf(`relative_path="%s"`, escapePromptXML(relativePath)))
}
if len(attrs) == 0 {
continue
}
entries = append(entries, "<file "+strings.Join(attrs, " ")+">\n"+file.GetContent()+"\n</file>")
}
if len(entries) == 0 {
return ""
}
return "<selected_files>\n" + strings.Join(entries, "\n\n") + "\n</selected_files>"
}
func buildSelectedImageContentParts(selectedContext *agentv1.SelectedContext) []ContentPart {
if selectedContext == nil {
return nil
}
parts := make([]ContentPart, 0, len(selectedContext.GetSelectedImages()))
for _, image := range selectedContext.GetSelectedImages() {
if image == nil {
continue
}
data := image.GetData()
if len(data) == 0 {
data = image.GetBlobIdWithData().GetData()
}
if len(data) == 0 && strings.TrimSpace(image.GetPath()) == "" {
continue
}
parts = append(parts, ContentPart{
Type: "image",
Image: &ImageContent{
MIMEType: strings.TrimSpace(image.GetMimeType()),
Path: strings.TrimSpace(image.GetPath()),
Data: data,
},
})
}
return parts
}
// EncodeReplayMessages 把 canonical replay message 编码为 root_prompt_messages_json。
func EncodeReplayMessages(messages []Message) ([][]byte, error) {
if len(messages) == 0 {
return nil, nil
}
encoded := make([][]byte, 0, len(messages))
for _, message := range messages {
payload, err := marshalReplayMessage(message)
if err != nil {
return nil, err
}
encoded = append(encoded, payload)
}
return encoded, nil
}
func marshalReplayMessage(message Message) ([]byte, error) {
payload := map[string]any{
"role": message.Role,
"content": message.Content,
}
if len(message.ContentParts) > 0 {
payload["content_parts"] = message.ContentParts
}
if strings.TrimSpace(message.ReasoningContent) != "" || (strings.TrimSpace(message.Role) == "assistant" && len(message.ToolCalls) > 0) {
payload["reasoning_content"] = message.ReasoningContent
}
if strings.TrimSpace(message.ReasoningSignature) != "" {
payload["reasoning_signature"] = message.ReasoningSignature
}
if strings.TrimSpace(message.ReasoningSignatureSource) != "" {
payload["reasoning_signature_source"] = strings.TrimSpace(message.ReasoningSignatureSource)
}
if strings.TrimSpace(message.OpenAIResponsesReasoningID) != "" {
payload["openai_responses_reasoning_id"] = strings.TrimSpace(message.OpenAIResponsesReasoningID)
}
if strings.TrimSpace(message.OpenAIResponsesReasoningStatus) != "" {
payload["openai_responses_reasoning_status"] = strings.TrimSpace(message.OpenAIResponsesReasoningStatus)
}
if len(message.OpenAIResponsesReasoningSummary) > 0 {
payload["openai_responses_reasoning_summary"] = json.RawMessage(append([]byte(nil), message.OpenAIResponsesReasoningSummary...))
}
if len(message.ToolCalls) > 0 {
payload["tool_calls"] = message.ToolCalls
}
if strings.TrimSpace(message.ToolCallID) != "" {
payload["tool_call_id"] = message.ToolCallID
}
if strings.TrimSpace(message.Name) != "" {
payload["name"] = message.Name
}
return json.Marshal(payload)
}
// DecodeReplayMessages 从 root_prompt_messages_json 解码 canonical replay message。
func DecodeReplayMessages(rawItems [][]byte) ([]Message, error) {
if len(rawItems) == 0 {
return nil, nil
}
messages := make([]Message, 0, len(rawItems))
for _, raw := range rawItems {
if len(raw) == 0 {
continue
}
var message Message
if err := json.Unmarshal(raw, &message); err != nil {
return nil, err
}
if strings.TrimSpace(message.Role) == "" {
continue
}
messages = append(messages, message)
}
return messages, nil
}
// BuildReplayMessagesFromPendingAssistantOutputs 把 pending assistant raw 还原为 canonical replay message。
func BuildReplayMessagesFromPendingAssistantOutputs(rawValues []string) []Message {
if len(rawValues) == 0 {
return nil
}
messages := make([]Message, 0, len(rawValues)*3)
for _, raw := range rawValues {
messages = append(messages, buildMessagesFromPendingAssistantRaw(raw)...)
}
return messages
}
// BuildLegacyMessagesFromConversationStep 使用 legacy XML 文本形状回放单个 step。
func BuildLegacyMessagesFromConversationStep(step *agentv1.ConversationStep) []Message {
if step == nil {
return nil
}
switch item := step.GetMessage().(type) {
case *agentv1.ConversationStep_AssistantMessage:
text := strings.TrimSpace(item.AssistantMessage.GetText())
if text == "" {
return nil
}
return []Message{{Role: "assistant", Content: formatMessageText(text)}}
case *agentv1.ConversationStep_ThinkingMessage:
text := strings.TrimSpace(item.ThinkingMessage.GetText())
if text == "" {
return nil
}
return []Message{{
Role: "assistant",
Content: formatMessageText(fmt.Sprintf("<thinking>\n%s\n</thinking>", text)),
}}
case *agentv1.ConversationStep_ToolCall:
text := strings.TrimSpace(compactProtoJSON(item.ToolCall))
if text == "" {
return nil
}
return []Message{{
Role: "assistant",
Content: formatMessageText(fmt.Sprintf("<tool_call>\n%s\n</tool_call>", text)),
}}
default:
return nil
}
}
// BuildToolCallReplayMessages 把已完成的 ToolCall step 还原为 native assistant/tool replay message。
func BuildToolCallReplayMessages(toolCallID string, toolCall *agentv1.ToolCall) ([]Message, bool) {
descriptor, toolResult, ok := extractToolCallReplay(toolCallID, toolCall)
if !ok {
return nil, false
}
return []Message{
{
Role: "assistant",
Content: "",
ToolCalls: []ToolCallDescriptor{descriptor},
},
toolResult,
}, true
}
// BuildToolResultReplayMessage 从已完成的 ToolCall 中提取 tool replay message,
// 用于 history 已经单独记录过 assistant tool_call 时仅回放真实工具结果。
func BuildToolResultReplayMessage(toolCallID string, toolCall *agentv1.ToolCall) (Message, bool) {
if toolCall == nil || strings.TrimSpace(toolCallID) == "" {
return Message{}, false
}
shape, ok := extractToolCallReplayShape(toolCall)
if !ok || !shape.HasResult {
return Message{}, false
}
return Message{
Role: "tool",
Content: shape.ResultJSON,
ToolCallID: strings.TrimSpace(toolCallID),
Name: shape.ToolName,
}, true
}
// BuildAssistantToolCallReplayMessage 把未完成或已完成的 ToolCall 还原为 assistant tool-call replay message。
func BuildAssistantToolCallReplayMessage(toolCallID string, toolCall *agentv1.ToolCall) (Message, bool) {
descriptor, ok := BuildToolCallReplayDescriptor(toolCallID, toolCall)
if !ok {
return Message{}, false
}
return Message{
Role: "assistant",
Content: "",
ToolCalls: []ToolCallDescriptor{descriptor},
}, true
}
// BuildToolCallReplayDescriptor 从 ToolCall proto 提取 assistant replay 所需的工具调用描述。
func BuildToolCallReplayDescriptor(toolCallID string, toolCall *agentv1.ToolCall) (ToolCallDescriptor, bool) {
if toolCall == nil || strings.TrimSpace(toolCallID) == "" {
return ToolCallDescriptor{}, false
}
shape, ok := extractToolCallReplayShape(toolCall)
if !ok {
return ToolCallDescriptor{}, false
}
return ToolCallDescriptor{
ID: strings.TrimSpace(toolCallID),
Type: "function",
Function: ToolCallFunctionShape{
Name: shape.ToolName,
Arguments: firstNonEmpty(shape.ArgsJSON, "{}"),
},
}, true
}
func extractToolCallReplay(toolCallID string, toolCall *agentv1.ToolCall) (ToolCallDescriptor, Message, bool) {
descriptor, ok := BuildToolCallReplayDescriptor(toolCallID, toolCall)
if !ok {
return ToolCallDescriptor{}, Message{}, false
}
toolResult, ok := BuildToolResultReplayMessage(toolCallID, toolCall)
if !ok {
return ToolCallDescriptor{}, Message{}, false
}
return descriptor, toolResult, true
}
type toolCallReplayShape struct {
ToolName string
ArgsJSON string
ResultJSON string
HasResult bool
}
func extractToolCallReplayShape(toolCall *agentv1.ToolCall) (toolCallReplayShape, bool) {
if toolCall == nil {
return toolCallReplayShape{}, false
}
value := toolCall.ProtoReflect()
oneof := value.Descriptor().Oneofs().ByName("tool")
if oneof == nil {
return toolCallReplayShape{}, false
}
selected := value.WhichOneof(oneof)
if selected == nil {
return toolCallReplayShape{}, false
}
selectedValue := value.Get(selected)
if !selectedValue.IsValid() {
return toolCallReplayShape{}, false
}
selectedMessage := selectedValue.Message()
if !selectedMessage.IsValid() {
return toolCallReplayShape{}, false
}
argsJSON, _ := extractReplayFieldJSON(selectedMessage, "args")
resultJSON, hasResult := extractReplayFieldJSON(selectedMessage, "result")
toolName := canonicalReplayToolName(string(selected.Name()), string(selectedMessage.Descriptor().Name()), argsJSON, resultJSON)
if toolName == "" {
return toolCallReplayShape{}, false
}
return toolCallReplayShape{
ToolName: toolName,
ArgsJSON: firstNonEmpty(argsJSON, "{}"),
ResultJSON: resultJSON,
HasResult: hasResult,
}, true
}
func extractReplayFieldJSON(message protoreflect.Message, fieldName string) (string, bool) {
if !message.IsValid() {
return "", false
}
field := message.Descriptor().Fields().ByName(protoreflect.Name(fieldName))
if field == nil || !message.Has(field) {
return "", false
}
value := message.Get(field)
if !value.IsValid() {
return "", false
}
child := value.Message()
if !child.IsValid() {
return "", false
}
item, ok := child.Interface().(proto.Message)
if !ok {
return "", false
}
return compactProtoJSON(item), true
}
func canonicalReplayToolName(fieldName string, messageName string, argsJSON string, resultJSON string) string {
switch strings.TrimSpace(fieldName) {
case "mcp_tool_call":
return "CallMcpTool"
case "read_mcp_resource_tool_call":
return "FetchMcpResource"
case "update_todos_tool_call":
return "TodoWrite"
case "read_todos_tool_call":
return "ReadTodos"
case "sem_search_tool_call":
return "SemanticSearch"
case "edit_tool_call":
return replayEditToolName(argsJSON, resultJSON)
}
trimmed := strings.TrimSuffix(strings.TrimSpace(messageName), "ToolCall")
return strings.TrimSpace(trimmed)
}
func replayEditToolName(argsJSON string, resultJSON string) string {
if replayEditResultLooksLikeStructuredEdit(resultJSON) {
return "Edit"
}
if editArgsIndicateWrite(argsJSON) {
return "Write"
}
return "Edit"
}
func editArgsIndicateWrite(argsJSON string) bool {
trimmed := strings.TrimSpace(argsJSON)
if trimmed == "" || trimmed == "{}" || trimmed == "null" {
return false
}
var args map[string]any
if err := json.Unmarshal([]byte(trimmed), &args); err != nil {
return false
}
for _, key := range []string{"stream_content", "streamContent"} {
if _, ok := args[key]; !ok {
continue
}
switch args[key].(type) {
case string:
return true
case nil:
return true
default:
return true
}
}
return false
}
func replayEditResultLooksLikeStructuredEdit(resultJSON string) bool {
trimmed := strings.TrimSpace(resultJSON)
if trimmed == "" || trimmed == "{}" || trimmed == "null" {
return false
}
var payload map[string]any
if err := json.Unmarshal([]byte(trimmed), &payload); err != nil {
return false
}
success, ok := payload["success"].(map[string]any)
if !ok || len(success) == 0 {
return false
}
if _, ok := success["beforeFullFileContent"]; ok {
return true
}
if _, ok := success["before_full_file_content"]; ok {
return true
}
if _, ok := success["diffString"]; ok {
return true
}
if _, ok := success["diff_string"]; ok {
return true
}
return false
}
+2
View File
@@ -0,0 +1,2 @@
// Package protocol 负责协议层解码、摘要与上行消息种类判断。
package protocol
+224
View File
@@ -0,0 +1,224 @@
// inbound.go 实现上行协议的解码、摘要与命令类型识别。
package protocol
import (
"encoding/hex"
"fmt"
"strings"
"unicode/utf8"
"cursor/gen/agentv1"
"cursor/gen/aiserverv1"
"cursor/internal/backend/agent/core"
"google.golang.org/protobuf/proto"
)
// ReadAppendRequestID 从 BidiAppendRequest 中读取 request_id 文本。
func ReadAppendRequestID(input *aiserverv1.BidiAppendRequest) string {
if input == nil {
return ""
}
return ReadBidiRequestID(input.GetRequestId())
}
// ReadBidiRequestID 从 BidiRequestId 结构中提取并去除首尾空白。
func ReadBidiRequestID(input *aiserverv1.BidiRequestId) string {
if input == nil {
return ""
}
return strings.TrimSpace(input.GetRequestId())
}
// NormalizeRequestID 规范化请求标识并去除首尾空白。
func NormalizeRequestID(requestID string) string {
return strings.TrimSpace(requestID)
}
// DecodeAgentClientMessage 解析 hex 文本为 AgentClientMessage,并返回消息类型标签。
func DecodeAgentClientMessage(hexData string) (*agentv1.AgentClientMessage, string, error) {
trimmed := strings.TrimSpace(hexData)
if trimmed == "" {
return nil, "", nil
}
payload, err := hex.DecodeString(trimmed)
if err != nil {
return nil, "", fmt.Errorf("bidi append data is not valid hex: %w", err)
}
clientMessage := &agentv1.AgentClientMessage{}
if err := proto.Unmarshal(payload, clientMessage); err != nil {
return nil, "", fmt.Errorf("decode agent client message failed: %w", err)
}
return clientMessage, detectClientMessageKind(clientMessage), nil
}
// MapClientMessageToCommandKind 将上行协议消息映射为运行时命令类型。
func MapClientMessageToCommandKind(message *agentv1.AgentClientMessage, clientKind string) (runtimecore.CommandKind, error) {
switch strings.TrimSpace(clientKind) {
case "run_request":
return runtimecore.CommandKindRunRequested, nil
case "prewarm_request":
return runtimecore.CommandKindPrewarmRequested, nil
case "conversation_action":
if message == nil || message.GetConversationAction() == nil {
return "", fmt.Errorf("conversation_action payload is required")
}
switch message.GetConversationAction().GetAction().(type) {
case *agentv1.ConversationAction_CancelAction:
return runtimecore.CommandKindCancelRequested, nil
case *agentv1.ConversationAction_UserMessageAction,
*agentv1.ConversationAction_ResumeAction,
*agentv1.ConversationAction_SummarizeAction,
*agentv1.ConversationAction_ShellCommandAction,
*agentv1.ConversationAction_StartPlanAction,
*agentv1.ConversationAction_ExecutePlanAction,
*agentv1.ConversationAction_AsyncAskQuestionCompletionAction,
*agentv1.ConversationAction_CancelSubagentAction,
*agentv1.ConversationAction_BackgroundShellAction,
*agentv1.ConversationAction_BackgroundTaskCompletionAction:
return runtimecore.CommandKindConversationActionRecordOnly, nil
default:
return "", fmt.Errorf("unsupported conversation_action payload")
}
case "exec_client_message":
return runtimecore.CommandKindExecClientMessage, nil
case "interaction_response":
return runtimecore.CommandKindInteractionResponse, nil
case "exec_client_control_message":
return runtimecore.CommandKindExecClientControlMessage, nil
case "client_heartbeat":
return runtimecore.CommandKindClientHeartbeat, nil
case "kv_client_message":
return runtimecore.CommandKindKVClientMessage, nil
default:
return "", fmt.Errorf("unsupported client message kind: %s", clientKind)
}
}
// IsResumeRunRequest 判断当前消息是否为带 `resume_action` 的 `run_request`。
func IsResumeRunRequest(message *agentv1.AgentClientMessage) bool {
if message == nil || message.GetRunRequest() == nil {
return false
}
action := message.GetRunRequest().GetAction()
if action == nil {
return false
}
_, ok := action.GetAction().(*agentv1.ConversationAction_ResumeAction)
return ok
}
// BuildClientHistoryEntry 将消息类型与负载摘要拼接成会话历史记录文本。
func BuildClientHistoryEntry(kind string, message *agentv1.AgentClientMessage) string {
normalizedKind := strings.TrimSpace(kind)
if normalizedKind == "" {
normalizedKind = "unknown"
}
payload := extractClientMessagePayload(message)
summary := summarizePayload(payload)
if summary == "" {
if normalizedKind == "unknown" {
return ""
}
return normalizedKind
}
return fmt.Sprintf("%s:%s", normalizedKind, summary)
}
// detectClientMessageKind 判断 oneof message 当前承载的消息分支类型。
func detectClientMessageKind(message *agentv1.AgentClientMessage) string {
if message == nil || message.GetMessage() == nil {
return ""
}
switch message.GetMessage().(type) {
case *agentv1.AgentClientMessage_RunRequest:
return "run_request"
case *agentv1.AgentClientMessage_PrewarmRequest:
return "prewarm_request"
case *agentv1.AgentClientMessage_ConversationAction:
return "conversation_action"
case *agentv1.AgentClientMessage_ExecClientMessage:
return "exec_client_message"
case *agentv1.AgentClientMessage_InteractionResponse:
return "interaction_response"
case *agentv1.AgentClientMessage_ExecClientControlMessage:
return "exec_client_control_message"
case *agentv1.AgentClientMessage_ClientHeartbeat:
return "client_heartbeat"
case *agentv1.AgentClientMessage_KvClientMessage:
return "kv_client_message"
default:
return ""
}
}
// extractClientMessagePayload 从 oneof 分支中提取原始 bytes 负载。
func extractClientMessagePayload(message *agentv1.AgentClientMessage) []byte {
if message == nil || message.GetMessage() == nil {
return nil
}
switch item := message.GetMessage().(type) {
case *agentv1.AgentClientMessage_RunRequest:
return marshalProtoMessage(item.RunRequest)
case *agentv1.AgentClientMessage_PrewarmRequest:
return marshalProtoMessage(item.PrewarmRequest)
case *agentv1.AgentClientMessage_ConversationAction:
return marshalProtoMessage(item.ConversationAction)
case *agentv1.AgentClientMessage_ExecClientMessage:
return marshalProtoMessage(item.ExecClientMessage)
case *agentv1.AgentClientMessage_InteractionResponse:
return marshalProtoMessage(item.InteractionResponse)
case *agentv1.AgentClientMessage_ExecClientControlMessage:
return marshalProtoMessage(item.ExecClientControlMessage)
case *agentv1.AgentClientMessage_ClientHeartbeat:
return marshalProtoMessage(item.ClientHeartbeat)
case *agentv1.AgentClientMessage_KvClientMessage:
return marshalProtoMessage(item.KvClientMessage)
default:
return nil
}
}
// marshalProtoMessage 将 proto 消息重新编码为 bytes,用于调试摘要展示。
func marshalProtoMessage(message proto.Message) []byte {
if message == nil {
return nil
}
payload, err := proto.Marshal(message)
if err != nil {
return nil
}
return payload
}
// summarizePayload 生成可读摘要:优先文本,无法直接读时回退为 hex 片段。
func summarizePayload(payload []byte) string {
if len(payload) == 0 {
return ""
}
if utf8.Valid(payload) {
text := strings.TrimSpace(string(payload))
text = strings.ReplaceAll(text, "\n", " ")
text = strings.ReplaceAll(text, "\r", " ")
text = strings.TrimSpace(text)
if text != "" {
return truncateText(text, 120)
}
}
return "hex:" + truncateText(hex.EncodeToString(payload), 120)
}
// truncateText 按 rune 数截断文本,避免在多字节字符中间截断导致乱码。
func truncateText(text string, maxRunes int) string {
if maxRunes <= 0 {
return ""
}
runes := []rune(text)
if len(runes) <= maxRunes {
return text
}
return string(runes[:maxRunes]) + "..."
}
+2
View File
@@ -0,0 +1,2 @@
// Package step 负责 assistant 输出记录的整理、解析与最小构造。
package step
+237
View File
@@ -0,0 +1,237 @@
// recorder.go 实现 pending assistant 输出记录的解析与构造。
package step
import (
"encoding/json"
"strings"
"cursor/internal/backend/agent/core"
)
const (
// defaultTextPreviewLimit 表示文本摘要允许保留的最大 rune 数。
defaultTextPreviewLimit = 120
)
// assistantMessage 表示 `pending_tool_calls` 中常见的 assistant message 结构。
type assistantMessage struct {
// ID 是 message 级别的标识,当前抓包中常见值为 1。
ID string `json:"id,omitempty"`
// Role 是当前 message 的角色,当前常见值为 assistant。
Role string `json:"role,omitempty"`
// Content 保存该 message 的内容块列表。
Content []assistantContent `json:"content,omitempty"`
}
// assistantContent 表示 assistant message 内的单个内容块。
type assistantContent struct {
// Type 表示内容块类型,例如 text、reasoning 或 tool-call。
Type string `json:"type,omitempty"`
// Text 保存文本内容块文本。
Text string `json:"text,omitempty"`
// ToolCallID 保存工具调用标识。
ToolCallID string `json:"toolCallId,omitempty"`
// ToolName 保存工具名称。
ToolName string `json:"toolName,omitempty"`
// Args 保存工具调用参数原文。
Args json.RawMessage `json:"args,omitempty"`
// Result 保存工具调用结果原文。
Result json.RawMessage `json:"result,omitempty"`
}
// Recorder 负责解析与构造 pending assistant 输出记录。
type Recorder struct {
}
// StepRecorder 定义运行时依赖的 step 记录接口。
type StepRecorder interface {
// ParsePendingAssistantOutputs 解析一组原始 pending assistant 输出记录。
ParsePendingAssistantOutputs(rawValues []string) []runtimecore.PendingAssistantOutput
// ParsePendingAssistantOutput 解析单条原始 assistant 输出记录。
ParsePendingAssistantOutput(raw string) runtimecore.PendingAssistantOutput
// BuildTextAssistantOutput 构造一条只包含文本的 assistant 输出记录。
BuildTextAssistantOutput(text string) (string, runtimecore.PendingAssistantOutput, error)
// StartAssistantOutput 创建一个新的 assistant 输出构造器。
StartAssistantOutput() *AssistantOutputBuilder
}
// NewRecorder 创建 assistant 输出记录整理器。
func NewRecorder() *Recorder {
return &Recorder{}
}
// AssistantOutputBuilder 表示一条 assistant 输出记录的构造器。
type AssistantOutputBuilder struct {
// message 保存当前正在构造的原始 assistant message。
message assistantMessage
}
// ParsePendingAssistantOutputs 解析一组原始 `pending_tool_calls` 字符串。
func (recorder *Recorder) ParsePendingAssistantOutputs(rawValues []string) []runtimecore.PendingAssistantOutput {
if len(rawValues) == 0 {
return nil
}
outputs := make([]runtimecore.PendingAssistantOutput, 0, len(rawValues))
for _, raw := range rawValues {
outputs = append(outputs, recorder.ParsePendingAssistantOutput(raw))
}
return outputs
}
// ParsePendingAssistantOutput 解析单条原始 assistant 输出记录。
func (recorder *Recorder) ParsePendingAssistantOutput(raw string) runtimecore.PendingAssistantOutput {
output := runtimecore.PendingAssistantOutput{
RawMessage: strings.TrimSpace(raw),
}
if output.RawMessage == "" {
return output
}
var message assistantMessage
if err := json.Unmarshal([]byte(output.RawMessage), &message); err != nil {
output.TextPreview = truncateText(output.RawMessage, defaultTextPreviewLimit)
return output
}
output.Role = strings.TrimSpace(message.Role)
output.ContentKinds = make([]string, 0, len(message.Content))
output.ToolCallIDs = make([]string, 0, len(message.Content))
output.ToolNames = make([]string, 0, len(message.Content))
textParts := make([]string, 0, len(message.Content))
for _, part := range message.Content {
kind := strings.TrimSpace(part.Type)
if kind == "" {
kind = "unknown"
}
output.ContentKinds = append(output.ContentKinds, kind)
if trimmedToolCallID := strings.TrimSpace(part.ToolCallID); trimmedToolCallID != "" {
output.ToolCallIDs = append(output.ToolCallIDs, trimmedToolCallID)
}
if trimmedToolName := strings.TrimSpace(part.ToolName); trimmedToolName != "" {
output.ToolNames = append(output.ToolNames, trimmedToolName)
}
if trimmedText := strings.TrimSpace(part.Text); trimmedText != "" {
textParts = append(textParts, trimmedText)
}
}
output.TextPreview = truncateText(strings.Join(textParts, "\n"), defaultTextPreviewLimit)
return output
}
// BuildTextAssistantOutput 构造一条只包含文本的 assistant 输出记录。
func (recorder *Recorder) BuildTextAssistantOutput(text string) (string, runtimecore.PendingAssistantOutput, error) {
message := assistantMessage{
ID: "1",
Role: "assistant",
Content: []assistantContent{
{
Type: "text",
Text: text,
},
},
}
payload, err := json.Marshal(message)
if err != nil {
return "", runtimecore.PendingAssistantOutput{}, err
}
raw := string(payload)
return raw, recorder.ParsePendingAssistantOutput(raw), nil
}
// StartAssistantOutput 创建一个新的 assistant 输出构造器。
func (recorder *Recorder) StartAssistantOutput() *AssistantOutputBuilder {
return &AssistantOutputBuilder{
message: assistantMessage{
ID: "1",
Role: "assistant",
Content: make([]assistantContent, 0, 4),
},
}
}
// AppendTextDelta 追加一段文本内容。
func (builder *AssistantOutputBuilder) AppendTextDelta(text string) {
if builder == nil {
return
}
builder.message.Content = append(builder.message.Content, assistantContent{
Type: "text",
Text: text,
})
}
// AppendReasoningDelta 追加一段推理内容,供 reasoning 模型在续跑时回放。
func (builder *AssistantOutputBuilder) AppendReasoningDelta(text string) {
if builder == nil {
return
}
builder.message.Content = append(builder.message.Content, assistantContent{
Type: "reasoning",
Text: text,
})
}
// OpenToolCall 追加一个尚未完成的工具调用块。
func (builder *AssistantOutputBuilder) OpenToolCall(toolCall runtimecore.ToolInvocation) {
if builder == nil {
return
}
builder.message.Content = append(builder.message.Content, assistantContent{
Type: "tool-call",
ToolCallID: strings.TrimSpace(toolCall.CallID),
ToolName: strings.TrimSpace(toolCall.ToolName),
Args: append(json.RawMessage(nil), toolCall.ArgsJSON...),
})
}
// CompleteToolCall 为指定工具调用补充结果内容。
func (builder *AssistantOutputBuilder) CompleteToolCall(toolCallID string, resultJSON []byte) {
if builder == nil {
return
}
for index := range builder.message.Content {
if builder.message.Content[index].Type != "tool-call" {
continue
}
if strings.TrimSpace(builder.message.Content[index].ToolCallID) != strings.TrimSpace(toolCallID) {
continue
}
builder.message.Content[index].Result = append(json.RawMessage(nil), resultJSON...)
return
}
}
// SnapshotRaw 输出当前 builder 的原始 JSON 和解析结果。
func (builder *AssistantOutputBuilder) SnapshotRaw(recorder *Recorder) (string, runtimecore.PendingAssistantOutput, error) {
if builder == nil {
return "", runtimecore.PendingAssistantOutput{}, nil
}
payload, err := json.Marshal(builder.message)
if err != nil {
return "", runtimecore.PendingAssistantOutput{}, err
}
raw := string(payload)
if recorder == nil {
recorder = NewRecorder()
}
return raw, recorder.ParsePendingAssistantOutput(raw), nil
}
// truncateText 按 rune 数截断文本,避免在多字节字符中间截断。
func truncateText(text string, maxRunes int) string {
trimmed := strings.TrimSpace(text)
if maxRunes <= 0 || trimmed == "" {
return ""
}
runes := []rune(trimmed)
if len(runes) <= maxRunes {
return trimmed
}
return string(runes[:maxRunes]) + "..."
}
File diff suppressed because it is too large Load Diff
+356
View File
@@ -0,0 +1,356 @@
package forwarder
import (
"context"
"encoding/json"
"net/http"
"strings"
"time"
"connectrpc.com/connect"
"cursor/gen/aiserverv1"
"cursor/gen/aiserverv1/aiserverv1connect"
)
type usageLookupRecord struct {
InputTokens int64
OutputTokens int64
CreatedAt time.Time
}
const (
dashboardServiceGetTokenUsageProcedure = "/aiserver.v1.DashboardService/GetTokenUsage"
dashboardServiceGetGlassEarlyPreviewEnrollmentProcedure = "/aiserver.v1.DashboardService/GetGlassEarlyPreviewEnrollment"
)
func newAIHandler(service *Service) http.Handler {
mux := http.NewServeMux()
mux.Handle(
dashboardServiceGetTokenUsageProcedure,
connect.NewUnaryHandler(dashboardServiceGetTokenUsageProcedure, service.GetTokenUsage),
)
mux.Handle(
dashboardServiceGetGlassEarlyPreviewEnrollmentProcedure,
connect.NewUnaryHandler(dashboardServiceGetGlassEarlyPreviewEnrollmentProcedure, service.GetGlassEarlyPreviewEnrollment),
)
mux.Handle(
aiserverv1connect.AiServiceCountTokensProcedure,
connect.NewUnaryHandler(aiserverv1connect.AiServiceCountTokensProcedure, service.CountTokens),
)
mux.Handle(
aiserverv1connect.AiServiceGetThoughtAnnotationProcedure,
connect.NewUnaryHandler(aiserverv1connect.AiServiceGetThoughtAnnotationProcedure, service.GetThoughtAnnotation),
)
mux.Handle(
aiserverv1connect.AiServiceWriteGitCommitMessageProcedure,
connect.NewUnaryHandler(aiserverv1connect.AiServiceWriteGitCommitMessageProcedure, service.WriteGitCommitMessage),
)
mux.Handle(
aiserverv1connect.AiServiceCreateExperimentalIndexProcedure,
connect.NewUnaryHandler(aiserverv1connect.AiServiceCreateExperimentalIndexProcedure, service.CreateExperimentalIndex),
)
mux.Handle(
aiserverv1connect.AiServiceListExperimentalIndexFilesProcedure,
connect.NewUnaryHandler(aiserverv1connect.AiServiceListExperimentalIndexFilesProcedure, service.ListExperimentalIndexFiles),
)
mux.Handle(
aiserverv1connect.AiServiceListenExperimentalIndexProcedure,
connect.NewServerStreamHandler(aiserverv1connect.AiServiceListenExperimentalIndexProcedure, service.ListenExperimentalIndex),
)
mux.Handle(
aiserverv1connect.AiServiceRegisterFileToIndexProcedure,
connect.NewUnaryHandler(aiserverv1connect.AiServiceRegisterFileToIndexProcedure, service.RegisterFileToIndex),
)
mux.Handle(
aiserverv1connect.AiServiceSetupIndexDependenciesProcedure,
connect.NewUnaryHandler(aiserverv1connect.AiServiceSetupIndexDependenciesProcedure, service.SetupIndexDependencies),
)
mux.Handle(
aiserverv1connect.AiServiceComputeIndexTopoSortProcedure,
connect.NewUnaryHandler(aiserverv1connect.AiServiceComputeIndexTopoSortProcedure, service.ComputeIndexTopoSort),
)
mux.Handle(
aiserverv1connect.AiServiceDocumentationQueryProcedure,
connect.NewUnaryHandler(aiserverv1connect.AiServiceDocumentationQueryProcedure, service.DocumentationQuery),
)
mux.Handle(
aiserverv1connect.AiServiceAvailableDocsProcedure,
connect.NewUnaryHandler(aiserverv1connect.AiServiceAvailableDocsProcedure, service.AvailableDocs),
)
mux.Handle(
aiserverv1connect.AiServiceKnowledgeBaseAddProcedure,
connect.NewUnaryHandler(aiserverv1connect.AiServiceKnowledgeBaseAddProcedure, service.KnowledgeBaseAdd),
)
mux.Handle(
aiserverv1connect.AiServiceKnowledgeBaseListProcedure,
connect.NewUnaryHandler(aiserverv1connect.AiServiceKnowledgeBaseListProcedure, service.KnowledgeBaseList),
)
mux.Handle(
aiserverv1connect.AiServiceKnowledgeBaseRemoveProcedure,
connect.NewUnaryHandler(aiserverv1connect.AiServiceKnowledgeBaseRemoveProcedure, service.KnowledgeBaseRemove),
)
mux.Handle(
aiserverv1connect.AiServiceKnowledgeBaseUpdateProcedure,
connect.NewUnaryHandler(aiserverv1connect.AiServiceKnowledgeBaseUpdateProcedure, service.KnowledgeBaseUpdate),
)
mux.Handle(
aiserverv1connect.AiServiceFetchRelevantKnowledgeForConversationProcedure,
connect.NewUnaryHandler(aiserverv1connect.AiServiceFetchRelevantKnowledgeForConversationProcedure, service.FetchRelevantKnowledgeForConversation),
)
mux.Handle("/", http.NotFoundHandler())
return mux
}
func (service *Service) GetThoughtAnnotation(_ context.Context, req *connect.Request[aiserverv1.GetThoughtAnnotationRequest]) (*connect.Response[aiserverv1.GetThoughtAnnotationResponse], error) {
requestID := strings.TrimSpace(req.Msg.GetRequestId())
thought, ok, err := service.lookupThoughtAnnotation(requestID)
if err != nil {
return nil, connect.NewError(connect.CodeInternal, err)
}
if !ok {
return connect.NewResponse(&aiserverv1.GetThoughtAnnotationResponse{}), nil
}
return connect.NewResponse(&aiserverv1.GetThoughtAnnotationResponse{
ThoughtAnnotation: &aiserverv1.AiThoughtAnnotation{
RequestId: requestID,
Thought: thought,
},
}), nil
}
func (service *Service) GetTokenUsage(_ context.Context, req *connect.Request[aiserverv1.GetTokenUsageRequest]) (*connect.Response[aiserverv1.GetTokenUsageResponse], error) {
record, ok, err := service.lookupUsageRecord(strings.TrimSpace(req.Msg.GetUsageUuid()))
if err != nil {
return nil, connect.NewError(connect.CodeInternal, err)
}
if !ok {
return connect.NewResponse(&aiserverv1.GetTokenUsageResponse{}), nil
}
return connect.NewResponse(&aiserverv1.GetTokenUsageResponse{
InputTokens: clampInt64ToInt32(record.InputTokens),
OutputTokens: clampInt64ToInt32(record.OutputTokens),
}), nil
}
func (service *Service) GetGlassEarlyPreviewEnrollment(context.Context, *connect.Request[aiserverv1.GetGlassEarlyPreviewEnrollmentRequest]) (*connect.Response[aiserverv1.GetGlassEarlyPreviewEnrollmentResponse], error) {
granted := true
return connect.NewResponse(&aiserverv1.GetGlassEarlyPreviewEnrollmentResponse{
Enabled: true,
EnterpriseGlassSelfEnrollEligible: &granted,
GlassAccessGranted: &granted,
}), nil
}
func (service *Service) CountTokens(_ context.Context, req *connect.Request[aiserverv1.CountTokensRequest]) (*connect.Response[aiserverv1.CountTokensResponse], error) {
total := int64(0)
details := make([]*aiserverv1.ContextItemTokenDetail, 0, len(req.Msg.GetContextItems()))
for _, item := range req.Msg.GetContextItems() {
count := estimateContextItemTokens(item)
total += count
details = append(details, &aiserverv1.ContextItemTokenDetail{
RelativeWorkspacePath: contextItemRelativeWorkspacePath(item),
Count: clampInt64ToInt32(count),
LineCount: lineCountForContextItem(item),
})
}
return connect.NewResponse(&aiserverv1.CountTokensResponse{
Count: clampInt64ToInt32(total),
TokenDetails: details,
}), nil
}
func (service *Service) lookupUsageRecord(usageUUID string) (usageLookupRecord, bool, error) {
if service == nil {
return usageLookupRecord{}, false, nil
}
if service.usageStore == nil {
return usageLookupRecord{}, false, nil
}
item, ok, err := service.usageStore.LookupEvent(strings.TrimSpace(usageUUID))
if err != nil || !ok {
return usageLookupRecord{}, ok, err
}
return usageLookupRecord{
InputTokens: item.InputTokens + item.CacheReadTokens + item.CacheWriteTokens,
OutputTokens: item.OutputTokens,
CreatedAt: item.At,
}, true, nil
}
func (service *Service) lookupThoughtAnnotation(requestID string) (string, bool, error) {
if service == nil || service.store == nil {
return "", false, nil
}
needle := strings.TrimSpace(requestID)
if needle == "" {
return "", false, nil
}
foundRequest := false
foundThought := false
latestThought := ""
latestThoughtAt := time.Time{}
conversationIDs, err := service.store.ListConversationIDs()
if err != nil {
return "", false, err
}
for _, conversationID := range conversationIDs {
conversation, err := service.store.LoadConversation(conversationID)
if err != nil {
return "", false, err
}
if conversation == nil {
continue
}
for _, entry := range conversation.Entries {
if strings.TrimSpace(entry.RequestID) != needle {
continue
}
foundRequest = true
if strings.TrimSpace(entry.Kind) != "metadata" {
continue
}
var payload metadataPayload
if err := json.Unmarshal(entry.Payload, &payload); err != nil {
continue
}
if strings.TrimSpace(payload.Type) != "thought_annotation" {
continue
}
if strings.TrimSpace(readStringValue(payload.Value["kind"])) != "summary_completed" {
continue
}
thought := strings.TrimSpace(readStringValue(payload.Value["thought"]))
if thought == "" {
continue
}
if !foundThought || entry.CreatedAt.After(latestThoughtAt) {
foundThought = true
latestThought = thought
latestThoughtAt = entry.CreatedAt
}
}
}
if foundThought {
return latestThought, true, nil
}
if foundRequest {
return defaultSummaryCompletedThought, true, nil
}
return "", false, nil
}
func lineCountForContextItem(item *aiserverv1.ContextItem) int32 {
if item == nil {
return 0
}
if chunk := item.GetFileChunk(); chunk != nil {
return countTextLines(chunk.GetChunkContents())
}
if outline := item.GetOutlineChunk(); outline != nil {
return lineCountForRange(outline.GetFullRange())
}
if selection := item.GetCmdKSelection(); selection != nil {
return int32(len(selection.GetLines()))
}
if sparse := item.GetSparseFileChunk(); sparse != nil {
return int32(len(sparse.GetLines()))
}
return 0
}
func contextItemRelativeWorkspacePath(item *aiserverv1.ContextItem) string {
if item == nil {
return ""
}
if chunk := item.GetFileChunk(); chunk != nil {
return strings.TrimSpace(chunk.GetRelativeWorkspacePath())
}
if outline := item.GetOutlineChunk(); outline != nil {
return strings.TrimSpace(outline.GetRelativeWorkspacePath())
}
if sparse := item.GetSparseFileChunk(); sparse != nil {
return strings.TrimSpace(sparse.GetRelativeWorkspacePath())
}
return ""
}
func lineCountForRange(lineRange *aiserverv1.LineRange) int32 {
if lineRange == nil {
return 0
}
start := lineRange.GetStartLineNumber()
end := lineRange.GetEndLineNumberInclusive()
if start <= 0 || end < start {
return 0
}
return end - start + 1
}
func countTextLines(text string) int32 {
if strings.TrimSpace(text) == "" {
return 0
}
return int32(strings.Count(text, "\n") + 1)
}
func readInt64Value(value any) int64 {
switch typed := value.(type) {
case float64:
return int64(typed)
case float32:
return int64(typed)
case int:
return int64(typed)
case int32:
return int64(typed)
case int64:
return typed
case uint32:
return int64(typed)
case uint64:
if typed > uint64(^uint32(0)>>1) {
return int64(^uint32(0) >> 1)
}
return int64(typed)
case json.Number:
value, err := typed.Int64()
if err == nil {
return value
}
}
return 0
}
func readStringValue(value any) string {
switch typed := value.(type) {
case string:
return strings.TrimSpace(typed)
default:
return ""
}
}
func readBoolValue(value any) bool {
switch typed := value.(type) {
case bool:
return typed
case string:
return strings.EqualFold(strings.TrimSpace(typed), "true")
default:
return false
}
}
func readTimeValue(value any) time.Time {
return parseRFC3339Time(readStringValue(value))
}
func firstNonZeroTime(values ...time.Time) time.Time {
for _, value := range values {
if !value.IsZero() {
return value.UTC()
}
}
return time.Time{}
}
+149
View File
@@ -0,0 +1,149 @@
package forwarder
import (
"context"
"strings"
"sync"
"time"
)
const appendSequenceRetention = 10 * time.Minute
type appendSequenceTracker struct {
mu sync.Mutex
states map[string]*appendSequenceState
}
type appendSequenceState struct {
mu sync.Mutex
next int64
processing bool
ready chan struct{}
updatedAt time.Time
}
type appendSequenceTicket struct {
state *appendSequenceState
seq int64
}
func newAppendSequenceTracker() *appendSequenceTracker {
return &appendSequenceTracker{
states: make(map[string]*appendSequenceState),
}
}
func (tracker *appendSequenceTracker) Acquire(ctx context.Context, requestID string, appendSeq int64) (appendSequenceTicket, bool, error) {
if tracker == nil || strings.TrimSpace(requestID) == "" || appendSeq <= 0 {
return appendSequenceTicket{}, false, nil
}
state := tracker.state(strings.TrimSpace(requestID))
stale, err := state.acquire(ctx, appendSeq)
if err != nil || stale {
return appendSequenceTicket{}, stale, err
}
return appendSequenceTicket{
state: state,
seq: appendSeq,
}, false, nil
}
func (tracker *appendSequenceTracker) state(requestID string) *appendSequenceState {
now := time.Now().UTC()
cutoff := now.Add(-appendSequenceRetention)
tracker.mu.Lock()
defer tracker.mu.Unlock()
for key, state := range tracker.states {
if state == nil || state.expired(cutoff) {
delete(tracker.states, key)
}
}
if state, ok := tracker.states[requestID]; ok && state != nil {
state.touch(now)
return state
}
state := &appendSequenceState{
next: 1,
ready: make(chan struct{}),
updatedAt: now,
}
tracker.states[requestID] = state
return state
}
func (state *appendSequenceState) acquire(ctx context.Context, appendSeq int64) (bool, error) {
for {
state.mu.Lock()
now := time.Now().UTC()
if state.next <= 0 {
state.next = 1
}
if state.ready == nil {
state.ready = make(chan struct{})
}
state.updatedAt = now
switch {
case appendSeq < state.next:
state.mu.Unlock()
return true, nil
case appendSeq == state.next && !state.processing:
state.processing = true
state.mu.Unlock()
return false, nil
default:
ready := state.ready
state.mu.Unlock()
select {
case <-ctx.Done():
return false, ctx.Err()
case <-ready:
}
}
}
}
func (state *appendSequenceState) Release(seq int64) {
if state == nil || seq <= 0 {
return
}
state.mu.Lock()
defer state.mu.Unlock()
if state.processing && state.next == seq {
state.processing = false
state.next++
close(state.ready)
state.ready = make(chan struct{})
}
state.updatedAt = time.Now().UTC()
}
func (state *appendSequenceState) expired(cutoff time.Time) bool {
if state == nil {
return true
}
state.mu.Lock()
defer state.mu.Unlock()
if state.processing {
return false
}
return !state.updatedAt.IsZero() && state.updatedAt.Before(cutoff)
}
func (state *appendSequenceState) touch(now time.Time) {
if state == nil {
return
}
state.mu.Lock()
state.updatedAt = now
state.mu.Unlock()
}
func (ticket appendSequenceTicket) Release() {
if ticket.state == nil || ticket.seq <= 0 {
return
}
ticket.state.Release(ticket.seq)
}
+310
View File
@@ -0,0 +1,310 @@
package forwarder
import (
"context"
"fmt"
"os"
"path/filepath"
"runtime"
"strings"
"sync"
"time"
)
// artifactRecorder 只缓存当前 provider call 的请求摘要,用于本轮内 usage/model 信息补齐。
// 对话恢复事实只存在 state.json 与 context.json。
type artifactRecorder struct {
store *ConversationFileStore
broker *StreamBroker
debug *debugRecorder
mu sync.Mutex
sessions map[string]artifactSession
}
type artifactSession struct {
conversationID string
turnSeq int64
requestPayload map[string]any
summaryPayload map[string]any
}
type requestArtifactPrefix struct {
Provider string
OpenAIEndpoint string
Model string
PromptTokensTotal int64
ReplayMessageCount int
CanonicalBodyHash string
FrontierHash string
FrontierPath string
BreakpointCount int
ExpectedCacheRead bool
PreviousFrontierMatched bool
}
func newArtifactRecorder(store *ConversationFileStore, broker *StreamBroker, debug *debugRecorder) *artifactRecorder {
return &artifactRecorder{
store: store,
broker: broker,
debug: debug,
sessions: make(map[string]artifactSession),
}
}
func (recorder *artifactRecorder) RecordLLMRequest(requestID string, _ string, modelCallID string, payload map[string]any) (string, error) {
session, err := recorder.ensureSession(requestID, modelCallID)
if err != nil {
return "", err
}
session.requestPayload = cloneStringAnyMap(payload)
recorder.mu.Lock()
recorder.sessions[artifactSessionKey(requestID, modelCallID)] = session
recorder.mu.Unlock()
if prefix, ok, err := decodeRequestPrefixPayload(session.requestPayload); err == nil && ok && prefix != nil {
recorder.persistLatestRequestPrefix(session.conversationID, requestID, modelCallID, prefix)
}
recorder.debug.LogProviderArtifact(context.Background(), requestID, session.conversationID, modelCallID, "llm_request", session.requestPayload)
return "", nil
}
func (recorder *artifactRecorder) AppendLLMResponseChunk(requestID string, runID string, modelCallID string, chunk string) (string, error) {
session, err := recorder.ensureSession(requestID, modelCallID)
if err != nil {
return "", err
}
recorder.debug.LogProviderArtifact(context.Background(), requestID, session.conversationID, modelCallID, "llm_response_chunk", map[string]any{
"run_id": strings.TrimSpace(runID),
"raw_chunk": chunk,
"byte_len": len([]byte(chunk)),
})
return "", nil
}
func (recorder *artifactRecorder) RecordLLMSummary(requestID string, _ string, modelCallID string, payload map[string]any) (string, error) {
session, err := recorder.ensureSession(requestID, modelCallID)
if err != nil {
return "", err
}
session.summaryPayload = cloneStringAnyMap(payload)
recorder.mu.Lock()
recorder.sessions[artifactSessionKey(requestID, modelCallID)] = session
recorder.mu.Unlock()
if prefix, ok, err := decodeRequestPrefixPayload(session.requestPayload); err == nil && ok && prefix != nil {
if tokens := readInt64Value(session.summaryPayload["prompt_tokens_total"]); tokens > 0 {
prefix.PromptTokensTotal = tokens
}
recorder.persistLatestRequestPrefix(session.conversationID, requestID, modelCallID, prefix)
}
recorder.debug.LogProviderArtifact(context.Background(), requestID, session.conversationID, modelCallID, "llm_summary", session.summaryPayload)
return "", nil
}
func (recorder *artifactRecorder) persistLatestRequestPrefix(conversationID string, requestID string, modelCallID string, prefix *requestArtifactPrefix) {
if recorder == nil || recorder.store == nil || prefix == nil || strings.TrimSpace(conversationID) == "" || strings.TrimSpace(conversationID) == "unknown" {
return
}
_, _ = recorder.store.UpdateConversationMeta(conversationID, func(item *ConversationFile) error {
if item == nil {
return nil
}
item.LatestRequestPrefix = &ConversationRequestPrefix{
RequestID: strings.TrimSpace(requestID),
ModelCallID: firstNonEmpty(strings.TrimSpace(modelCallID), strings.TrimSpace(requestID)),
Provider: strings.TrimSpace(prefix.Provider),
OpenAIEndpoint: strings.TrimSpace(prefix.OpenAIEndpoint),
Model: strings.TrimSpace(prefix.Model),
PromptTokensTotal: prefix.PromptTokensTotal,
ReplayMessageCount: prefix.ReplayMessageCount,
CanonicalBodyHash: strings.TrimSpace(prefix.CanonicalBodyHash),
FrontierHash: strings.TrimSpace(prefix.FrontierHash),
FrontierPath: strings.TrimSpace(prefix.FrontierPath),
BreakpointCount: prefix.BreakpointCount,
ExpectedCacheRead: prefix.ExpectedCacheRead,
PreviousFrontierMatched: prefix.PreviousFrontierMatched,
UpdatedAt: time.Now().UTC(),
}
return nil
})
}
func (recorder *artifactRecorder) ClearActiveArtifacts(requestID string, modelCallID string) {
if recorder == nil {
return
}
recorder.mu.Lock()
delete(recorder.sessions, artifactSessionKey(requestID, modelCallID))
recorder.mu.Unlock()
}
func (recorder *artifactRecorder) ensureSession(requestID string, modelCallID string) (artifactSession, error) {
if recorder == nil {
return artifactSession{}, fmt.Errorf("artifact recorder is nil")
}
key := artifactSessionKey(requestID, modelCallID)
recorder.mu.Lock()
defer recorder.mu.Unlock()
if session, ok := recorder.sessions[key]; ok {
return session, nil
}
conversationID, turnSeq := recorder.resolveConversationContext(requestID)
session := artifactSession{
conversationID: conversationID,
turnSeq: turnSeq,
}
recorder.sessions[key] = session
return session, nil
}
func artifactSessionKey(requestID string, modelCallID string) string {
return strings.TrimSpace(requestID) + "::" + strings.TrimSpace(modelCallID)
}
func (recorder *artifactRecorder) resolveConversationContext(requestID string) (string, int64) {
if recorder == nil || recorder.broker == nil {
return "unknown", 0
}
stream, ok := recorder.broker.Get(requestID)
if !ok || stream == nil {
return "unknown", 0
}
stream.mu.Lock()
defer stream.mu.Unlock()
return sanitizeArtifactName(firstNonEmpty(stream.ConversationID, "unknown")), stream.TurnSeq
}
func decodeRequestPrefixPayload(payload map[string]any) (*requestArtifactPrefix, bool, error) {
if len(payload) == 0 {
return nil, false, nil
}
provider := strings.TrimSpace(readStringValue(payload["provider"]))
model := strings.TrimSpace(readStringValue(payload["model"]))
openAIEndpoint := strings.TrimSpace(readStringValue(payload["openai_endpoint"]))
if provider == "" && model == "" && openAIEndpoint == "" {
return nil, false, nil
}
replayMessageCount := replayMessageCountFromRequestPayload(payload)
frontier := cacheFrontierFromRequestPayload(payload)
return &requestArtifactPrefix{
Provider: provider,
OpenAIEndpoint: openAIEndpoint,
Model: model,
ReplayMessageCount: replayMessageCount,
CanonicalBodyHash: frontier.CanonicalBodyHash,
FrontierHash: frontier.FrontierHash,
FrontierPath: frontier.FrontierPath,
BreakpointCount: frontier.BreakpointCount,
ExpectedCacheRead: frontier.ExpectedCacheRead,
PreviousFrontierMatched: frontier.PreviousFrontierMatched,
}, true, nil
}
type requestCacheFrontierPayload struct {
CanonicalBodyHash string
FrontierHash string
FrontierPath string
BreakpointCount int
ExpectedCacheRead bool
PreviousFrontierMatched bool
}
func cacheFrontierFromRequestPayload(payload map[string]any) requestCacheFrontierPayload {
var frontier requestCacheFrontierPayload
knobs, ok := payload["request_knobs"].(map[string]any)
if !ok {
return frontier
}
rawFrontier, ok := knobs["cache_frontier"].(map[string]any)
if !ok {
return frontier
}
frontier.CanonicalBodyHash = strings.TrimSpace(readStringValue(rawFrontier["canonical_body_hash"]))
frontier.FrontierHash = strings.TrimSpace(readStringValue(rawFrontier["frontier_hash"]))
frontier.FrontierPath = strings.TrimSpace(readStringValue(rawFrontier["frontier_path"]))
frontier.BreakpointCount = int(readInt64Value(rawFrontier["breakpoint_count"]))
frontier.ExpectedCacheRead = readBoolValue(rawFrontier["expected_cache_read"])
frontier.PreviousFrontierMatched = readBoolValue(rawFrontier["previous_frontier_matched"])
return frontier
}
func replayMessageCountFromRequestPayload(payload map[string]any) int {
switch messages := payload["messages_summary"].(type) {
case []map[string]any:
return replayMessageCountFromSummaryObjects(messages)
case []any:
items := make([]map[string]any, 0, len(messages))
for _, item := range messages {
message, ok := item.(map[string]any)
if ok {
items = append(items, message)
}
}
return replayMessageCountFromSummaryObjects(items)
default:
return 0
}
}
func replayMessageCountFromSummaryObjects(messages []map[string]any) int {
count := 0
for _, message := range messages {
if strings.TrimSpace(readStringValue(message["role"])) == "system" {
continue
}
count++
}
return count
}
func sanitizeArtifactName(value string) string {
replacer := strings.NewReplacer("/", "_", "\\", "_", ":", "_", " ", "_")
normalized := replacer.Replace(strings.TrimSpace(value))
if normalized == "" {
return "unknown"
}
return normalized
}
func artifactConversationDir(historyRoot string, conversationID string) string {
return filepath.Join(strings.TrimSpace(historyRoot), sanitizeArtifactName(conversationID))
}
func openUniqueArtifactTempFile(path string) (*os.File, string, error) {
dir := filepath.Dir(path)
base := filepath.Base(path)
for attempt := 0; attempt < 32; attempt++ {
tempPath := filepath.Join(dir, fmt.Sprintf("%s.tmp-%d-%d", base, os.Getpid(), time.Now().UnixNano()))
file, err := os.OpenFile(tempPath, os.O_CREATE|os.O_EXCL|os.O_WRONLY, 0o644)
if err == nil {
return file, tempPath, nil
}
if !os.IsExist(err) {
return nil, "", err
}
}
return nil, "", fmt.Errorf("exhausted temp file attempts")
}
func renameArtifactTempFile(tempPath string, path string) error {
attempts := 1
delay := 10 * time.Millisecond
if runtime.GOOS == "windows" {
attempts = 12
}
var lastErr error
for attempt := 0; attempt < attempts; attempt++ {
if err := os.Rename(tempPath, path); err == nil {
return nil
} else {
lastErr = err
}
if runtime.GOOS != "windows" || os.IsNotExist(lastErr) || attempt == attempts-1 {
break
}
time.Sleep(delay)
if delay < 200*time.Millisecond {
delay *= 2
}
}
return lastErr
}
@@ -0,0 +1,954 @@
package forwarder
import (
"encoding/json"
"fmt"
"os"
"path/filepath"
"regexp"
"strconv"
"strings"
"time"
"cursor/gen/agentv1"
execbridge "cursor/internal/backend/agent/bridge/exec"
runtimecore "cursor/internal/backend/agent/core"
)
const awaitShellOutputLimit = 16 * 1024
const (
backgroundShellStatusBackgrounded = "backgrounded"
backgroundShellStatusRunning = "running"
backgroundShellStatusCompleted = "completed"
backgroundShellStatusRejected = "rejected"
backgroundShellStatusPermissionDenied = "permission_denied"
backgroundShellStatusTransportClosed = "transport_closed"
backgroundShellStatusUnknown = "unknown"
backgroundShellActionSourceClient = "client"
backgroundShellActionSourceLocalBackgrounded = "local_shell_backgrounded"
)
type awaitShellArgs struct {
ShellID string `json:"shell_id,omitempty"`
TaskID string `json:"task_id,omitempty"`
BlockUntilMS *int64 `json:"block_until_ms,omitempty"`
Pattern string `json:"pattern,omitempty"`
}
type awaitShellResult struct {
ShellID string `json:"shell_id,omitempty"`
Status string `json:"status"`
Matched bool `json:"matched"`
TimedOut bool `json:"timed_out"`
ExitCode *int64 `json:"exit_code,omitempty"`
Stdout string `json:"stdout,omitempty"`
Stderr string `json:"stderr,omitempty"`
StdoutOffset int `json:"stdout_offset,omitempty"`
StderrOffset int `json:"stderr_offset,omitempty"`
RuntimeMS uint64 `json:"runtime_ms,omitempty"`
OutputLength uint64 `json:"output_length,omitempty"`
RegexRequested bool `json:"regex_requested,omitempty"`
RegexMatch *string `json:"regex_match,omitempty"`
Message string `json:"message,omitempty"`
}
func (service *Service) handleAwaitShellToolInvocation(stream *ActiveStream, invocation runtimecore.ToolInvocation) error {
args, err := decodeAwaitShellArgs(invocation.ArgsJSON)
if err != nil {
result := awaitShellResult{Status: "error", Message: err.Error()}
payload, encodeErr := json.Marshal(result)
if encodeErr != nil {
return encodeErr
}
return service.completeImmediateToolResult(stream, invocation, string(payload), buildAwaitShellToolCall(buildAwaitArgsFromAwaitShellArgs(args), buildAwaitShellProtoResult(result)))
}
result := service.awaitShellSnapshot(stream, args)
payload, err := json.Marshal(result)
if err != nil {
return err
}
return service.completeImmediateToolResult(stream, invocation, string(payload), buildAwaitShellToolCall(buildAwaitArgsFromAwaitShellArgs(args), buildAwaitShellProtoResult(result)))
}
func decodeAwaitShellArgs(raw []byte) (awaitShellArgs, error) {
argsMap, err := runtimecore.DecodeArgsMap(raw)
if err != nil {
return awaitShellArgs{}, fmt.Errorf("decode AwaitShell args failed: %w", err)
}
result := awaitShellArgs{
ShellID: strings.TrimSpace(runtimecore.ReadStringArg(argsMap, "shell_id")),
TaskID: strings.TrimSpace(runtimecore.ReadStringArg(argsMap, "task_id")),
Pattern: strings.TrimSpace(runtimecore.ReadStringArg(argsMap, "pattern")),
}
if result.ShellID == "" {
result.ShellID = result.TaskID
}
if value, found, err := runtimecore.ReadInt64Arg(argsMap, "block_until_ms"); err != nil {
return result, err
} else if found {
if value < 0 {
value = 0
}
result.BlockUntilMS = &value
}
if result.BlockUntilMS != nil && *result.BlockUntilMS == 0 && strings.TrimSpace(result.ShellID) == "" {
return result, fmt.Errorf("AwaitShell shell_id is required when block_until_ms is 0")
}
return result, nil
}
func (service *Service) awaitShellSnapshot(stream *ActiveStream, args awaitShellArgs) awaitShellResult {
blockUntilMS := int64(30000)
if args.BlockUntilMS != nil {
blockUntilMS = *args.BlockUntilMS
}
shellID := strings.TrimSpace(args.ShellID)
if shellID == "" {
return awaitShellResult{
Status: "waited",
TimedOut: false,
Message: fmt.Sprintf("waited %dms", blockUntilMS),
}
}
service.refreshBackgroundShellFromTerminalFile(stream, shellID)
stream.mu.Lock()
state, ok := stream.BackgroundShells[shellID]
if !ok || state == nil {
stream.mu.Unlock()
return awaitShellResult{
ShellID: shellID,
Status: backgroundShellStatusUnknown,
TimedOut: false,
Message: "unknown or expired shell_id",
}
}
stdoutStart := clampOffset(state.AwaitStdoutOffset, len(state.StdoutBuffer))
stderrStart := clampOffset(state.AwaitStderrOffset, len(state.StderrBuffer))
stdout := state.StdoutBuffer[stdoutStart:]
stderr := state.StderrBuffer[stderrStart:]
stdoutEnd := len(state.StdoutBuffer)
stderrEnd := len(state.StderrBuffer)
state.AwaitStdoutOffset = stdoutEnd
state.AwaitStderrOffset = stderrEnd
status := strings.TrimSpace(state.Status)
if status == "" {
status = backgroundShellStatusUnknown
}
var exitCode *int64
if state.ExitCode != nil {
value := int64(*state.ExitCode)
exitCode = &value
}
createdAt := state.CreatedAt
completedAt := state.CompletedAt
combinedOutput := state.StdoutBuffer + "\n" + state.StderrBuffer
stream.UpdatedAt = time.Now().UTC()
stream.mu.Unlock()
matched, matchText, patternErr := awaitShellPatternMatched(args.Pattern, combinedOutput)
message := ""
if patternErr != nil {
message = patternErr.Error()
}
timedOut := blockUntilMS > 0 && !matched && !isBackgroundShellTerminalStatus(status)
stdout = truncateAwaitShellOutput(stdout)
stderr = truncateAwaitShellOutput(stderr)
return awaitShellResult{
ShellID: shellID,
Status: status,
Matched: matched,
TimedOut: timedOut,
ExitCode: exitCode,
Stdout: stdout,
Stderr: stderr,
StdoutOffset: stdoutEnd,
StderrOffset: stderrEnd,
RuntimeMS: backgroundShellRuntimeMS(createdAt, completedAt),
OutputLength: uint64(len(combinedOutput)),
RegexRequested: strings.TrimSpace(args.Pattern) != "",
RegexMatch: matchText,
Message: message,
}
}
func buildAwaitArgsFromAwaitShellArgs(args awaitShellArgs) *agentv1.AwaitArgs {
awaitArgs := &agentv1.AwaitArgs{TaskId: strings.TrimSpace(firstNonEmpty(args.ShellID, args.TaskID))}
if args.BlockUntilMS != nil {
value := uint32(0)
if *args.BlockUntilMS > 0 {
value = uint32(*args.BlockUntilMS)
}
awaitArgs.BlockUntilMs = &value
}
if pattern := strings.TrimSpace(args.Pattern); pattern != "" {
awaitArgs.Regex = &pattern
}
return awaitArgs
}
func buildAwaitShellToolCall(args *agentv1.AwaitArgs, result *agentv1.AwaitResult) *agentv1.ToolCall {
return &agentv1.ToolCall{
Tool: &agentv1.ToolCall_AwaitToolCall{
AwaitToolCall: &agentv1.AwaitToolCall{
Args: args,
Result: result,
},
},
}
}
func buildAwaitShellProtoResult(result awaitShellResult) *agentv1.AwaitResult {
switch strings.TrimSpace(result.Status) {
case "error", backgroundShellStatusUnknown:
message := firstNonEmpty(strings.TrimSpace(result.Message), "unknown or expired shell_id")
return &agentv1.AwaitResult{Result: &agentv1.AwaitResult_Error{Error: &agentv1.AwaitError{Error: message}}}
}
if isBackgroundShellTerminalStatus(result.Status) || result.Matched {
task := &agentv1.AwaitTaskComplete{
TaskId: strings.TrimSpace(result.ShellID),
RuntimeMs: result.RuntimeMS,
OutputLength: result.OutputLength,
RegexRequested: result.RegexRequested,
RegexMatch: result.RegexMatch,
}
if result.ExitCode != nil {
exitCode := int32(*result.ExitCode)
task.ExitCode = &exitCode
}
return &agentv1.AwaitResult{Result: &agentv1.AwaitResult_Complete{Complete: task}}
}
task := &agentv1.AwaitTaskStillRunning{
TaskId: strings.TrimSpace(result.ShellID),
RuntimeMs: result.RuntimeMS,
OutputLength: result.OutputLength,
RegexRequested: result.RegexRequested,
RegexMatch: result.RegexMatch,
}
return &agentv1.AwaitResult{Result: &agentv1.AwaitResult_StillRunning{StillRunning: task}}
}
func backgroundShellRuntimeMS(createdAt time.Time, completedAt time.Time) uint64 {
if createdAt.IsZero() {
return 0
}
end := completedAt
if end.IsZero() {
end = time.Now().UTC()
}
if end.Before(createdAt) {
return 0
}
return uint64(end.Sub(createdAt).Milliseconds())
}
func awaitShellPatternMatched(pattern string, output string) (bool, *string, error) {
trimmed := strings.TrimSpace(pattern)
if trimmed == "" {
return false, nil, nil
}
expr, err := regexp.Compile("(?m)" + trimmed)
if err != nil {
return false, nil, fmt.Errorf("invalid AwaitShell pattern: %w", err)
}
match := expr.FindString(output)
if match == "" {
return false, nil, nil
}
return true, &match, nil
}
func truncateAwaitShellOutput(value string) string {
if len(value) <= awaitShellOutputLimit {
return value
}
return value[len(value)-awaitShellOutputLimit:]
}
func clampOffset(offset int, length int) int {
if offset < 0 {
return 0
}
if offset > length {
return length
}
return offset
}
type terminalShellFileSnapshot struct {
PID *uint32
Command string
CWD string
StartedAt time.Time
Output string
ExitCode *int32
EndedAt time.Time
}
func (service *Service) refreshBackgroundShellFromTerminalFile(stream *ActiveStream, shellID string) {
if stream == nil {
return
}
stream.mu.Lock()
terminalsFolder := strings.TrimSpace(stream.TerminalsFolder)
stream.mu.Unlock()
terminalSnapshot, ok := readTerminalShellFileSnapshot(terminalsFolder, shellID)
if !ok {
return
}
now := time.Now().UTC()
stream.mu.Lock()
defer stream.mu.Unlock()
state := ensureBackgroundShellStateLocked(stream, shellID, runtimecore.PendingExec{}, firstNonZeroTime(terminalSnapshot.StartedAt, now))
if state == nil {
return
}
if state.Command == "" {
state.Command = terminalSnapshot.Command
}
if state.WorkingDirectory == "" {
state.WorkingDirectory = terminalSnapshot.CWD
}
if state.PID == nil && terminalSnapshot.PID != nil {
pid := *terminalSnapshot.PID
state.PID = &pid
}
if state.CreatedAt.IsZero() {
state.CreatedAt = firstNonZeroTime(terminalSnapshot.StartedAt, now)
}
if state.StdoutBuffer == "" && state.StderrBuffer == "" && terminalSnapshot.Output != "" {
state.StdoutBuffer = terminalSnapshot.Output
if !isBackgroundShellTerminalStatus(state.Status) {
state.Status = backgroundShellStatusRunning
}
state.LastActivityAt = now
}
if !isBackgroundShellTerminalStatus(state.Status) && terminalSnapshot.ExitCode != nil {
exitCode := *terminalSnapshot.ExitCode
state.ExitCode = &exitCode
state.Status = backgroundShellStatusCompleted
state.CompletedAt = firstNonZeroTime(terminalSnapshot.EndedAt, now)
state.LastActivityAt = state.CompletedAt
}
if strings.TrimSpace(state.Status) == "" {
state.Status = backgroundShellStatusBackgrounded
}
if state.LastActivityAt.IsZero() {
state.LastActivityAt = now
}
stream.UpdatedAt = now
}
func readTerminalShellFileSnapshot(terminalsFolder string, shellID string) (terminalShellFileSnapshot, bool) {
path, ok := terminalShellFilePath(terminalsFolder, shellID)
if !ok {
return terminalShellFileSnapshot{}, false
}
data, err := os.ReadFile(path)
if err != nil || len(data) == 0 {
return terminalShellFileSnapshot{}, false
}
return parseTerminalShellFileSnapshot(string(data)), true
}
func terminalShellFilePath(terminalsFolder string, shellID string) (string, bool) {
folder := filepath.Clean(strings.TrimSpace(terminalsFolder))
if folder == "." || !filepath.IsAbs(folder) {
return "", false
}
id, err := strconv.ParseUint(strings.TrimSpace(shellID), 10, 32)
if err != nil {
return "", false
}
filename := strconv.FormatUint(id, 10) + ".txt"
return filepath.Join(folder, filename), true
}
func parseTerminalShellFileSnapshot(raw string) terminalShellFileSnapshot {
lines := strings.Split(strings.ReplaceAll(raw, "\r\n", "\n"), "\n")
separators := terminalShellSeparatorIndexes(lines)
snapshot := terminalShellFileSnapshot{}
if len(separators) >= 2 {
parseTerminalShellMetadata(lines[separators[0]+1:separators[1]], &snapshot)
}
outputStart := 0
if len(separators) >= 2 {
outputStart = separators[1] + 1
}
outputEnd := len(lines)
for index := len(separators) - 2; index >= 1; index-- {
block := lines[separators[index]+1 : separators[index+1]]
if terminalMetadataBlockHasKey(block, "exit_code") || terminalMetadataBlockHasKey(block, "ended_at") {
parseTerminalShellMetadata(block, &snapshot)
outputEnd = separators[index]
break
}
}
if outputStart < outputEnd {
snapshot.Output = strings.Join(lines[outputStart:outputEnd], "\n")
snapshot.Output = strings.TrimSuffix(snapshot.Output, "\n")
}
return snapshot
}
func terminalShellSeparatorIndexes(lines []string) []int {
indexes := make([]int, 0, 4)
for index, line := range lines {
if strings.TrimSpace(line) == "---" {
indexes = append(indexes, index)
}
}
return indexes
}
func parseTerminalShellMetadata(lines []string, snapshot *terminalShellFileSnapshot) {
if snapshot == nil {
return
}
for _, line := range lines {
key, value, ok := strings.Cut(line, ":")
if !ok {
continue
}
key = strings.TrimSpace(key)
value = parseTerminalShellMetadataValue(value)
switch key {
case "pid":
if parsed, err := strconv.ParseUint(value, 10, 32); err == nil {
pid := uint32(parsed)
snapshot.PID = &pid
}
case "cwd":
snapshot.CWD = value
case "command":
snapshot.Command = value
case "started_at":
snapshot.StartedAt = parseTerminalShellTime(value)
case "exit_code":
if parsed, err := strconv.ParseInt(value, 10, 32); err == nil {
exitCode := int32(parsed)
snapshot.ExitCode = &exitCode
}
case "ended_at":
snapshot.EndedAt = parseTerminalShellTime(value)
}
}
}
func parseTerminalShellMetadataValue(value string) string {
trimmed := strings.TrimSpace(value)
if unquoted, err := strconv.Unquote(trimmed); err == nil {
return unquoted
}
return trimmed
}
func parseTerminalShellTime(value string) time.Time {
parsed, err := time.Parse(time.RFC3339Nano, strings.TrimSpace(value))
if err != nil {
return time.Time{}
}
return parsed.UTC()
}
func terminalMetadataBlockHasKey(lines []string, key string) bool {
for _, line := range lines {
current, _, ok := strings.Cut(line, ":")
if ok && strings.TrimSpace(current) == key {
return true
}
}
return false
}
func isBackgroundShellTerminalStatus(status string) bool {
switch strings.TrimSpace(status) {
case backgroundShellStatusCompleted, backgroundShellStatusRejected, backgroundShellStatusPermissionDenied, backgroundShellStatusTransportClosed:
return true
default:
return false
}
}
func (service *Service) observeBackgroundShellExecClientMessage(stream *ActiveStream, pending runtimecore.PendingExec, message *agentv1.ExecClientMessage) {
if stream == nil || message == nil {
return
}
now := time.Now().UTC()
stream.mu.Lock()
defer stream.mu.Unlock()
service.observeBackgroundShellSpawnResultLocked(stream, pending, message.GetBackgroundShellSpawnResult(), now)
service.observeForceBackgroundShellResultLocked(stream, pending, message.GetForceBackgroundShellResult(), now)
}
func (service *Service) observeMissingBackgroundShellExecClientMessage(stream *ActiveStream, message *agentv1.ExecClientMessage) bool {
if stream == nil || message == nil {
return false
}
if message.GetBackgroundShellSpawnResult() == nil && message.GetForceBackgroundShellResult() == nil {
return false
}
now := time.Now().UTC()
stream.mu.Lock()
defer stream.mu.Unlock()
pending := runtimecore.PendingExec{
MessageID: message.GetId(),
ExecID: strings.TrimSpace(message.GetExecId()),
ExecKind: "shell",
}
return service.observeBackgroundShellSpawnResultLocked(stream, pending, message.GetBackgroundShellSpawnResult(), now) || service.observeForceBackgroundShellResultLocked(stream, pending, message.GetForceBackgroundShellResult(), now)
}
func (service *Service) observeBackgroundShellSpawnResultLocked(stream *ActiveStream, pending runtimecore.PendingExec, result *agentv1.BackgroundShellSpawnResult, now time.Time) bool {
if stream == nil || result == nil {
return false
}
if stream.BackgroundShells == nil {
stream.BackgroundShells = make(map[string]*BackgroundShellState)
}
if stream.BackgroundShellsByMessageID == nil {
stream.BackgroundShellsByMessageID = make(map[uint32]string)
}
if stream.BackgroundShellsByExecID == nil {
stream.BackgroundShellsByExecID = make(map[string]string)
}
success := result.GetSuccess()
if success != nil {
shellID := strconv.FormatUint(uint64(success.GetShellId()), 10)
state := stream.BackgroundShells[shellID]
if state == nil {
state = &BackgroundShellState{ShellID: shellID, CreatedAt: now}
stream.BackgroundShells[shellID] = state
}
pid := success.Pid
state.Command = firstNonEmpty(strings.TrimSpace(success.GetCommand()), state.Command)
state.WorkingDirectory = firstNonEmpty(strings.TrimSpace(success.GetWorkingDirectory()), state.WorkingDirectory)
state.PID = pid
state.Status = backgroundShellStatusBackgrounded
state.OriginalToolCallID = firstNonEmpty(strings.TrimSpace(pending.ToolCallID), state.OriginalToolCallID)
state.OriginalExecID = firstNonEmpty(strings.TrimSpace(pending.ExecID), state.OriginalExecID)
if state.OriginalMessageID == 0 {
state.OriginalMessageID = pending.MessageID
}
if len(state.ArgsJSON) == 0 && len(pending.ArgsJSON) > 0 {
state.ArgsJSON = append([]byte(nil), pending.ArgsJSON...)
}
state.ModelCallID = firstNonEmpty(strings.TrimSpace(pending.ModelCallID), state.ModelCallID)
state.LastActivityAt = now
if pending.MessageID != 0 {
stream.BackgroundShellsByMessageID[pending.MessageID] = shellID
}
if strings.TrimSpace(pending.ExecID) != "" {
stream.BackgroundShellsByExecID[strings.TrimSpace(pending.ExecID)] = shellID
}
stream.UpdatedAt = now
return true
}
shellID := backgroundShellIDForMessageLocked(stream, pending.MessageID, pending.ExecID)
if shellID == "" {
return false
}
state := ensureBackgroundShellStateLocked(stream, shellID, pending, now)
if state == nil {
return false
}
switch {
case result.GetRejected() != nil:
state.Status = backgroundShellStatusRejected
state.StderrBuffer += strings.TrimSpace(result.GetRejected().GetReason())
case result.GetPermissionDenied() != nil:
state.Status = backgroundShellStatusPermissionDenied
state.StderrBuffer += strings.TrimSpace(result.GetPermissionDenied().GetError())
case result.GetError() != nil:
state.Status = backgroundShellStatusTransportClosed
state.StderrBuffer += strings.TrimSpace(result.GetError().GetError())
default:
return false
}
state.LastActivityAt = now
state.CompletedAt = now
stream.UpdatedAt = now
return true
}
func (service *Service) observeForceBackgroundShellResultLocked(stream *ActiveStream, pending runtimecore.PendingExec, result *agentv1.ForceBackgroundShellResult, now time.Time) bool {
if stream == nil || result == nil || result.GetShellResult() == nil {
return false
}
shellResult := result.GetShellResult()
shellIDValue := uint32(0)
if success := shellResult.GetSuccess(); success != nil {
shellIDValue = success.GetShellId()
}
if shellIDValue == 0 {
return false
}
shellID := strconv.FormatUint(uint64(shellIDValue), 10)
state := ensureBackgroundShellStateLocked(stream, shellID, pending, now)
if state == nil {
return false
}
state.Status = backgroundShellStatusBackgrounded
state.LastActivityAt = now
stream.UpdatedAt = now
return true
}
func (service *Service) observeShellExecClientMessage(stream *ActiveStream, pending runtimecore.PendingExec, message *agentv1.ExecClientMessage) {
if stream == nil || message == nil || strings.TrimSpace(pending.ExecKind) != "shell" {
return
}
shellStream := message.GetShellStream()
if shellStream == nil {
return
}
now := time.Now().UTC()
stream.mu.Lock()
defer stream.mu.Unlock()
service.observeShellStreamLocked(stream, pending, shellStream, now)
}
func (service *Service) observeMissingShellExecClientMessage(stream *ActiveStream, message *agentv1.ExecClientMessage) bool {
if stream == nil || message == nil || message.GetShellStream() == nil {
return false
}
now := time.Now().UTC()
stream.mu.Lock()
defer stream.mu.Unlock()
shellID := backgroundShellIDForMessageLocked(stream, message.GetId(), message.GetExecId())
if shellID == "" {
return false
}
pending := runtimecore.PendingExec{
MessageID: message.GetId(),
ExecID: strings.TrimSpace(message.GetExecId()),
ExecKind: "shell",
}
if state := stream.BackgroundShells[shellID]; state != nil {
pending.ToolCallID = state.OriginalToolCallID
pending.ArgsJSON = append([]byte(nil), state.ArgsJSON...)
pending.ModelCallID = state.ModelCallID
}
service.observeShellStreamLocked(stream, pending, message.GetShellStream(), now)
return true
}
func (service *Service) observeShellStreamLocked(stream *ActiveStream, pending runtimecore.PendingExec, shellStream *agentv1.ShellStream, now time.Time) {
if stream.BackgroundShells == nil {
stream.BackgroundShells = make(map[string]*BackgroundShellState)
}
if stream.BackgroundShellsByMessageID == nil {
stream.BackgroundShellsByMessageID = make(map[uint32]string)
}
if stream.BackgroundShellsByExecID == nil {
stream.BackgroundShellsByExecID = make(map[string]string)
}
shellID := backgroundShellIDForMessageLocked(stream, pending.MessageID, pending.ExecID)
switch event := shellStream.GetEvent().(type) {
case *agentv1.ShellStream_Backgrounded:
shellID = strconv.FormatUint(uint64(event.Backgrounded.GetShellId()), 10)
state := stream.BackgroundShells[shellID]
if state == nil {
state = &BackgroundShellState{ShellID: shellID, CreatedAt: now}
stream.BackgroundShells[shellID] = state
}
state.Command = firstNonEmpty(strings.TrimSpace(event.Backgrounded.GetCommand()), state.Command)
state.WorkingDirectory = firstNonEmpty(strings.TrimSpace(event.Backgrounded.GetWorkingDirectory()), state.WorkingDirectory)
state.PID = event.Backgrounded.Pid
state.Status = backgroundShellStatusBackgrounded
state.OriginalToolCallID = strings.TrimSpace(pending.ToolCallID)
state.OriginalExecID = strings.TrimSpace(pending.ExecID)
state.OriginalMessageID = pending.MessageID
state.ArgsJSON = append([]byte(nil), pending.ArgsJSON...)
state.ModelCallID = strings.TrimSpace(pending.ModelCallID)
state.LastActivityAt = now
if pending.MessageID != 0 {
stream.BackgroundShellsByMessageID[pending.MessageID] = shellID
}
if strings.TrimSpace(pending.ExecID) != "" {
stream.BackgroundShellsByExecID[strings.TrimSpace(pending.ExecID)] = shellID
}
case *agentv1.ShellStream_Stdout:
state := ensureBackgroundShellStateLocked(stream, shellID, pending, now)
if state == nil {
return
}
state.StdoutBuffer += execbridge.DecodeShellStdout(event.Stdout)
state.Status = backgroundShellStatusRunning
state.LastActivityAt = now
case *agentv1.ShellStream_Stderr:
state := ensureBackgroundShellStateLocked(stream, shellID, pending, now)
if state == nil {
return
}
state.StderrBuffer += event.Stderr.GetData()
state.Status = backgroundShellStatusRunning
state.LastActivityAt = now
case *agentv1.ShellStream_Exit:
state := ensureBackgroundShellStateLocked(stream, shellID, pending, now)
if state == nil {
return
}
exitCode := int32(event.Exit.GetCode())
state.ExitCode = &exitCode
state.WorkingDirectory = firstNonEmpty(strings.TrimSpace(event.Exit.GetCwd()), state.WorkingDirectory)
state.Status = backgroundShellStatusCompleted
state.LastActivityAt = now
state.CompletedAt = now
case *agentv1.ShellStream_Rejected:
state := ensureBackgroundShellStateLocked(stream, shellID, pending, now)
if state == nil {
return
}
state.Status = backgroundShellStatusRejected
state.StderrBuffer += strings.TrimSpace(event.Rejected.GetReason())
state.LastActivityAt = now
state.CompletedAt = now
case *agentv1.ShellStream_PermissionDenied:
state := ensureBackgroundShellStateLocked(stream, shellID, pending, now)
if state == nil {
return
}
state.Status = backgroundShellStatusPermissionDenied
state.StderrBuffer += strings.TrimSpace(event.PermissionDenied.GetError())
state.LastActivityAt = now
state.CompletedAt = now
}
stream.UpdatedAt = now
}
func observeBackgroundTaskCompletionAction(stream *ActiveStream, message *agentv1.AgentClientMessage) {
if stream == nil || message == nil {
return
}
action := message.GetConversationAction()
if action == nil {
return
}
item, ok := action.GetAction().(*agentv1.ConversationAction_BackgroundTaskCompletionAction)
if !ok || item.BackgroundTaskCompletionAction == nil {
return
}
now := time.Now().UTC()
stream.mu.Lock()
defer stream.mu.Unlock()
for _, completion := range item.BackgroundTaskCompletionAction.GetCompletions() {
observeBackgroundTaskCompletionLocked(stream, completion, now)
}
}
func observeBackgroundTaskCompletionLocked(stream *ActiveStream, completion *agentv1.BackgroundTaskCompletion, now time.Time) bool {
if stream == nil || completion == nil || completion.GetKind() != agentv1.BackgroundTaskKind_BACKGROUND_TASK_KIND_SHELL {
return false
}
shellID := strings.TrimSpace(completion.GetTaskId())
if shellID == "" {
return false
}
state := ensureBackgroundShellStateLocked(stream, shellID, runtimecore.PendingExec{}, now)
if state == nil {
return false
}
detail := strings.TrimSpace(completion.GetDetail())
switch completion.GetStatus() {
case agentv1.BackgroundTaskStatus_BACKGROUND_TASK_STATUS_SUCCESS:
state.Status = backgroundShellStatusCompleted
if detail != "" {
state.StdoutBuffer = appendBackgroundShellBuffer(state.StdoutBuffer, detail)
}
case agentv1.BackgroundTaskStatus_BACKGROUND_TASK_STATUS_ERROR:
state.Status = backgroundShellStatusTransportClosed
if detail != "" {
state.StderrBuffer = appendBackgroundShellBuffer(state.StderrBuffer, detail)
}
case agentv1.BackgroundTaskStatus_BACKGROUND_TASK_STATUS_ABORTED:
state.Status = backgroundShellStatusTransportClosed
if detail != "" {
state.StderrBuffer = appendBackgroundShellBuffer(state.StderrBuffer, detail)
}
default:
if completion.GetReason() == agentv1.BackgroundTaskCompletionReason_BACKGROUND_TASK_COMPLETION_REASON_TASK_PROGRESS {
state.Status = backgroundShellStatusRunning
} else {
return false
}
}
state.LastActivityAt = now
if completion.GetReason() == agentv1.BackgroundTaskCompletionReason_BACKGROUND_TASK_COMPLETION_REASON_TASK_FINISHED || isBackgroundShellTerminalStatus(state.Status) {
state.CompletedAt = now
}
stream.UpdatedAt = now
return true
}
func observeBackgroundShellAction(stream *ActiveStream, message *agentv1.AgentClientMessage) (string, bool) {
if stream == nil || message == nil {
return "", false
}
action := message.GetConversationAction()
if action == nil {
return "", false
}
item, ok := action.GetAction().(*agentv1.ConversationAction_BackgroundShellAction)
if !ok || item.BackgroundShellAction == nil {
return "", false
}
return recordBackgroundShellActionMemory(stream, item.BackgroundShellAction.GetToolCallId(), time.Now().UTC())
}
func recordBackgroundShellActionMemory(stream *ActiveStream, toolCallID string, now time.Time) (string, bool) {
trimmedToolCallID := strings.TrimSpace(toolCallID)
if stream == nil || trimmedToolCallID == "" {
return "", false
}
stream.mu.Lock()
defer stream.mu.Unlock()
return recordBackgroundShellActionLocked(stream, trimmedToolCallID, now)
}
func recordBackgroundShellActionLocked(stream *ActiveStream, toolCallID string, now time.Time) (string, bool) {
trimmedToolCallID := strings.TrimSpace(toolCallID)
if stream == nil || trimmedToolCallID == "" {
return "", false
}
if stream.BackgroundShellActions == nil {
stream.BackgroundShellActions = make(map[string]time.Time)
}
if _, exists := stream.BackgroundShellActions[trimmedToolCallID]; exists {
return trimmedToolCallID, false
}
stream.BackgroundShellActions[trimmedToolCallID] = now
stream.UpdatedAt = now
return trimmedToolCallID, true
}
func newBackgroundShellActionMetadataEntry(turnSeq int64, requestID string, toolCallID string, source string) HistoryEntry {
return newMetadataEntry(turnSeq, requestID, "background_shell_action", map[string]any{
"tool_call_id": strings.TrimSpace(toolCallID),
"source": strings.TrimSpace(source),
})
}
func backgroundTaskCompletionMetadataEntries(turnSeq int64, requestID string, message *agentv1.AgentClientMessage) []HistoryEntry {
if message == nil || message.GetConversationAction() == nil {
return nil
}
item, ok := message.GetConversationAction().GetAction().(*agentv1.ConversationAction_BackgroundTaskCompletionAction)
if !ok || item.BackgroundTaskCompletionAction == nil {
return nil
}
completions := item.BackgroundTaskCompletionAction.GetCompletions()
if len(completions) == 0 {
return nil
}
entries := make([]HistoryEntry, 0, len(completions))
for _, completion := range completions {
if completion == nil {
continue
}
values := map[string]any{
"task_id": strings.TrimSpace(completion.GetTaskId()),
"kind": completion.GetKind().String(),
"status": completion.GetStatus().String(),
"reason": completion.GetReason().String(),
}
if title := strings.TrimSpace(completion.GetTitle()); title != "" {
values["title"] = title
}
if detail := strings.TrimSpace(completion.GetDetail()); detail != "" {
values["detail"] = detail
}
if outputPath := strings.TrimSpace(completion.GetOutputPath()); outputPath != "" {
values["output_path"] = outputPath
}
if threadID := strings.TrimSpace(completion.GetThreadId()); threadID != "" {
values["thread_id"] = threadID
}
entries = append(entries, newMetadataEntry(turnSeq, requestID, "background_task_completion_action", values))
}
return entries
}
func shellToolCallIsBackgrounded(toolCall *agentv1.ToolCall) bool {
if toolCall == nil {
return false
}
shellToolCall := toolCall.GetShellToolCall()
if shellToolCall == nil || shellToolCall.GetResult() == nil {
return false
}
return shellToolCall.GetResult().GetIsBackground()
}
func appendBackgroundShellBuffer(current string, value string) string {
trimmed := strings.TrimSpace(value)
if trimmed == "" {
return current
}
if current == "" || strings.HasSuffix(current, "\n") {
return current + trimmed
}
return current + "\n" + trimmed
}
func ensureBackgroundShellStateLocked(stream *ActiveStream, shellID string, pending runtimecore.PendingExec, now time.Time) *BackgroundShellState {
trimmedShellID := strings.TrimSpace(shellID)
if trimmedShellID == "" {
return nil
}
if stream.BackgroundShells == nil {
stream.BackgroundShells = make(map[string]*BackgroundShellState)
}
if stream.BackgroundShellsByMessageID == nil {
stream.BackgroundShellsByMessageID = make(map[uint32]string)
}
if stream.BackgroundShellsByExecID == nil {
stream.BackgroundShellsByExecID = make(map[string]string)
}
state := stream.BackgroundShells[trimmedShellID]
if state == nil {
state = &BackgroundShellState{ShellID: trimmedShellID, Status: backgroundShellStatusRunning, CreatedAt: now}
stream.BackgroundShells[trimmedShellID] = state
}
if pending.MessageID != 0 {
stream.BackgroundShellsByMessageID[pending.MessageID] = trimmedShellID
}
if strings.TrimSpace(pending.ExecID) != "" {
stream.BackgroundShellsByExecID[strings.TrimSpace(pending.ExecID)] = trimmedShellID
}
if state.OriginalToolCallID == "" {
state.OriginalToolCallID = strings.TrimSpace(pending.ToolCallID)
}
if state.OriginalExecID == "" {
state.OriginalExecID = strings.TrimSpace(pending.ExecID)
}
if state.OriginalMessageID == 0 {
state.OriginalMessageID = pending.MessageID
}
if len(state.ArgsJSON) == 0 && len(pending.ArgsJSON) > 0 {
state.ArgsJSON = append([]byte(nil), pending.ArgsJSON...)
}
if state.ModelCallID == "" {
state.ModelCallID = strings.TrimSpace(pending.ModelCallID)
}
return state
}
func backgroundShellIDForMessageLocked(stream *ActiveStream, messageID uint32, execID string) string {
if stream == nil {
return ""
}
if strings.TrimSpace(execID) != "" {
if shellID := strings.TrimSpace(stream.BackgroundShellsByExecID[strings.TrimSpace(execID)]); shellID != "" {
return shellID
}
}
if messageID != 0 {
return strings.TrimSpace(stream.BackgroundShellsByMessageID[messageID])
}
return ""
}
+442
View File
@@ -0,0 +1,442 @@
// broker.go 负责 request 维度活动流的订阅、广播、取消和终态收口。
package forwarder
import (
"fmt"
"strings"
"sync"
"sync/atomic"
"time"
"cursor/gen/agentv1"
runtimecore "cursor/internal/backend/agent/core"
)
const subscriberSignalBufferSize = 1
const orphanSubscriberGracePeriod = 30 * time.Second
const terminalStreamRetentionPeriod = 30 * time.Second
type StreamBroker struct {
mu sync.RWMutex
streams map[string]*ActiveStream
nextID atomic.Uint64
}
// NewStreamBroker 创建活动流注册表。
func NewStreamBroker() *StreamBroker {
return &StreamBroker{
streams: make(map[string]*ActiveStream),
}
}
// OpenStream 打开或复用指定 request 的活动流,并刷新其最新上下文。
func (broker *StreamBroker) OpenStream(requestID string, conversationID string, turnSeq int64, modelID string, modelName string, mode agentv1.AgentMode, latestUserText string) (*ActiveStream, error) {
normalizedRequestID := strings.TrimSpace(requestID)
if normalizedRequestID == "" {
return nil, nil
}
normalizedMode, err := validateSupportedActiveMode(mode)
if err != nil {
return nil, err
}
broker.mu.Lock()
defer broker.mu.Unlock()
if existing, ok := broker.streams[normalizedRequestID]; ok {
existing.mu.Lock()
existing.ConversationID = strings.TrimSpace(conversationID)
existing.TurnSeq = turnSeq
existing.ModelID = strings.TrimSpace(modelID)
existing.ModelName = strings.TrimSpace(modelName)
existing.Mode = normalizedMode
existing.LatestUserText = strings.TrimSpace(latestUserText)
if existing.Status == "" {
existing.Status = StreamStatusCreated
}
if existing.PendingExecs == nil {
existing.PendingExecs = make(map[string]runtimecore.PendingExec)
}
if existing.PendingInteractions == nil {
existing.PendingInteractions = make(map[string]runtimecore.PendingInteraction)
}
if existing.PartialToolCallIDs == nil {
existing.PartialToolCallIDs = make(map[string]struct{})
}
if existing.PatchEditQueues == nil {
existing.PatchEditQueues = make(map[string][]queuedPatchEditOperation)
}
if existing.BackgroundShells == nil {
existing.BackgroundShells = make(map[string]*BackgroundShellState)
}
if existing.BackgroundShellsByMessageID == nil {
existing.BackgroundShellsByMessageID = make(map[uint32]string)
}
if existing.BackgroundShellsByExecID == nil {
existing.BackgroundShellsByExecID = make(map[string]string)
}
if existing.BackgroundShellActions == nil {
existing.BackgroundShellActions = make(map[string]time.Time)
}
existing.UpdatedAt = time.Now().UTC()
existing.mu.Unlock()
return existing, nil
}
now := time.Now().UTC()
stream := &ActiveStream{
RequestID: normalizedRequestID,
ConversationID: strings.TrimSpace(conversationID),
TurnSeq: turnSeq,
ModelID: strings.TrimSpace(modelID),
ModelName: strings.TrimSpace(modelName),
Mode: normalizedMode,
LatestUserText: strings.TrimSpace(latestUserText),
Status: StreamStatusCreated,
Backlog: make([]StreamEvent, 0, 64),
Subscribers: make(map[string]*StreamSubscriber),
PendingExecs: make(map[string]runtimecore.PendingExec),
PendingInteractions: make(map[string]runtimecore.PendingInteraction),
PartialToolCallIDs: make(map[string]struct{}),
PatchEditQueues: make(map[string][]queuedPatchEditOperation),
MCPToolServers: make(map[string]string),
RecentCompletedExecs: make(map[uint32]time.Time),
BackgroundShells: make(map[string]*BackgroundShellState),
BackgroundShellsByMessageID: make(map[uint32]string),
BackgroundShellsByExecID: make(map[string]string),
BackgroundShellActions: make(map[string]time.Time),
CreatedAt: now,
UpdatedAt: now,
}
broker.streams[normalizedRequestID] = stream
return stream, nil
}
// Get 返回指定 request 对应的活动流句柄。
func (broker *StreamBroker) Get(requestID string) (*ActiveStream, bool) {
if broker == nil {
return nil, false
}
broker.mu.RLock()
defer broker.mu.RUnlock()
stream, ok := broker.streams[strings.TrimSpace(requestID)]
return stream, ok
}
// Subscribe 为指定 request 注册一个新订阅者,并返回用于唤醒 backlog 消费的信号通道。
func (broker *StreamBroker) Subscribe(requestID string) (string, <-chan struct{}, error) {
normalizedRequestID := strings.TrimSpace(requestID)
if normalizedRequestID == "" {
return "", nil, fmt.Errorf("request_id is required")
}
stream, ok := broker.Get(normalizedRequestID)
if !ok || stream == nil {
// RunSSE 可能先于 BidiAppend 到达。此时先创建一个占位活动流,
// 等待后续上行把真实 conversation/model/mode 信息补齐。
var err error
stream, err = broker.OpenStream(normalizedRequestID, "", 0, "", "", agentv1.AgentMode_AGENT_MODE_AGENT, "")
if err != nil {
return "", nil, err
}
}
subscriberID := fmt.Sprintf("sub-%d", broker.nextID.Add(1))
subscriber := &StreamSubscriber{Signal: make(chan struct{}, subscriberSignalBufferSize)}
stream.mu.Lock()
broker.stopTerminalCleanupTimerLocked(stream)
stream.Subscribers[subscriberID] = subscriber
stream.UpdatedAt = time.Now().UTC()
stream.mu.Unlock()
return subscriberID, subscriber.Signal, nil
}
func (broker *StreamBroker) stopTerminalCleanupTimerLocked(stream *ActiveStream) {
if stream == nil {
return
}
stream.TerminalCleanupSeq.Add(1)
if stream.TerminalCleanupTimer != nil {
stream.TerminalCleanupTimer.Stop()
stream.TerminalCleanupTimer = nil
}
}
// Unsubscribe 移除并关闭指定订阅者,并返回移除后的剩余订阅者数量。
func (broker *StreamBroker) Unsubscribe(requestID string, subscriberID string) int {
stream, ok := broker.Get(requestID)
if !ok || stream == nil {
return 0
}
remaining := 0
stream.mu.Lock()
if _, ok := stream.Subscribers[strings.TrimSpace(subscriberID)]; ok {
delete(stream.Subscribers, strings.TrimSpace(subscriberID))
}
remaining = len(stream.Subscribers)
stream.mu.Unlock()
return remaining
}
func (broker *StreamBroker) OtherConversationRequestIDs(conversationID string, keepRequestID string) []string {
normalizedConversationID := strings.TrimSpace(conversationID)
normalizedKeepRequestID := strings.TrimSpace(keepRequestID)
if normalizedConversationID == "" {
return nil
}
type requestStream struct {
requestID string
stream *ActiveStream
}
candidates := make([]requestStream, 0, 2)
broker.mu.RLock()
for requestID, stream := range broker.streams {
if stream == nil || strings.TrimSpace(requestID) == normalizedKeepRequestID {
continue
}
candidates = append(candidates, requestStream{
requestID: requestID,
stream: stream,
})
}
broker.mu.RUnlock()
requestIDs := make([]string, 0, 2)
for _, candidate := range candidates {
stream := candidate.stream
stream.mu.Lock()
sameConversation := strings.TrimSpace(stream.ConversationID) == normalizedConversationID
status := stream.Status
phase := stream.Phase
stream.mu.Unlock()
terminalPhase := phase == TurnPhaseCanceled || phase == TurnPhaseCompleted || phase == TurnPhaseFailed
if !sameConversation || isTerminalStreamStatus(status) || terminalPhase {
continue
}
requestIDs = append(requestIDs, candidate.requestID)
}
return requestIDs
}
func (broker *StreamBroker) scheduleTerminalCleanup(requestID string) bool {
stream, ok := broker.Get(requestID)
if !ok || stream == nil {
return false
}
stream.mu.Lock()
defer stream.mu.Unlock()
if len(stream.Subscribers) > 0 {
broker.stopTerminalCleanupTimerLocked(stream)
return false
}
if stream.Status != StreamStatusCompleted && stream.Status != StreamStatusCanceled && stream.Status != StreamStatusFailed {
broker.stopTerminalCleanupTimerLocked(stream)
return false
}
sequence := stream.TerminalCleanupSeq.Add(1)
if stream.TerminalCleanupTimer != nil {
stream.TerminalCleanupTimer.Stop()
}
stream.TerminalCleanupTimer = time.AfterFunc(terminalStreamRetentionPeriod, func() {
broker.runScheduledTerminalCleanup(requestID, sequence)
})
stream.UpdatedAt = time.Now().UTC()
return true
}
func (broker *StreamBroker) runScheduledTerminalCleanup(requestID string, sequence uint64) {
stream, ok := broker.Get(requestID)
if !ok || stream == nil {
return
}
stream.mu.Lock()
if stream.TerminalCleanupSeq.Load() != sequence {
stream.mu.Unlock()
return
}
stream.TerminalCleanupTimer = nil
if len(stream.Subscribers) > 0 {
stream.mu.Unlock()
return
}
if stream.Status != StreamStatusCompleted && stream.Status != StreamStatusCanceled && stream.Status != StreamStatusFailed {
stream.mu.Unlock()
return
}
stream.mu.Unlock()
broker.RemoveIfIdle(requestID)
}
// RemoveIfIdle 在没有订阅者时移除终态流,或移除仍为空壳的占位流。
func (broker *StreamBroker) RemoveIfIdle(requestID string) bool {
normalizedRequestID := strings.TrimSpace(requestID)
if normalizedRequestID == "" {
return false
}
broker.mu.Lock()
defer broker.mu.Unlock()
stream, ok := broker.streams[normalizedRequestID]
if !ok || stream == nil {
return false
}
stream.mu.Lock()
subscriberCount := len(stream.Subscribers)
isActive := stream.ProviderActive
hasBacklog := len(stream.Backlog) > 0
hasConversation := strings.TrimSpace(stream.ConversationID) != ""
status := stream.Status
if status == StreamStatusCompleted || status == StreamStatusCanceled || status == StreamStatusFailed {
broker.stopTerminalCleanupTimerLocked(stream)
}
stream.mu.Unlock()
if subscriberCount > 0 {
return false
}
if status == StreamStatusCompleted || status == StreamStatusCanceled || status == StreamStatusFailed {
delete(broker.streams, normalizedRequestID)
return true
}
if isActive || hasBacklog || hasConversation {
return false
}
delete(broker.streams, normalizedRequestID)
return true
}
// Publish 把一个事件写入 backlog,并唤醒当前所有订阅者读取 backlog。
func (broker *StreamBroker) Publish(requestID string, event StreamEvent) error {
stream, ok := broker.Get(requestID)
if !ok || stream == nil {
return fmt.Errorf("request is not active: %s", strings.TrimSpace(requestID))
}
stream.mu.Lock()
if !event.End && isTerminalStreamStatus(stream.Status) {
stream.mu.Unlock()
return nil
}
stream.Backlog = append(stream.Backlog, event)
stream.UpdatedAt = time.Now().UTC()
subscribers := make([]*StreamSubscriber, 0, len(stream.Subscribers))
for _, subscriber := range stream.Subscribers {
subscribers = append(subscribers, subscriber)
}
stream.mu.Unlock()
for _, subscriber := range subscribers {
if subscriber == nil {
continue
}
select {
case subscriber.Signal <- struct{}{}:
default:
}
}
return nil
}
// ReadFromCursor 返回从 cursor 开始尚未消费的 backlog 事件副本。
func (broker *StreamBroker) ReadFromCursor(requestID string, cursor int) ([]StreamEvent, error) {
stream, ok := broker.Get(requestID)
if !ok || stream == nil {
return nil, fmt.Errorf("request is not active: %s", strings.TrimSpace(requestID))
}
stream.mu.Lock()
defer stream.mu.Unlock()
if cursor < 0 {
cursor = 0
}
if cursor >= len(stream.Backlog) {
return nil, nil
}
return append([]StreamEvent(nil), stream.Backlog[cursor:]...), nil
}
// Complete 把活动流标记为成功完成,并发布一个成功 endstream 事件。
func (broker *StreamBroker) Complete(requestID string, terminalCode string, terminalMessage string) error {
stream, ok := broker.Get(requestID)
if !ok || stream == nil {
return fmt.Errorf("request is not active: %s", strings.TrimSpace(requestID))
}
stream.mu.Lock()
if stream.Status == StreamStatusCanceled || stream.Status == StreamStatusFailed || stream.Status == StreamStatusCompleted {
stream.mu.Unlock()
return nil
}
broker.stopTerminalCleanupTimerLocked(stream)
stream.Status = StreamStatusCompleted
subscriberCount := len(stream.Subscribers)
stream.UpdatedAt = time.Now().UTC()
stream.mu.Unlock()
if err := broker.Publish(requestID, StreamEvent{
End: true,
TerminalErrorCode: strings.TrimSpace(terminalCode),
TerminalErrorMessage: strings.TrimSpace(terminalMessage),
}); err != nil {
return err
}
if subscriberCount == 0 {
broker.scheduleTerminalCleanup(requestID)
}
return nil
}
// Fail 把活动流标记为失败,并发布一个失败 endstream 事件。
func (broker *StreamBroker) Fail(requestID string, terminalCode string, terminalMessage string) error {
stream, ok := broker.Get(requestID)
if !ok || stream == nil {
return fmt.Errorf("request is not active: %s", strings.TrimSpace(requestID))
}
stream.mu.Lock()
broker.stopTerminalCleanupTimerLocked(stream)
stream.Status = StreamStatusFailed
subscriberCount := len(stream.Subscribers)
stream.UpdatedAt = time.Now().UTC()
stream.mu.Unlock()
if err := broker.Publish(requestID, StreamEvent{
End: true,
TerminalErrorCode: strings.TrimSpace(terminalCode),
TerminalErrorMessage: strings.TrimSpace(terminalMessage),
}); err != nil {
return err
}
if subscriberCount == 0 {
broker.scheduleTerminalCleanup(requestID)
}
return nil
}
// Cancel 主动取消活动流,并发布 canceled endstream。
func (broker *StreamBroker) Cancel(requestID string, terminalMessage string) error {
stream, ok := broker.Get(requestID)
if !ok || stream == nil {
return fmt.Errorf("request is not active: %s", strings.TrimSpace(requestID))
}
stream.mu.Lock()
broker.stopTerminalCleanupTimerLocked(stream)
if stream.ProviderCancel != nil {
stream.ProviderCancel()
stream.ProviderCancel = nil
}
stream.ProviderActive = false
stream.Status = StreamStatusCanceled
subscriberCount := len(stream.Subscribers)
stream.UpdatedAt = time.Now().UTC()
stream.mu.Unlock()
if err := broker.Publish(requestID, StreamEvent{
End: true,
TerminalErrorCode: "canceled",
TerminalErrorMessage: firstNonEmpty(strings.TrimSpace(terminalMessage), "[canceled] User aborted request"),
}); err != nil {
return err
}
if subscriberCount == 0 {
broker.scheduleTerminalCleanup(requestID)
}
return nil
}
// firstNonEmpty 返回第一个非空白字符串。
func firstNonEmpty(values ...string) string {
for _, value := range values {
if strings.TrimSpace(value) != "" {
return strings.TrimSpace(value)
}
}
return ""
}
@@ -0,0 +1,147 @@
package forwarder
import (
"fmt"
"strings"
"time"
runtimecore "cursor/internal/backend/agent/core"
)
func (service *Service) replaceCheckpointConversation(stream *ActiveStream, conversation *ConversationFile) error {
if stream == nil {
return fmt.Errorf("active stream is required")
}
if conversation == nil {
return fmt.Errorf("checkpoint conversation is required")
}
stream.mu.Lock()
stream.CheckpointConversation = cloneConversationFile(conversation)
stream.mu.Unlock()
return nil
}
func checkpointConversationInitialized(stream *ActiveStream) bool {
if stream == nil {
return false
}
stream.mu.Lock()
defer stream.mu.Unlock()
return stream.CheckpointConversation != nil
}
func (service *Service) appendCheckpointEntries(stream *ActiveStream, entries []HistoryEntry) error {
if len(entries) == 0 {
return nil
}
if stream == nil {
return fmt.Errorf("active stream is required")
}
stream.mu.Lock()
defer stream.mu.Unlock()
if stream.CheckpointConversation == nil {
return fmt.Errorf("checkpoint conversation is not initialized")
}
appendEntriesInPlace(stream.CheckpointConversation, entries)
return nil
}
func (service *Service) snapshotCheckpointConversation(stream *ActiveStream) (*ConversationFile, []runtimecore.PendingExec, []runtimecore.PendingInteraction, error) {
if stream == nil {
return nil, nil, nil, fmt.Errorf("active stream is required")
}
stream.mu.Lock()
defer stream.mu.Unlock()
if stream.CheckpointConversation == nil {
return nil, nil, nil, fmt.Errorf("checkpoint conversation is not initialized")
}
conversation := cloneConversationFile(stream.CheckpointConversation)
pendingExecs := make([]runtimecore.PendingExec, 0, len(stream.PendingExecs))
for _, pending := range stream.PendingExecs {
pendingExecs = append(pendingExecs, pending)
}
pendingInteractions := make([]runtimecore.PendingInteraction, 0, len(stream.PendingInteractions))
for _, pending := range stream.PendingInteractions {
pendingInteractions = append(pendingInteractions, pending)
}
return conversation, pendingExecs, pendingInteractions, nil
}
func (service *Service) appendConversationEntries(stream *ActiveStream, conversationID string, entries []HistoryEntry) ([]HistoryEntry, error) {
if len(entries) == 0 {
return nil, nil
}
if stream == nil {
return nil, fmt.Errorf("active stream is required")
}
stream.mu.Lock()
if stream.CheckpointConversation == nil {
stream.mu.Unlock()
return nil, fmt.Errorf("checkpoint conversation is not initialized")
}
working := cloneConversationFile(stream.CheckpointConversation)
if service.store != nil {
persisted, assigned, err := service.store.AppendEntries(conversationID, resetEntrySequences(entries))
if err != nil {
stream.mu.Unlock()
return nil, err
}
if persisted != nil {
working = persisted
} else {
appendEntriesInPlace(working, assigned)
}
stream.CheckpointConversation = working
stream.UpdatedAt = time.Now().UTC()
stream.mu.Unlock()
return assigned, nil
}
assigned := appendEntriesInPlace(working, entries)
stream.CheckpointConversation = working
stream.UpdatedAt = time.Now().UTC()
stream.mu.Unlock()
return assigned, nil
}
func resetEntrySequences(entries []HistoryEntry) []HistoryEntry {
if len(entries) == 0 {
return nil
}
reset := make([]HistoryEntry, 0, len(entries))
for _, entry := range entries {
next := entry
next.Seq = 0
reset = append(reset, next)
}
return reset
}
func firstWorkspacePath(paths []string) string {
for _, path := range paths {
if strings.TrimSpace(path) != "" {
return strings.TrimSpace(path)
}
}
return ""
}
func (service *Service) updateConversationMetaAndCheckpoint(stream *ActiveStream, conversationID string, update func(*ConversationFile) error) (*ConversationFile, error) {
if stream == nil {
return nil, fmt.Errorf("active stream is required")
}
stream.mu.Lock()
defer stream.mu.Unlock()
if stream.CheckpointConversation == nil {
return nil, fmt.Errorf("checkpoint conversation is not initialized")
}
conversation := cloneConversationFile(stream.CheckpointConversation)
if err := update(conversation); err != nil {
return nil, err
}
if err := service.syncConversationRecord(conversationID, conversation); err != nil {
return nil, err
}
stream.CheckpointConversation = conversation
stream.UpdatedAt = time.Now().UTC()
return cloneConversationFile(conversation), nil
}
@@ -0,0 +1,200 @@
package forwarder
import (
"context"
"crypto/sha256"
"encoding/hex"
"fmt"
"strings"
"connectrpc.com/connect"
"github.com/google/uuid"
"cursor/gen/aiserverv1"
)
func (service *Service) CreateExperimentalIndex(_ context.Context, req *connect.Request[aiserverv1.CreateExperimentalIndexRequest]) (*connect.Response[aiserverv1.CreateExperimentalIndexResponse], error) {
store, err := service.requireCodebaseIndexStore()
if err != nil {
return nil, connect.NewError(connect.CodeInternal, err)
}
record, err := store.EnsureIndexed(CodebaseIndexRecord{
Repo: strings.TrimSpace(req.Msg.GetRepo()),
TargetDir: strings.TrimSpace(req.Msg.GetTargetDir()),
Files: append([]string(nil), req.Msg.GetFiles()...),
})
if err != nil {
return nil, connect.NewError(connect.CodeInternal, err)
}
return connect.NewResponse(&aiserverv1.CreateExperimentalIndexResponse{IndexId: record.IndexID}), nil
}
func (service *Service) ListExperimentalIndexFiles(_ context.Context, req *connect.Request[aiserverv1.ListExperimentalIndexFilesRequest]) (*connect.Response[aiserverv1.ListExperimentalIndexFilesResponse], error) {
store, err := service.requireCodebaseIndexStore()
if err != nil {
return nil, connect.NewError(connect.CodeInternal, err)
}
indexID := strings.TrimSpace(req.Msg.GetIndexId())
files, err := store.ListFiles(indexID)
if err != nil {
return nil, connect.NewError(connect.CodeInternal, err)
}
return connect.NewResponse(&aiserverv1.ListExperimentalIndexFilesResponse{
IndexId: indexID,
Files: codebaseIndexFileData(indexID, files),
}), nil
}
func (service *Service) ListenExperimentalIndex(ctx context.Context, req *connect.Request[aiserverv1.ListenExperimentalIndexRequest], stream *connect.ServerStream[aiserverv1.ListenExperimentalIndexResponse]) error {
store, err := service.requireCodebaseIndexStore()
if err != nil {
return connect.NewError(connect.CodeInternal, err)
}
indexID := strings.TrimSpace(req.Msg.GetIndexId())
if indexID == "" {
return connect.NewError(connect.CodeInvalidArgument, fmt.Errorf("index_id is required"))
}
_, err = store.EnsureIndexed(CodebaseIndexRecord{IndexID: indexID})
if err != nil {
return connect.NewError(connect.CodeInternal, err)
}
subscription, err := store.Subscribe(indexID)
if err != nil {
return connect.NewError(connect.CodeInternal, err)
}
defer store.Unsubscribe(subscription)
if err := stream.Send(&aiserverv1.ListenExperimentalIndexResponse{
IndexId: indexID,
Item: &aiserverv1.ListenExperimentalIndexResponse_Ready{
Ready: &aiserverv1.ListenExperimentalIndexResponse_ReadyItem{
IndexId: indexID,
Request: &aiserverv1.ListenExperimentalIndexRequest{
IndexId: indexID,
},
},
},
}); err != nil {
return err
}
var lastEventID int64
for {
events, err := store.ListEvents(indexID, lastEventID)
if err != nil {
return connect.NewError(connect.CodeInternal, err)
}
for _, event := range events {
if event.ID > lastEventID {
lastEventID = event.ID
}
if event.Kind != codebaseIndexEventKindRegister {
continue
}
if err := stream.Send(codebaseIndexRegisterEventResponse(indexID, event)); err != nil {
return err
}
}
select {
case <-ctx.Done():
return nil
case <-subscription.Signal:
}
}
}
func (service *Service) RegisterFileToIndex(_ context.Context, req *connect.Request[aiserverv1.RegisterFileToIndexRequest]) (*connect.Response[aiserverv1.RequestReceivedResponse], error) {
store, err := service.requireCodebaseIndexStore()
if err != nil {
return nil, connect.NewError(connect.CodeInternal, err)
}
reqUUID := uuid.NewString()
_, _, err = store.RegisterFileWithRequest(req.Msg.GetIndexId(), reqUUID, CodebaseIndexFileRecord{
WorkspaceRelativePath: strings.TrimSpace(req.Msg.GetWorkspaceRelativePath()),
ContentHash: codebaseFileContentHash(strings.Join(req.Msg.GetContent(), "\n")),
Stage: codebaseIndexFileStage,
})
if err != nil {
return nil, connect.NewError(connect.CodeInvalidArgument, err)
}
return connect.NewResponse(&aiserverv1.RequestReceivedResponse{ReqUuid: reqUUID}), nil
}
func (service *Service) SetupIndexDependencies(_ context.Context, req *connect.Request[aiserverv1.SetupIndexDependenciesRequest]) (*connect.Response[aiserverv1.SetupIndexDependenciesResponse], error) {
store, err := service.requireCodebaseIndexStore()
if err != nil {
return nil, connect.NewError(connect.CodeInternal, err)
}
if err := store.MarkDependenciesReady(req.Msg.GetIndexId()); err != nil {
return nil, connect.NewError(connect.CodeInvalidArgument, err)
}
return connect.NewResponse(&aiserverv1.SetupIndexDependenciesResponse{}), nil
}
func (service *Service) ComputeIndexTopoSort(_ context.Context, req *connect.Request[aiserverv1.ComputeIndexTopoSortRequest]) (*connect.Response[aiserverv1.ComputeIndexTopoSortResponse], error) {
store, err := service.requireCodebaseIndexStore()
if err != nil {
return nil, connect.NewError(connect.CodeInternal, err)
}
if err := store.MarkTopoSortReady(req.Msg.GetIndexId()); err != nil {
return nil, connect.NewError(connect.CodeInvalidArgument, err)
}
return connect.NewResponse(&aiserverv1.ComputeIndexTopoSortResponse{}), nil
}
func (service *Service) requireCodebaseIndexStore() (*CodebaseIndexStore, error) {
if service == nil || service.codebaseIndexStore == nil {
return nil, fmt.Errorf("codebase index store is not initialized")
}
return service.codebaseIndexStore, nil
}
func codebaseIndexFileData(indexID string, files []CodebaseIndexFileRecord) []*aiserverv1.IndexFileData {
result := make([]*aiserverv1.IndexFileData, 0, len(files))
for index, file := range files {
result = append(result, codebaseIndexFileDatum(indexID, file, int32(index+1)))
}
return result
}
func codebaseIndexFileDatum(indexID string, file CodebaseIndexFileRecord, fallbackOrder int32) *aiserverv1.IndexFileData {
order := file.Order
if order == 0 {
order = fallbackOrder
}
return &aiserverv1.IndexFileData{
IndexId: indexID,
WorkspaceRelativePath: file.WorkspaceRelativePath,
Stage: file.Stage,
Order: order,
}
}
func codebaseIndexRegisterEventResponse(indexID string, event CodebaseIndexEventRecord) *aiserverv1.ListenExperimentalIndexResponse {
fileData := codebaseIndexFileDatum(indexID, event.File, 1)
return &aiserverv1.ListenExperimentalIndexResponse{
IndexId: indexID,
Item: &aiserverv1.ListenExperimentalIndexResponse_Register{
Register: &aiserverv1.ListenExperimentalIndexResponse_RegisterItem{
ReqUuid: strings.TrimSpace(event.ReqUUID),
Request: &aiserverv1.RegisterFileToIndexRequest{
IndexId: indexID,
WorkspaceRelativePath: event.File.WorkspaceRelativePath,
},
Response: &aiserverv1.RegisterFileToIndexResponse{
FileId: codebaseIndexFileID(indexID, event.File.WorkspaceRelativePath),
RootContextNodeId: codebaseIndexFileID(indexID, event.File.WorkspaceRelativePath),
FileData: fileData,
},
},
},
}
}
func codebaseIndexFileID(indexID string, workspaceRelativePath string) string {
trimmedIndexID := strings.TrimSpace(indexID)
trimmedPath := strings.TrimSpace(workspaceRelativePath)
if trimmedIndexID == "" || trimmedPath == "" {
return ""
}
sum := sha256.Sum256([]byte(trimmedIndexID + "\x00" + trimmedPath))
return "file_" + hex.EncodeToString(sum[:])[:24]
}
@@ -0,0 +1,606 @@
package forwarder
import (
"crypto/sha256"
"encoding/hex"
"encoding/json"
"fmt"
"os"
"path/filepath"
"sort"
"strings"
"sync"
"time"
"github.com/google/uuid"
)
const (
codebaseIndexStatusReady = "ready"
codebaseIndexFileStage = "indexed"
codebaseIndexEventKindRegister = "register"
codebaseIndexSubscriberSignalCapacity = 1
codebaseIndexMaxPersistedEvents = 2000
codebaseIndexSourceRepositoryService = "repository_service"
repositoryIndexStatusReady = "ready"
repositoryIndexSyncStatusCompleted = "sync_completed"
repositoryIndexCopyStatusCompleted = "copy_completed"
)
type CodebaseSearchBackend interface {
EnsureIndexed(record CodebaseIndexRecord) (CodebaseIndexRecord, error)
RegisterFile(indexID string, file CodebaseIndexFileRecord) (CodebaseIndexFileRecord, error)
ListFiles(indexID string) ([]CodebaseIndexFileRecord, error)
}
type CodebaseIndexStore struct {
mu sync.Mutex
root string
path string
loaded bool
nextEventID int64
state codebaseIndexState
subscribers map[string]map[string]chan struct{}
}
type codebaseIndexState struct {
SchemaVersion int `json:"schema_version"`
Indexes map[string]CodebaseIndexRecord `json:"indexes"`
Events []CodebaseIndexEventRecord `json:"events,omitempty"`
}
type CodebaseIndexRecord struct {
IndexID string `json:"index_id"`
Repo string `json:"repo,omitempty"`
TargetDir string `json:"target_dir,omitempty"`
Files []string `json:"files,omitempty"`
Status string `json:"status"`
Source string `json:"source,omitempty"`
DependenciesReady bool `json:"dependencies_ready,omitempty"`
TopoSortReady bool `json:"topo_sort_ready,omitempty"`
RepositoryState *RepositoryIndexStateRecord `json:"repository_state,omitempty"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
RegisteredFiles map[string]CodebaseIndexFileRecord `json:"registered_files,omitempty"`
}
type RepositoryIndexStateRecord struct {
CodebaseID string `json:"codebase_id"`
Repo string `json:"repo,omitempty"`
TargetDir string `json:"target_dir,omitempty"`
Status string `json:"status"`
CopyStatus string `json:"copy_status,omitempty"`
SyncStatus string `json:"sync_status,omitempty"`
CopyTaskHandle string `json:"copy_task_handle,omitempty"`
PathKeyHash string `json:"path_key_hash,omitempty"`
Source string `json:"source"`
LastHandshakeAt time.Time `json:"last_handshake_at,omitempty"`
LastSyncAt time.Time `json:"last_sync_at,omitempty"`
LastCopyStatusAt time.Time `json:"last_copy_status_at,omitempty"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
type CodebaseIndexFileRecord struct {
WorkspaceRelativePath string `json:"workspace_relative_path"`
ContentHash string `json:"content_hash,omitempty"`
Stage string `json:"stage"`
Order int32 `json:"order,omitempty"`
NodeCount int32 `json:"node_count,omitempty"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
type CodebaseIndexEventRecord struct {
ID int64 `json:"id"`
IndexID string `json:"index_id"`
Kind string `json:"kind"`
ReqUUID string `json:"req_uuid,omitempty"`
File CodebaseIndexFileRecord `json:"file,omitempty"`
At time.Time `json:"at"`
}
type CodebaseIndexSubscription struct {
IndexID string
SubscriberID string
Signal <-chan struct{}
}
func NewCodebaseIndexStore(root string) *CodebaseIndexStore {
trimmed := strings.TrimSpace(root)
if trimmed == "" {
trimmed = "codebase-index"
}
return &CodebaseIndexStore{
root: trimmed,
path: filepath.Join(trimmed, "index.json"),
subscribers: make(map[string]map[string]chan struct{}),
state: codebaseIndexState{
SchemaVersion: 1,
Indexes: make(map[string]CodebaseIndexRecord),
},
}
}
func (store *CodebaseIndexStore) EnsureIndexed(record CodebaseIndexRecord) (CodebaseIndexRecord, error) {
if store == nil {
return CodebaseIndexRecord{}, fmt.Errorf("codebase index store is nil")
}
store.mu.Lock()
defer store.mu.Unlock()
if err := store.loadLocked(); err != nil {
return CodebaseIndexRecord{}, err
}
now := time.Now().UTC()
indexID := strings.TrimSpace(record.IndexID)
if indexID == "" {
indexID = stableCodebaseIndexID(record.Repo, record.TargetDir, record.Files)
}
existing, ok := store.state.Indexes[indexID]
if ok {
existing.Repo = firstNonEmptyCodebaseIndex(record.Repo, existing.Repo)
existing.TargetDir = firstNonEmptyCodebaseIndex(record.TargetDir, existing.TargetDir)
if len(record.Files) > 0 {
existing.Files = append([]string(nil), record.Files...)
}
if strings.TrimSpace(record.Source) != "" {
existing.Source = strings.TrimSpace(record.Source)
}
if record.RepositoryState != nil {
existing.RepositoryState = mergeRepositoryIndexState(existing.RepositoryState, record.RepositoryState, now)
}
existing.Status = codebaseIndexStatusReady
existing.UpdatedAt = now
if existing.RegisteredFiles == nil {
existing.RegisteredFiles = make(map[string]CodebaseIndexFileRecord)
}
store.state.Indexes[indexID] = existing
return existing, store.saveLocked()
}
record.IndexID = indexID
record.Status = codebaseIndexStatusReady
record.Source = strings.TrimSpace(record.Source)
record.CreatedAt = now
record.UpdatedAt = now
if record.RepositoryState != nil {
record.RepositoryState = mergeRepositoryIndexState(nil, record.RepositoryState, now)
}
record.Files = append([]string(nil), record.Files...)
record.RegisteredFiles = make(map[string]CodebaseIndexFileRecord)
store.state.Indexes[indexID] = record
return record, store.saveLocked()
}
func (store *CodebaseIndexStore) EnsureRepositoryIndexed(state RepositoryIndexStateRecord) (CodebaseIndexRecord, error) {
if store == nil {
return CodebaseIndexRecord{}, fmt.Errorf("codebase index store is nil")
}
codebaseID := strings.TrimSpace(state.CodebaseID)
if codebaseID == "" {
return CodebaseIndexRecord{}, fmt.Errorf("codebase_id is required")
}
state.CodebaseID = codebaseID
state.Source = codebaseIndexSourceRepositoryService
if strings.TrimSpace(state.Status) == "" {
state.Status = repositoryIndexStatusReady
}
return store.EnsureIndexed(CodebaseIndexRecord{
IndexID: codebaseID,
Repo: strings.TrimSpace(state.Repo),
TargetDir: strings.TrimSpace(state.TargetDir),
Source: codebaseIndexSourceRepositoryService,
RepositoryState: &state,
})
}
func (store *CodebaseIndexStore) MarkRepositorySyncComplete(codebaseID string, pathKeyHash string) error {
_, err := store.EnsureRepositoryIndexed(RepositoryIndexStateRecord{
CodebaseID: strings.TrimSpace(codebaseID),
Status: repositoryIndexStatusReady,
SyncStatus: repositoryIndexSyncStatusCompleted,
PathKeyHash: strings.TrimSpace(pathKeyHash),
LastSyncAt: time.Now().UTC(),
})
return err
}
func (store *CodebaseIndexStore) MarkRepositoryCopyComplete(codebaseID string, copyTaskHandle string) error {
_, err := store.EnsureRepositoryIndexed(RepositoryIndexStateRecord{
CodebaseID: strings.TrimSpace(codebaseID),
Status: repositoryIndexStatusReady,
CopyStatus: repositoryIndexCopyStatusCompleted,
CopyTaskHandle: strings.TrimSpace(copyTaskHandle),
LastCopyStatusAt: time.Now().UTC(),
})
return err
}
func (store *CodebaseIndexStore) RegisterFile(indexID string, file CodebaseIndexFileRecord) (CodebaseIndexFileRecord, error) {
registered, _, err := store.RegisterFileWithRequest(indexID, "", file)
return registered, err
}
func (store *CodebaseIndexStore) RegisterFileWithRequest(indexID string, reqUUID string, file CodebaseIndexFileRecord) (CodebaseIndexFileRecord, CodebaseIndexEventRecord, error) {
if store == nil {
return CodebaseIndexFileRecord{}, CodebaseIndexEventRecord{}, fmt.Errorf("codebase index store is nil")
}
store.mu.Lock()
defer store.mu.Unlock()
if err := store.loadLocked(); err != nil {
return CodebaseIndexFileRecord{}, CodebaseIndexEventRecord{}, err
}
trimmedIndexID := strings.TrimSpace(indexID)
if trimmedIndexID == "" {
return CodebaseIndexFileRecord{}, CodebaseIndexEventRecord{}, fmt.Errorf("index_id is required")
}
record, ok := store.state.Indexes[trimmedIndexID]
if !ok {
now := time.Now().UTC()
record = CodebaseIndexRecord{
IndexID: trimmedIndexID,
Status: codebaseIndexStatusReady,
CreatedAt: now,
UpdatedAt: now,
RegisteredFiles: make(map[string]CodebaseIndexFileRecord),
}
}
if record.RegisteredFiles == nil {
record.RegisteredFiles = make(map[string]CodebaseIndexFileRecord)
}
path := strings.TrimSpace(file.WorkspaceRelativePath)
if path == "" {
return CodebaseIndexFileRecord{}, CodebaseIndexEventRecord{}, fmt.Errorf("workspace_relative_path is required")
}
now := time.Now().UTC()
current, exists := record.RegisteredFiles[path]
if exists {
file.CreatedAt = current.CreatedAt
} else {
file.CreatedAt = now
}
file.WorkspaceRelativePath = path
if strings.TrimSpace(file.Stage) == "" {
file.Stage = codebaseIndexFileStage
}
if file.Order == 0 {
if exists && current.Order != 0 {
file.Order = current.Order
} else {
file.Order = int32(len(record.RegisteredFiles) + 1)
}
}
file.UpdatedAt = now
record.RegisteredFiles[path] = file
record.Status = codebaseIndexStatusReady
record.UpdatedAt = now
store.state.Indexes[trimmedIndexID] = record
event := CodebaseIndexEventRecord{
ID: store.nextCodebaseIndexEventIDLocked(),
IndexID: trimmedIndexID,
Kind: codebaseIndexEventKindRegister,
ReqUUID: strings.TrimSpace(reqUUID),
File: file,
At: now,
}
store.state.Events = append(store.state.Events, event)
store.compactEventsLocked()
if err := store.saveLocked(); err != nil {
return CodebaseIndexFileRecord{}, CodebaseIndexEventRecord{}, err
}
store.notifyIndexSubscribersLocked(trimmedIndexID)
return file, event, nil
}
func (store *CodebaseIndexStore) ListFiles(indexID string) ([]CodebaseIndexFileRecord, error) {
if store == nil {
return nil, fmt.Errorf("codebase index store is nil")
}
store.mu.Lock()
defer store.mu.Unlock()
if err := store.loadLocked(); err != nil {
return nil, err
}
record, ok := store.state.Indexes[strings.TrimSpace(indexID)]
if !ok || len(record.RegisteredFiles) == 0 {
return nil, nil
}
files := make([]CodebaseIndexFileRecord, 0, len(record.RegisteredFiles))
for _, file := range record.RegisteredFiles {
files = append(files, file)
}
sort.Slice(files, func(i, j int) bool {
if files[i].Order != files[j].Order {
return files[i].Order < files[j].Order
}
return files[i].WorkspaceRelativePath < files[j].WorkspaceRelativePath
})
return files, nil
}
func (store *CodebaseIndexStore) ListEvents(indexID string, afterEventID int64) ([]CodebaseIndexEventRecord, error) {
if store == nil {
return nil, fmt.Errorf("codebase index store is nil")
}
store.mu.Lock()
defer store.mu.Unlock()
if err := store.loadLocked(); err != nil {
return nil, err
}
trimmedIndexID := strings.TrimSpace(indexID)
if trimmedIndexID == "" {
return nil, fmt.Errorf("index_id is required")
}
events := make([]CodebaseIndexEventRecord, 0)
for _, event := range store.state.Events {
if event.ID <= afterEventID || event.IndexID != trimmedIndexID {
continue
}
events = append(events, event)
}
return events, nil
}
func (store *CodebaseIndexStore) Subscribe(indexID string) (CodebaseIndexSubscription, error) {
if store == nil {
return CodebaseIndexSubscription{}, fmt.Errorf("codebase index store is nil")
}
store.mu.Lock()
defer store.mu.Unlock()
if err := store.loadLocked(); err != nil {
return CodebaseIndexSubscription{}, err
}
trimmedIndexID := strings.TrimSpace(indexID)
if trimmedIndexID == "" {
return CodebaseIndexSubscription{}, fmt.Errorf("index_id is required")
}
if store.subscribers == nil {
store.subscribers = make(map[string]map[string]chan struct{})
}
subscriberID := uuid.NewString()
signal := make(chan struct{}, codebaseIndexSubscriberSignalCapacity)
if store.subscribers[trimmedIndexID] == nil {
store.subscribers[trimmedIndexID] = make(map[string]chan struct{})
}
store.subscribers[trimmedIndexID][subscriberID] = signal
return CodebaseIndexSubscription{
IndexID: trimmedIndexID,
SubscriberID: subscriberID,
Signal: signal,
}, nil
}
func (store *CodebaseIndexStore) Unsubscribe(subscription CodebaseIndexSubscription) {
if store == nil {
return
}
store.mu.Lock()
defer store.mu.Unlock()
indexID := strings.TrimSpace(subscription.IndexID)
subscriberID := strings.TrimSpace(subscription.SubscriberID)
if indexID == "" || subscriberID == "" {
return
}
byIndex := store.subscribers[indexID]
if byIndex == nil {
return
}
delete(byIndex, subscriberID)
if len(byIndex) == 0 {
delete(store.subscribers, indexID)
}
}
func (store *CodebaseIndexStore) MarkDependenciesReady(indexID string) error {
return store.updateIndex(indexID, func(record *CodebaseIndexRecord) {
record.DependenciesReady = true
})
}
func (store *CodebaseIndexStore) MarkTopoSortReady(indexID string) error {
return store.updateIndex(indexID, func(record *CodebaseIndexRecord) {
record.TopoSortReady = true
})
}
func (store *CodebaseIndexStore) updateIndex(indexID string, update func(*CodebaseIndexRecord)) error {
if store == nil {
return fmt.Errorf("codebase index store is nil")
}
store.mu.Lock()
defer store.mu.Unlock()
if err := store.loadLocked(); err != nil {
return err
}
trimmedIndexID := strings.TrimSpace(indexID)
if trimmedIndexID == "" {
return fmt.Errorf("index_id is required")
}
now := time.Now().UTC()
record, ok := store.state.Indexes[trimmedIndexID]
if !ok {
record = CodebaseIndexRecord{
IndexID: trimmedIndexID,
Status: codebaseIndexStatusReady,
CreatedAt: now,
RegisteredFiles: make(map[string]CodebaseIndexFileRecord),
}
}
if record.RegisteredFiles == nil {
record.RegisteredFiles = make(map[string]CodebaseIndexFileRecord)
}
update(&record)
record.Status = codebaseIndexStatusReady
record.UpdatedAt = now
store.state.Indexes[trimmedIndexID] = record
return store.saveLocked()
}
func (store *CodebaseIndexStore) loadLocked() error {
if store.loaded {
return nil
}
state := codebaseIndexState{SchemaVersion: 1, Indexes: make(map[string]CodebaseIndexRecord)}
data, err := os.ReadFile(store.path)
if err != nil {
if os.IsNotExist(err) {
store.state = state
store.nextEventID = 1
store.loaded = true
return nil
}
return err
}
if len(data) != 0 {
if err := json.Unmarshal(data, &state); err != nil {
return err
}
}
if state.SchemaVersion <= 0 {
state.SchemaVersion = 1
}
if state.Indexes == nil {
state.Indexes = make(map[string]CodebaseIndexRecord)
}
var maxEventID int64
for id, record := range state.Indexes {
if record.RegisteredFiles == nil {
record.RegisteredFiles = make(map[string]CodebaseIndexFileRecord)
state.Indexes[id] = record
}
}
for _, event := range state.Events {
if event.ID > maxEventID {
maxEventID = event.ID
}
}
store.state = state
store.nextEventID = maxEventID + 1
store.loaded = true
return nil
}
func (store *CodebaseIndexStore) saveLocked() error {
if err := os.MkdirAll(store.root, 0o755); err != nil {
return err
}
data, err := json.MarshalIndent(store.state, "", " ")
if err != nil {
return err
}
tmp, err := os.CreateTemp(store.root, ".index-*.json.tmp")
if err != nil {
return err
}
tmpPath := tmp.Name()
defer func() {
_ = os.Remove(tmpPath)
}()
if _, err := tmp.Write(data); err != nil {
_ = tmp.Close()
return err
}
if err := tmp.Close(); err != nil {
return err
}
if err := os.Chmod(tmpPath, 0o600); err != nil {
return err
}
return os.Rename(tmpPath, store.path)
}
func (store *CodebaseIndexStore) nextCodebaseIndexEventIDLocked() int64 {
if store.nextEventID <= 0 {
store.nextEventID = 1
}
eventID := store.nextEventID
store.nextEventID++
return eventID
}
func (store *CodebaseIndexStore) notifyIndexSubscribersLocked(indexID string) {
if store.subscribers == nil {
return
}
for _, signal := range store.subscribers[strings.TrimSpace(indexID)] {
select {
case signal <- struct{}{}:
default:
}
}
}
func (store *CodebaseIndexStore) compactEventsLocked() {
if len(store.state.Events) <= codebaseIndexMaxPersistedEvents {
return
}
store.state.Events = append([]CodebaseIndexEventRecord(nil), store.state.Events[len(store.state.Events)-codebaseIndexMaxPersistedEvents:]...)
}
func stableCodebaseIndexID(repo string, targetDir string, files []string) string {
parts := append([]string{strings.TrimSpace(repo), strings.TrimSpace(targetDir)}, files...)
normalized := make([]string, 0, len(parts))
for _, part := range parts {
trimmed := strings.TrimSpace(part)
if trimmed != "" {
normalized = append(normalized, trimmed)
}
}
if len(normalized) == 0 {
return "idx_local"
}
sort.Strings(normalized)
sum := sha256.Sum256([]byte(strings.Join(normalized, "\x00")))
return "idx_" + hex.EncodeToString(sum[:])[:24]
}
func codebaseFileContentHash(content string) string {
if content == "" {
return ""
}
sum := sha256.Sum256([]byte(content))
return hex.EncodeToString(sum[:])
}
func firstNonEmptyCodebaseIndex(values ...string) string {
for _, value := range values {
trimmed := strings.TrimSpace(value)
if trimmed != "" {
return trimmed
}
}
return ""
}
func mergeRepositoryIndexState(existing *RepositoryIndexStateRecord, next *RepositoryIndexStateRecord, now time.Time) *RepositoryIndexStateRecord {
if next == nil {
return existing
}
merged := RepositoryIndexStateRecord{}
if existing != nil {
merged = *existing
}
merged.CodebaseID = firstNonEmptyCodebaseIndex(next.CodebaseID, merged.CodebaseID)
merged.Repo = firstNonEmptyCodebaseIndex(next.Repo, merged.Repo)
merged.TargetDir = firstNonEmptyCodebaseIndex(next.TargetDir, merged.TargetDir)
merged.Status = firstNonEmptyCodebaseIndex(next.Status, merged.Status, repositoryIndexStatusReady)
merged.CopyStatus = firstNonEmptyCodebaseIndex(next.CopyStatus, merged.CopyStatus)
merged.SyncStatus = firstNonEmptyCodebaseIndex(next.SyncStatus, merged.SyncStatus)
merged.CopyTaskHandle = firstNonEmptyCodebaseIndex(next.CopyTaskHandle, merged.CopyTaskHandle)
merged.PathKeyHash = firstNonEmptyCodebaseIndex(next.PathKeyHash, merged.PathKeyHash)
merged.Source = firstNonEmptyCodebaseIndex(next.Source, merged.Source, codebaseIndexSourceRepositoryService)
merged.LastHandshakeAt = firstNonZeroCodebaseIndexTime(next.LastHandshakeAt, merged.LastHandshakeAt)
merged.LastSyncAt = firstNonZeroCodebaseIndexTime(next.LastSyncAt, merged.LastSyncAt)
merged.LastCopyStatusAt = firstNonZeroCodebaseIndexTime(next.LastCopyStatusAt, merged.LastCopyStatusAt)
merged.CreatedAt = firstNonZeroCodebaseIndexTime(merged.CreatedAt, next.CreatedAt, now)
merged.UpdatedAt = now
return &merged
}
func firstNonZeroCodebaseIndexTime(values ...time.Time) time.Time {
for _, value := range values {
if !value.IsZero() {
return value.UTC()
}
}
return time.Time{}
}
@@ -0,0 +1,306 @@
package forwarder
import (
"context"
"errors"
"fmt"
"log"
"strings"
"connectrpc.com/connect"
"github.com/google/uuid"
"google.golang.org/protobuf/encoding/protojson"
"cursor/gen/agentv1"
"cursor/gen/aiserverv1"
modeladapter "cursor/internal/backend/agent/model"
promptassets "cursor/prompt"
)
const (
commitMessageDiffTotalLimit = 120_000
commitMessageSingleDiffLimit = 40_000
commitMessagePreviousCommitLimit = 12
commitMessageExplicitContextLimit = 20_000
commitMessageMaxOutputTokens = 512
commitMessageGeneratedRequestIDPrefix = "commit-message-"
)
var errCommitMessageToolInvocation = errors.New("commit message generation must not invoke tools")
// WriteGitCommitMessage handles Cursor's SCM "Generate Commit Message" action.
func (service *Service) WriteGitCommitMessage(ctx context.Context, req *connect.Request[aiserverv1.WriteGitCommitMessageRequest]) (*connect.Response[aiserverv1.WriteGitCommitMessageResponse], error) {
if service == nil {
return nil, connect.NewError(connect.CodeInternal, fmt.Errorf("forwarder service is nil"))
}
requestID := commitMessageGeneratedRequestIDPrefix + uuid.NewString()
recorder, err := newCommitMessageLogRecorder(service.commitMessageHistoryRoot(), requestID)
if err != nil {
return nil, connect.NewError(connect.CodeInternal, err)
}
if req == nil || req.Msg == nil {
return nil, commitMessageConnectError(recorder, connect.CodeInvalidArgument, fmt.Errorf("write git commit message request is required"))
}
if err := recordCommitMessageIncomingRequest(recorder, requestID, req.Msg); err != nil {
return nil, connect.NewError(connect.CodeInternal, err)
}
diffs := truncateCommitMessageDiffs(req.Msg.GetDiffs(), commitMessageDiffTotalLimit, commitMessageSingleDiffLimit)
if len(diffs) == 0 {
return nil, commitMessageConnectError(recorder, connect.CodeInvalidArgument, fmt.Errorf("diffs are required"))
}
if service.provider == nil {
return nil, commitMessageConnectError(recorder, connect.CodeInternal, fmt.Errorf("provider gateway is not initialized"))
}
messages, err := buildCommitMessagePrompt(req.Msg, diffs)
if err != nil {
return nil, commitMessageConnectError(recorder, connect.CodeInternal, err)
}
modelCallID := requestID + "-model"
modelID, modelSource, lastAgentModelHash := service.resolveCommitMessageModelID(ctx)
accumulated := ""
artifactPaths := &modeladapter.LLMArtifactPaths{}
err = service.provider.StartStream(ctx, ProviderRequest{
RequestID: requestID,
RunID: requestID,
ModelCallID: modelCallID,
ModelID: modelID,
Mode: agentv1.AgentMode_AGENT_MODE_AGENT,
Messages: messages,
Tools: nil,
MaxTokens: commitMessageMaxOutputTokens,
CompileSummary: fmt.Sprintf("generate commit message diffs=%d previous_commits=%d model_source=%s last_agent_model_hash=%s", len(diffs), len(req.Msg.GetPreviousCommitMessages()), modelSource, lastAgentModelHash),
Observer: recorder,
ArtifactPaths: artifactPaths,
}, func(event modeladapter.ModelEvent) error {
switch event.Kind {
case modeladapter.ModelEventKindTextDelta:
accumulated += event.Text
return nil
case modeladapter.ModelEventKindThinkingDelta, modeladapter.ModelEventKindThinkingCompleted, modeladapter.ModelEventKindTurnFinished:
return nil
case modeladapter.ModelEventKindToolLikeCompleted, modeladapter.ModelEventKindPartialToolCall, modeladapter.ModelEventKindToolCallDelta:
return errCommitMessageToolInvocation
case modeladapter.ModelEventKindProviderError:
if event.Err != nil {
return providerTerminalError{cause: event.Err}
}
return providerTerminalError{cause: fmt.Errorf("provider error")}
default:
return nil
}
})
if err != nil {
if errors.Is(err, errCommitMessageToolInvocation) {
return nil, commitMessageConnectError(recorder, connect.CodeInternal, errCommitMessageToolInvocation)
}
return nil, commitMessageConnectError(recorder, connect.CodeUnknown, err)
}
commitMessage := cleanGeneratedCommitMessage(accumulated)
if firstCommitMessageLine(commitMessage) == "" {
return nil, commitMessageConnectError(recorder, connect.CodeInternal, fmt.Errorf("generated commit message is empty"))
}
if _, err := recorder.appendEvent("final_response", map[string]any{
"request_id": requestID,
"model_call_id": modelCallID,
"model_source": modelSource,
"model_id": modelID,
"raw_text": accumulated,
"commit_message": commitMessage,
"artifact_request": artifactPaths.RequestPath,
"artifact_sse": artifactPaths.ResponsePath,
"artifact_summary": artifactPaths.SummaryPath,
}); err != nil {
log.Printf("forwarder failed to write commit message final log request_id=%s error=%v", requestID, err)
}
return connect.NewResponse(&aiserverv1.WriteGitCommitMessageResponse{
CommitMessage: commitMessage,
}), nil
}
func buildCommitMessagePrompt(req *aiserverv1.WriteGitCommitMessageRequest, diffs []string) ([]modeladapter.Message, error) {
system, err := promptassets.ReadCommitPrompt()
if err != nil {
return nil, err
}
if strings.TrimSpace(system) == "" {
return nil, fmt.Errorf("commit prompt asset is empty")
}
sections := []string{"Generate a Git commit message for the following changes."}
previous := truncateCommitMessagePreviousCommits(req.GetPreviousCommitMessages(), commitMessagePreviousCommitLimit)
if len(previous) > 0 {
sections = append(sections, "Recent commit messages:\n"+strings.Join(previous, "\n"))
}
if contextJSON, err := marshalCommitMessageExplicitContext(req.GetExplicitContext()); err != nil {
return nil, err
} else if contextJSON != "" {
sections = append(sections, "Explicit context:\n"+contextJSON)
}
diffSections := make([]string, 0, len(diffs))
for index, diff := range diffs {
diffSections = append(diffSections, fmt.Sprintf("--- Diff %d ---\n%s", index+1, diff))
}
sections = append(sections, "Diffs:\n"+strings.Join(diffSections, "\n\n"))
return []modeladapter.Message{
{Role: "system", Content: system},
{Role: "user", Content: strings.Join(sections, "\n\n")},
}, nil
}
func (service *Service) commitMessageHistoryRoot() string {
if service == nil || service.store == nil {
return ""
}
return strings.TrimSpace(service.store.HistoryDir())
}
func (service *Service) resolveCommitMessageModelID(ctx context.Context) (string, string, string) {
if service == nil || service.modelMemory == nil || service.resolver == nil {
return "", "default_fallback", ""
}
hash := strings.TrimSpace(service.modelMemory.LastAgentModelHash())
if hash == "" {
return "", "default_fallback", ""
}
channel, err := service.resolver.SelectChannelForModel(ctx, hash)
if err != nil || channel == nil || strings.TrimSpace(channel.ID) != hash {
if err != nil {
log.Printf("forwarder commit message ignored invalid last agent model hash=%s error=%v", hash, err)
}
return "", "default_fallback", hash
}
return hash, "last_agent_model_hash", hash
}
func recordCommitMessageIncomingRequest(recorder *commitMessageLogRecorder, requestID string, request *aiserverv1.WriteGitCommitMessageRequest) error {
_, err := recorder.appendEvent("incoming_request", map[string]any{
"request_id": requestID,
})
return err
}
func commitMessageConnectError(recorder *commitMessageLogRecorder, code connect.Code, err error) error {
if err == nil {
err = fmt.Errorf("commit message generation failed")
}
if _, logErr := recorder.appendEvent("error", map[string]any{
"code": code.String(),
"error": err.Error(),
}); logErr != nil {
log.Printf("forwarder failed to write commit message error log error=%v original_error=%v", logErr, err)
}
return connect.NewError(code, err)
}
func truncateCommitMessageDiffs(input []string, totalLimit int, singleLimit int) []string {
result := make([]string, 0, len(input))
remaining := totalLimit
for _, raw := range input {
diff := strings.TrimSpace(raw)
if diff == "" || remaining <= 0 {
continue
}
limit := singleLimit
if remaining < limit {
limit = remaining
}
truncated := truncateCommitMessageText(diff, limit)
if truncated == "" {
continue
}
result = append(result, truncated)
remaining -= len([]rune(truncated))
}
return result
}
func truncateCommitMessagePreviousCommits(input []string, limit int) []string {
if limit <= 0 {
return nil
}
result := make([]string, 0, limit)
for _, raw := range input {
value := strings.TrimSpace(raw)
if value == "" {
continue
}
result = append(result, "- "+value)
if len(result) >= limit {
break
}
}
return result
}
func marshalCommitMessageExplicitContext(contextValue *aiserverv1.ExplicitContext) (string, error) {
if contextValue == nil {
return "", nil
}
payload, err := protojson.MarshalOptions{UseProtoNames: true}.Marshal(contextValue)
if err != nil {
return "", fmt.Errorf("marshal explicit context failed: %w", err)
}
return truncateCommitMessageText(string(payload), commitMessageExplicitContextLimit), nil
}
func truncateCommitMessageText(value string, limit int) string {
trimmed := strings.TrimSpace(value)
if trimmed == "" || limit <= 0 {
return ""
}
runes := []rune(trimmed)
if len(runes) <= limit {
return trimmed
}
return strings.TrimSpace(string(runes[:limit])) + "\n...[truncated]"
}
func cleanGeneratedCommitMessage(value string) string {
result := strings.TrimSpace(value)
result = stripCommitMessageCodeFence(result)
prefixes := []string{
"commit message:",
"git commit message:",
"message:",
}
for {
lower := strings.ToLower(strings.TrimSpace(result))
matched := ""
for _, prefix := range prefixes {
if strings.HasPrefix(lower, prefix) {
matched = prefix
break
}
}
if matched == "" {
break
}
result = strings.TrimSpace(result[len(matched):])
}
return strings.TrimSpace(result)
}
func stripCommitMessageCodeFence(value string) string {
trimmed := strings.TrimSpace(value)
if !strings.HasPrefix(trimmed, "```") {
return trimmed
}
lines := strings.Split(trimmed, "\n")
if len(lines) < 2 {
return trimmed
}
lines = lines[1:]
if len(lines) > 0 && strings.HasPrefix(strings.TrimSpace(lines[len(lines)-1]), "```") {
lines = lines[:len(lines)-1]
}
return strings.TrimSpace(strings.Join(lines, "\n"))
}
func firstCommitMessageLine(value string) string {
for _, line := range strings.Split(value, "\n") {
line = strings.TrimSpace(line)
if line != "" {
return line
}
}
return ""
}
@@ -0,0 +1,23 @@
package forwarder
type commitMessageLogRecorder struct{}
func newCommitMessageLogRecorder(_ string, _ string) (*commitMessageLogRecorder, error) {
return &commitMessageLogRecorder{}, nil
}
func (recorder *commitMessageLogRecorder) RecordLLMRequest(_ string, _ string, _ string, _ map[string]any) (string, error) {
return "", nil
}
func (recorder *commitMessageLogRecorder) AppendLLMResponseChunk(_ string, _ string, _ string, _ string) (string, error) {
return "", nil
}
func (recorder *commitMessageLogRecorder) RecordLLMSummary(_ string, _ string, _ string, _ map[string]any) (string, error) {
return "", nil
}
func (recorder *commitMessageLogRecorder) appendEvent(_ string, _ map[string]any) (string, error) {
return "", nil
}
File diff suppressed because it is too large Load Diff
+255
View File
@@ -0,0 +1,255 @@
// compiler.go 负责把固定 prompt、自然历史和 tool catalog 编译成 provider 请求。
package forwarder
import (
"fmt"
"strings"
"cursor/gen/agentv1"
modeladapter "cursor/internal/backend/agent/model"
promptassets "cursor/prompt"
)
type PromptCompiler interface {
Compile(conversation *ConversationFile, mode agentv1.AgentMode, latestUserText string, modelName string) (CompiledConversation, error)
DerivePromptContexts(conversation *ConversationFile, mode agentv1.AgentMode, latestUserText string) ([]PromptContextMessage, error)
}
type DefaultPromptCompiler struct {
projector *HistoryProjector
catalog ToolCatalog
reminders ReminderInjector
rules *UserRuleStore
}
// NewPromptCompiler 创建默认 prompt 编译器。
func NewPromptCompiler(projector *HistoryProjector, catalog ToolCatalog, reminders ReminderInjector, rules *UserRuleStore) *DefaultPromptCompiler {
return &DefaultPromptCompiler{
projector: projector,
catalog: catalog,
reminders: reminders,
rules: rules,
}
}
// Compile 生成当前 turn 应发送给 provider 的消息和工具集合。
func (compiler *DefaultPromptCompiler) Compile(conversation *ConversationFile, mode agentv1.AgentMode, latestUserText string, modelName string) (CompiledConversation, error) {
if compiler == nil || compiler.projector == nil || compiler.catalog == nil {
return CompiledConversation{}, fmt.Errorf("prompt compiler dependencies are not initialized")
}
normalizedMode, err := validateSupportedActiveMode(mode)
if err != nil {
return CompiledConversation{}, err
}
subagentTypeName := ""
if conversation != nil {
subagentTypeName = conversation.SubagentTypeName
}
assetMode, err := promptAssetModeForConversation(normalizedMode, subagentTypeName)
if err != nil {
return CompiledConversation{}, err
}
systemPrompt, err := promptassets.ReadPrompt(assetMode)
if err != nil {
return CompiledConversation{}, err
}
tools, _, err := compiler.catalog.Load(normalizedMode, subagentTypeName)
if err != nil {
return CompiledConversation{}, err
}
replayMessages, err := compiler.projector.ProjectPromptReplay(conversation)
if err != nil {
return CompiledConversation{}, err
}
sharedRulesPrompt := ""
sharedRuleCount := 0
sharedRuleTotal := 0
if compiler.rules != nil && normalizedMode != agentv1.AgentMode_AGENT_MODE_DEBUG {
sharedRulesPrompt, sharedRuleTotal, sharedRuleCount, err = compiler.rules.BuildSystemPromptSection()
if err != nil {
return CompiledConversation{}, err
}
}
messages := make([]modeladapter.Message, 0, len(replayMessages)+1)
systemParts := []string{sanitizePromptAsset(systemPrompt, modelName)}
if strings.TrimSpace(sharedRulesPrompt) != "" {
systemParts = append(systemParts, sharedRulesPrompt)
}
systemText := strings.TrimSpace(strings.Join(filterNonEmpty(systemParts), "\n\n"))
if systemText != "" {
messages = append(messages, modeladapter.Message{
Role: "system",
Content: systemText,
})
}
stableReplayCount, err := compiler.stableReplayMessageCount(conversation, replayMessages)
if err != nil {
return CompiledConversation{}, err
}
messages = append(messages, replayMessages...)
return CompiledConversation{
Mode: normalizedMode,
Messages: messages,
StableMessageCount: stableReplayCount,
Tools: tools,
CompileSummary: fmt.Sprintf("mode=%s asset_mode=%s child=%t messages=%d tools=%d shared_rules_total=%d shared_rules_deduped=%d", normalizedMode.String(), string(assetMode), isChildConversationSubagentTypeName(subagentTypeName), len(messages), len(tools), sharedRuleTotal, sharedRuleCount),
}, nil
}
func (compiler *DefaultPromptCompiler) DerivePromptContexts(conversation *ConversationFile, mode agentv1.AgentMode, latestUserText string) ([]PromptContextMessage, error) {
if compiler == nil || compiler.projector == nil || compiler.catalog == nil || compiler.reminders == nil {
return nil, fmt.Errorf("prompt compiler dependencies are not initialized")
}
normalizedMode, err := validateSupportedActiveMode(mode)
if err != nil {
return nil, err
}
subagentTypeName := ""
if conversation != nil {
subagentTypeName = conversation.SubagentTypeName
}
_, toolNames, err := compiler.catalog.Load(normalizedMode, subagentTypeName)
if err != nil {
return nil, err
}
replayMessages, err := compiler.projector.ProjectPromptReplay(conversation)
if err != nil {
return nil, err
}
structuredStatePromptContexts, structuredStateTailMessages, err := buildStructuredStatePromptContexts(conversation)
if err != nil {
return nil, err
}
promptReminders := compiler.reminders.Inject(normalizedMode, conversation, replayMessages, latestUserText, toolNames)
candidates := make([]PromptContextMessage, 0, len(structuredStatePromptContexts)+len(structuredStateTailMessages)+len(promptReminders.PromptContexts)+len(promptReminders.TailMessages))
candidates = append(candidates, structuredStatePromptContexts...)
for _, message := range structuredStateTailMessages {
candidates = append(candidates, newPromptContextMessage(promptContextSourceStructuredTodoReminder, message, true))
}
candidates = append(candidates, promptReminders.PromptContexts...)
for index, message := range promptReminders.TailMessages {
candidates = append(candidates, newPromptContextMessage(fmt.Sprintf("tail_reminder/%d", index), message, true))
}
for index := range candidates {
candidates[index].Persist = true
}
return filterCurrentTurnPromptContexts(conversation, candidates), nil
}
func (compiler *DefaultPromptCompiler) stableReplayMessageCount(conversation *ConversationFile, replayMessages []modeladapter.Message) (int, error) {
if compiler == nil || compiler.projector == nil || conversation == nil || len(replayMessages) == 0 {
return 0, nil
}
currentTurnSeq := conversation.CurrentTurnSeq
if currentTurnSeq <= 0 {
currentTurnSeq = conversation.NextTurnSeq - 1
}
stableCount := 0
if currentTurnSeq > 0 {
stableConversation := cloneConversationFile(conversation)
stableConversation.Entries = stableReplayEntriesBeforeTurn(conversation.Entries, currentTurnSeq)
stableMessages, err := compiler.projector.ProjectPromptReplay(stableConversation)
if err != nil {
return 0, err
}
stableCount = len(stableMessages)
}
if requestPrefixReplayCount := replayMessageCountFromRequestPrefix(conversation); requestPrefixReplayCount > stableCount {
stableCount = requestPrefixReplayCount
}
if stableCount > len(replayMessages) {
return len(replayMessages), nil
}
return stableCount, nil
}
func replayMessageCountFromRequestPrefix(conversation *ConversationFile) int {
if conversation == nil || conversation.LatestRequestPrefix == nil {
return 0
}
requestID := strings.TrimSpace(conversation.CurrentRequestID)
if requestID == "" || strings.TrimSpace(conversation.LatestRequestPrefix.RequestID) != requestID {
return 0
}
return conversation.LatestRequestPrefix.ReplayMessageCount
}
func stableReplayEntriesBeforeTurn(entries []HistoryEntry, currentTurnSeq int64) []HistoryEntry {
if len(entries) == 0 || currentTurnSeq <= 0 {
return nil
}
filtered := make([]HistoryEntry, 0, len(entries))
for _, entry := range entries {
if entry.TurnSeq > 0 && entry.TurnSeq >= currentTurnSeq {
continue
}
filtered = append(filtered, entry)
}
return filtered
}
// filterNonEmpty 过滤掉空白字符串,便于安全拼接 system prompt 片段。
func filterNonEmpty(items []string) []string {
filtered := make([]string, 0, len(items))
for _, item := range items {
if strings.TrimSpace(item) != "" {
filtered = append(filtered, strings.TrimSpace(item))
}
}
return filtered
}
func mergeAdjacentPlainUserMessages(messages []modeladapter.Message) []modeladapter.Message {
if len(messages) == 0 {
return nil
}
merged := make([]modeladapter.Message, 0, len(messages))
for _, message := range messages {
text, ok := plainUserMessageText(message)
if !ok {
merged = append(merged, message)
continue
}
if len(merged) == 0 {
message.Role = "user"
message.Content = text
merged = append(merged, message)
continue
}
last := &merged[len(merged)-1]
if lastText, ok := plainUserMessageText(*last); ok {
last.Role = "user"
last.Content = lastText + "\n\n" + text
continue
}
message.Role = "user"
message.Content = text
merged = append(merged, message)
}
return merged
}
func plainUserMessageText(message modeladapter.Message) (string, bool) {
if strings.TrimSpace(message.Role) != "user" {
return "", false
}
if len(message.ContentParts) > 0 || len(message.ToolCalls) > 0 {
return "", false
}
if strings.TrimSpace(message.ToolCallID) != "" || strings.TrimSpace(message.Name) != "" {
return "", false
}
if strings.TrimSpace(message.ReasoningContent) != "" ||
strings.TrimSpace(message.ReasoningSignature) != "" ||
strings.TrimSpace(message.ReasoningSignatureSource) != "" ||
strings.TrimSpace(message.OpenAIResponsesReasoningID) != "" ||
strings.TrimSpace(message.OpenAIResponsesReasoningStatus) != "" ||
len(message.OpenAIResponsesReasoningSummary) > 0 {
return "", false
}
text := strings.TrimSpace(message.Content)
if text == "" {
return "", false
}
return text, true
}
@@ -0,0 +1,36 @@
package forwarder
import (
"encoding/json"
"fmt"
"strings"
runtimecore "cursor/internal/backend/agent/core"
)
func (service *Service) sanitizeCreatePlanInvocationForCurrentPlan(stream *ActiveStream, invocation runtimecore.ToolInvocation) (runtimecore.ToolInvocation, error) {
if strings.TrimSpace(invocation.ToolName) != "CreatePlan" {
return invocation, nil
}
args, err := runtimecore.DecodeCreatePlanArgsJSON(invocation.ArgsJSON)
if err != nil {
return invocation, newRecoverableToolInvocationError(fmt.Errorf("decode CreatePlan args failed: %w", err))
}
if strings.TrimSpace(args.GetName()) == "" {
return invocation, nil
}
conversation, _, _, err := service.snapshotCheckpointConversation(stream)
if err != nil {
return invocation, err
}
if !hasCurrentPlan(conversation) {
return invocation, nil
}
args.Name = ""
sanitized, err := json.Marshal(args)
if err != nil {
return invocation, newRecoverableToolInvocationError(fmt.Errorf("encode sanitized CreatePlan args failed: %w", err))
}
invocation.ArgsJSON = sanitized
return invocation, nil
}
@@ -0,0 +1,394 @@
package forwarder
import (
"context"
"encoding/json"
"fmt"
"os"
"path/filepath"
"strings"
"sync"
"time"
"google.golang.org/protobuf/encoding/protojson"
"google.golang.org/protobuf/proto"
"cursor/gen/agentv1"
)
type debugLogConfig interface {
IsObservabilityLogEnabled(context.Context) bool
}
type debugRecorder struct {
historyRoot string
broker *StreamBroker
config debugLogConfig
mu sync.Mutex
}
func newDebugRecorder(historyRoot string, broker *StreamBroker, config debugLogConfig) *debugRecorder {
return &debugRecorder{
historyRoot: strings.TrimSpace(historyRoot),
broker: broker,
config: config,
}
}
func (recorder *debugRecorder) enabled(ctx context.Context) bool {
if recorder == nil || recorder.config == nil {
return false
}
if ctx == nil {
ctx = context.Background()
}
return recorder.config.IsObservabilityLogEnabled(ctx)
}
func (recorder *debugRecorder) LogBidiRaw(ctx context.Context, requestID string, conversationID string, appendSeqno int64, dataHex string, status string, extra map[string]any) {
if !recorder.enabled(ctx) {
return
}
event := recorder.baseEvent("bidi_raw", requestID, conversationID)
event["direction"] = "client_to_backend"
event["procedure"] = "/aiserver.v1.BidiService/BidiAppend"
event["append_seqno"] = appendSeqno
event["status"] = strings.TrimSpace(status)
event["data_hex"] = dataHex
event["data_len"] = len(dataHex)
for key, value := range extra {
event[key] = value
}
recorder.appendJSONL(ctx, requestID, conversationID, "bidi.raw.jsonl", event)
}
func (recorder *debugRecorder) LogBidiDecoded(ctx context.Context, requestID string, conversationID string, appendSeqno int64, clientKind string, message *agentv1.AgentClientMessage, intent InboundIntent, extra map[string]any) {
if !recorder.enabled(ctx) {
return
}
event := recorder.baseEvent("bidi_decoded", requestID, conversationID)
event["schema_version"] = 2
event["append_seqno"] = appendSeqno
event["client_kind"] = strings.TrimSpace(clientKind)
event["message_case"] = agentClientMessageCase(message)
event["message"] = protoJSONDebugPayload(message)
event["intent"] = inboundIntentDebugPayload(intent)
if requestedModel := requestedModelDebugPayload(message); requestedModel != nil {
event["requested_model"] = requestedModel
}
if actionCase := conversationActionCase(message); actionCase != "" {
event["conversation_action"] = actionCase
}
for key, value := range extra {
event[key] = value
}
recorder.appendJSONL(ctx, requestID, firstNonEmpty(conversationID, intent.ConversationID), "bidi.decoded.jsonl", event)
}
func (recorder *debugRecorder) LogRuntime(ctx context.Context, requestID string, conversationID string, eventName string, fields map[string]any) {
if !recorder.enabled(ctx) {
return
}
event := recorder.baseEvent("runtime", requestID, conversationID)
event["event"] = strings.TrimSpace(eventName)
for key, value := range fields {
event[key] = value
}
recorder.appendJSONL(ctx, requestID, conversationID, "runtime.jsonl", event)
}
func (recorder *debugRecorder) LogRunSSE(ctx context.Context, requestID string, conversationID string, eventName string, fields map[string]any) {
if !recorder.enabled(ctx) {
return
}
event := recorder.baseEvent("runsse", requestID, conversationID)
event["event"] = strings.TrimSpace(eventName)
for key, value := range fields {
event[key] = value
}
recorder.appendJSONL(ctx, requestID, conversationID, "runsse.jsonl", event)
}
func (recorder *debugRecorder) LogProvider(ctx context.Context, requestID string, conversationID string, eventName string, fields map[string]any) {
if !recorder.enabled(ctx) {
return
}
event := recorder.baseEvent("provider", requestID, conversationID)
event["event"] = strings.TrimSpace(eventName)
for key, value := range fields {
event[key] = value
}
recorder.appendJSONL(ctx, requestID, conversationID, "provider.jsonl", event)
}
func (recorder *debugRecorder) LogProviderArtifact(ctx context.Context, requestID string, conversationID string, modelCallID string, eventName string, payload map[string]any) {
if !recorder.enabled(ctx) {
return
}
recorder.LogProvider(ctx, requestID, conversationID, eventName, map[string]any{
"model_call_id": strings.TrimSpace(modelCallID),
"payload": payload,
})
}
func (recorder *debugRecorder) baseEvent(layer string, requestID string, conversationID string) map[string]any {
resolvedConversationID := firstNonEmpty(strings.TrimSpace(conversationID), recorder.conversationIDForRequest(requestID))
return map[string]any{
"schema_version": 1,
"at": time.Now().UTC().Format(time.RFC3339Nano),
"layer": strings.TrimSpace(layer),
"request_id": strings.TrimSpace(requestID),
"conversation_id": resolvedConversationID,
}
}
func (recorder *debugRecorder) appendJSONL(ctx context.Context, requestID string, conversationID string, filename string, event map[string]any) {
if !recorder.enabled(ctx) || len(event) == 0 {
return
}
dir := recorder.debugDir(requestID, conversationID)
if strings.TrimSpace(dir) == "" {
return
}
payload, err := json.Marshal(event)
if err != nil {
return
}
recorder.mu.Lock()
defer recorder.mu.Unlock()
if err := os.MkdirAll(dir, 0o755); err != nil {
return
}
file, err := os.OpenFile(filepath.Join(dir, filename), os.O_CREATE|os.O_WRONLY|os.O_APPEND, 0o644)
if err != nil {
return
}
defer file.Close()
_, _ = file.Write(append(payload, '\n'))
}
func (recorder *debugRecorder) debugDir(requestID string, conversationID string) string {
if recorder == nil || strings.TrimSpace(recorder.historyRoot) == "" {
return ""
}
conversationID = firstNonEmpty(strings.TrimSpace(conversationID), recorder.conversationIDForRequest(requestID))
if conversationID != "" && conversationID != "unknown" {
return filepath.Join(recorder.historyRoot, sanitizeArtifactName(conversationID), "debug")
}
requestID = firstNonEmpty(strings.TrimSpace(requestID), "unknown")
return filepath.Join(recorder.historyRoot, "_debug", "orphan", sanitizeArtifactName(requestID))
}
func (recorder *debugRecorder) conversationIDForRequest(requestID string) string {
if recorder == nil || recorder.broker == nil {
return ""
}
stream, ok := recorder.broker.Get(requestID)
if !ok || stream == nil {
return ""
}
stream.mu.Lock()
defer stream.mu.Unlock()
return strings.TrimSpace(stream.ConversationID)
}
func agentClientMessageCase(message *agentv1.AgentClientMessage) string {
if message == nil {
return ""
}
switch message.GetMessage().(type) {
case *agentv1.AgentClientMessage_RunRequest:
return "run_request"
case *agentv1.AgentClientMessage_PrewarmRequest:
return "prewarm_request"
case *agentv1.AgentClientMessage_ConversationAction:
return "conversation_action"
case *agentv1.AgentClientMessage_ExecClientMessage:
return "exec_client_message"
case *agentv1.AgentClientMessage_ExecClientControlMessage:
return "exec_client_control_message"
case *agentv1.AgentClientMessage_InteractionResponse:
return "interaction_response"
case *agentv1.AgentClientMessage_ClientHeartbeat:
return "client_heartbeat"
case *agentv1.AgentClientMessage_KvClientMessage:
return "kv_client_message"
default:
return fmt.Sprintf("%T", message.GetMessage())
}
}
func agentServerMessageCase(message *agentv1.AgentServerMessage) string {
if message == nil {
return ""
}
switch message.GetMessage().(type) {
case *agentv1.AgentServerMessage_InteractionUpdate:
return "interaction_update"
case *agentv1.AgentServerMessage_ExecServerMessage:
return "exec_server_message"
case *agentv1.AgentServerMessage_ExecServerControlMessage:
return "exec_server_control_message"
case *agentv1.AgentServerMessage_ConversationCheckpointUpdate:
return "conversation_checkpoint_update"
case *agentv1.AgentServerMessage_KvServerMessage:
return "kv_server_message"
case *agentv1.AgentServerMessage_InteractionQuery:
return "interaction_query"
default:
return fmt.Sprintf("%T", message.GetMessage())
}
}
func conversationActionCase(message *agentv1.AgentClientMessage) string {
if message == nil {
return ""
}
action := message.GetConversationAction()
if action == nil && message.GetRunRequest() != nil {
action = message.GetRunRequest().GetAction()
}
if action == nil {
return ""
}
return conversationActionKind(action)
}
func requestedModelDebugPayload(message *agentv1.AgentClientMessage) map[string]any {
if message == nil {
return nil
}
if runRequest := message.GetRunRequest(); runRequest != nil {
return requestedModelPayload(runRequest.GetRequestedModel())
}
if prewarm := message.GetPrewarmRequest(); prewarm != nil {
return requestedModelPayload(prewarm.GetRequestedModel())
}
return nil
}
func requestedModelPayload(model *agentv1.RequestedModel) map[string]any {
if model == nil {
return nil
}
parameters := make([]map[string]string, 0, len(model.GetParameters()))
for _, parameter := range model.GetParameters() {
if parameter == nil {
continue
}
parameters = append(parameters, map[string]string{
"id": parameter.GetId(),
"value": parameter.GetValue(),
})
}
return map[string]any{
"model_id": strings.TrimSpace(model.GetModelId()),
"max_mode": model.GetMaxMode(),
"built_in_model": model.GetBuiltInModel(),
"is_variant_string_representation": model.GetIsVariantStringRepresentation(),
"parameters": parameters,
}
}
func protoJSONDebugPayload(message proto.Message) any {
if message == nil {
return nil
}
payload, err := protojson.MarshalOptions{
UseProtoNames: true,
EmitUnpopulated: false,
}.Marshal(message)
if err != nil {
return map[string]any{"marshal_error": err.Error()}
}
var decoded any
if err := json.Unmarshal(payload, &decoded); err != nil {
return string(payload)
}
return decoded
}
func inboundIntentDebugPayload(intent InboundIntent) map[string]any {
payload := map[string]any{
"kind": strings.TrimSpace(intent.Kind),
"request_id": strings.TrimSpace(intent.RequestID),
"conversation_id": strings.TrimSpace(intent.ConversationID),
"model_id": strings.TrimSpace(intent.ModelID),
"model_name": strings.TrimSpace(intent.ModelName),
"thinking_effort": strings.TrimSpace(intent.ThinkingEffort),
"mode": intent.Mode.String(),
"has_explicit_mode": intent.HasExplicitMode,
"mode_source": string(intent.ModeSource),
"starts_run": intent.StartsRun,
"subagent_type_name": strings.TrimSpace(intent.SubagentTypeName),
"cancel_reason": strings.TrimSpace(intent.CancelReason),
"prewarm": intent.Prewarm,
}
if len(intent.SubagentModelOverrides) > 0 {
payload["subagent_model_overrides"] = subagentModelOverrideSummaries(intent.SubagentModelOverrides)
payload["subagent_model_override_count"] = len(intent.SubagentModelOverrides)
}
if intent.ClientMessage != nil {
payload["client_message"] = protoJSONDebugPayload(intent.ClientMessage)
}
if intent.ConversationState != nil {
payload["conversation_state"] = protoJSONDebugPayload(intent.ConversationState)
}
if intent.UserMessage != nil {
payload["user_message"] = protoJSONDebugPayload(intent.UserMessage)
}
if intent.RequestContext != nil {
payload["request_context"] = protoJSONDebugPayload(intent.RequestContext)
}
if strings.TrimSpace(intent.IgnoredReason) != "" {
payload["ignored_reason"] = strings.TrimSpace(intent.IgnoredReason)
payload["ignored_empty_resume"] = strings.TrimSpace(intent.IgnoredReason) == "empty_resume_without_pending_continuation"
}
if intent.ExecClientMessage != nil {
payload["exec_client_message"] = protoJSONDebugPayload(intent.ExecClientMessage)
}
if intent.ExecClientControlMessage != nil {
payload["exec_client_control_message"] = protoJSONDebugPayload(intent.ExecClientControlMessage)
}
if intent.InteractionResponse != nil {
payload["interaction_response"] = protoJSONDebugPayload(intent.InteractionResponse)
}
if intent.KVClientMessage != nil {
payload["kv_client_message"] = protoJSONDebugPayload(intent.KVClientMessage)
}
return payload
}
func conversationActionKind(action *agentv1.ConversationAction) string {
if action == nil {
return ""
}
switch action.GetAction().(type) {
case *agentv1.ConversationAction_UserMessageAction:
return "user_message_action"
case *agentv1.ConversationAction_ResumeAction:
return "resume_action"
case *agentv1.ConversationAction_CancelAction:
return "cancel_action"
case *agentv1.ConversationAction_SummarizeAction:
return "summarize_action"
case *agentv1.ConversationAction_ShellCommandAction:
return "shell_command_action"
case *agentv1.ConversationAction_StartPlanAction:
return "start_plan_action"
case *agentv1.ConversationAction_ExecutePlanAction:
return "execute_plan_action"
case *agentv1.ConversationAction_AsyncAskQuestionCompletionAction:
return "async_ask_question_completion_action"
case *agentv1.ConversationAction_CancelSubagentAction:
return "cancel_subagent_action"
case *agentv1.ConversationAction_BackgroundTaskCompletionAction:
return "background_task_completion_action"
case *agentv1.ConversationAction_BackgroundShellAction:
return "background_shell_action"
case *agentv1.ConversationAction_BackgroundSubagentAction:
return "background_subagent_action"
default:
return fmt.Sprintf("%T", action.GetAction())
}
}
@@ -0,0 +1,210 @@
package forwarder
import (
"context"
"fmt"
"strings"
"connectrpc.com/connect"
"cursor/gen/aiserverv1"
)
func (service *Service) AvailableDocs(_ context.Context, req *connect.Request[aiserverv1.AvailableDocsRequest]) (*connect.Response[aiserverv1.AvailableDocsResponse], error) {
store, err := service.requireDocsIndexStore()
if err != nil {
return nil, connect.NewError(connect.CodeInternal, err)
}
records, err := store.List("", 0)
if err != nil {
return nil, connect.NewError(connect.CodeInternal, err)
}
records = filterDocsIndexRecords(records, req.Msg)
records, err = appendMissingAdditionalDocs(store, records, req.Msg)
if err != nil {
return nil, connect.NewError(connect.CodeInternal, err)
}
docs := make([]*aiserverv1.DocumentationInfo, 0, len(records))
for _, record := range records {
docs = append(docs, docsIndexDocumentationInfo(record))
}
return connect.NewResponse(&aiserverv1.AvailableDocsResponse{Docs: docs}), nil
}
func (service *Service) DocumentationQuery(_ context.Context, req *connect.Request[aiserverv1.DocumentationQueryRequest]) (*connect.Response[aiserverv1.DocumentationQueryResponse], error) {
store, err := service.requireDocsIndexStore()
if err != nil {
return nil, connect.NewError(connect.CodeInternal, err)
}
identifier := strings.TrimSpace(req.Msg.GetDocIdentifier())
if identifier == "" {
return connect.NewResponse(&aiserverv1.DocumentationQueryResponse{
Status: aiserverv1.DocumentationQueryResponse_STATUS_NOT_FOUND,
}), nil
}
record, ok, err := store.Get(identifier)
if err != nil {
return nil, connect.NewError(connect.CodeInternal, err)
}
if !ok {
return connect.NewResponse(&aiserverv1.DocumentationQueryResponse{
DocIdentifier: identifier,
Status: aiserverv1.DocumentationQueryResponse_STATUS_NOT_FOUND,
}), nil
}
return connect.NewResponse(&aiserverv1.DocumentationQueryResponse{
DocIdentifier: record.Identifier,
DocName: record.Title,
DocChunks: docsIndexChunks(record, req.Msg.GetTopK()),
Status: aiserverv1.DocumentationQueryResponse_STATUS_SUCCESS,
}), nil
}
func (service *Service) FetchRelevantKnowledgeForConversation(_ context.Context, req *connect.Request[aiserverv1.FetchRelevantKnowledgeForConversationRequest]) (*connect.Response[aiserverv1.FetchRelevantKnowledgeForConversationResponse], error) {
store, err := service.requireDocsIndexStore()
if err != nil {
return nil, connect.NewError(connect.CodeInternal, err)
}
gitOrigin := strings.TrimSpace(req.Msg.GetGitOrigin())
if gitOrigin == "" {
return connect.NewResponse(&aiserverv1.FetchRelevantKnowledgeForConversationResponse{}), nil
}
records, err := store.List(gitOrigin, req.Msg.GetLimit())
if err != nil {
return nil, connect.NewError(connect.CodeInternal, err)
}
items := make([]*aiserverv1.ConversationMessage_KnowledgeItem, 0, len(records))
for _, record := range records {
items = append(items, &aiserverv1.ConversationMessage_KnowledgeItem{
Title: record.Title,
Knowledge: docsIndexKnowledgeText(record),
KnowledgeId: record.Identifier,
IsGenerated: false,
})
}
return connect.NewResponse(&aiserverv1.FetchRelevantKnowledgeForConversationResponse{KnowledgeItems: items}), nil
}
func (service *Service) requireDocsIndexStore() (*DocsIndexStore, error) {
if service == nil || service.docsIndexStore == nil {
return nil, fmt.Errorf("docs index store is not initialized")
}
return service.docsIndexStore, nil
}
func filterDocsIndexRecords(records []DocsIndexRecord, req *aiserverv1.AvailableDocsRequest) []DocsIndexRecord {
if req == nil || req.GetGetAll() || (req.GetPartialUrl() == "" && req.GetPartialDocName() == "" && len(req.GetAdditionalDocIdentifiers()) == 0) {
return records
}
partialURL := strings.ToLower(strings.TrimSpace(req.GetPartialUrl()))
partialName := strings.ToLower(strings.TrimSpace(req.GetPartialDocName()))
allowed := make(map[string]struct{})
for _, identifier := range req.GetAdditionalDocIdentifiers() {
trimmed := strings.TrimSpace(identifier)
if trimmed != "" {
allowed[trimmed] = struct{}{}
}
}
filtered := make([]DocsIndexRecord, 0, len(records))
for _, record := range records {
_, explicitlyAllowed := allowed[record.Identifier]
matchesURL := partialURL != "" && strings.Contains(strings.ToLower(record.URL), partialURL)
matchesName := partialName != "" && strings.Contains(strings.ToLower(record.Title), partialName)
if explicitlyAllowed || matchesURL || matchesName {
filtered = append(filtered, record)
}
}
return filtered
}
func appendMissingAdditionalDocs(store *DocsIndexStore, records []DocsIndexRecord, req *aiserverv1.AvailableDocsRequest) ([]DocsIndexRecord, error) {
if req == nil {
return records, nil
}
seen := make(map[string]struct{}, len(records))
for _, record := range records {
seen[record.Identifier] = struct{}{}
}
for _, identifier := range req.GetAdditionalDocIdentifiers() {
trimmed := strings.TrimSpace(identifier)
if trimmed == "" {
continue
}
if _, ok := seen[trimmed]; ok {
continue
}
record := DocsIndexRecord{
ID: trimmed,
Identifier: trimmed,
Title: trimmed,
Status: docsIndexStatusIndexed,
Source: docsIndexSourceAdditional,
}
if store != nil {
persisted, err := store.Upsert(record)
if err != nil {
return nil, err
}
record = persisted
}
records = append(records, record)
seen[trimmed] = struct{}{}
}
return records, nil
}
func docsIndexDocumentationInfo(record DocsIndexRecord) *aiserverv1.DocumentationInfo {
return &aiserverv1.DocumentationInfo{
DocIdentifier: record.Identifier,
Metadata: &aiserverv1.DocumentationMetadata{
PrefixUrl: record.URL,
DocName: record.Title,
TruePrefixUrl: record.URL,
Public: false,
},
}
}
func docsIndexChunks(record DocsIndexRecord, topK uint32) []*aiserverv1.DocumentationChunk {
content := strings.TrimSpace(record.Content)
if content == "" {
return nil
}
if topK == 0 {
topK = 1
}
return []*aiserverv1.DocumentationChunk{
{
DocName: record.Title,
PageUrl: record.URL,
DocumentationChunk: content,
Score: 1,
PageTitle: record.Title,
},
}
}
func docsIndexKnowledgeText(record DocsIndexRecord) string {
parts := make([]string, 0, 3)
if strings.TrimSpace(record.URL) != "" {
parts = append(parts, "URL: "+strings.TrimSpace(record.URL))
}
if strings.TrimSpace(record.Title) != "" {
parts = append(parts, "Title: "+strings.TrimSpace(record.Title))
}
if strings.TrimSpace(record.Content) != "" {
parts = append(parts, strings.TrimSpace(record.Content))
}
return strings.Join(parts, "\n")
}
func docsIndexURLCandidate(value string) string {
trimmed := strings.TrimSpace(value)
if strings.HasPrefix(trimmed, "http://") || strings.HasPrefix(trimmed, "https://") {
fields := strings.Fields(trimmed)
if len(fields) > 0 {
return fields[0]
}
}
return ""
}
@@ -0,0 +1,241 @@
package forwarder
import (
"crypto/sha256"
"encoding/hex"
"encoding/json"
"fmt"
"os"
"path/filepath"
"sort"
"strings"
"sync"
"time"
)
const (
docsIndexStatusIndexed = "indexed"
docsIndexSourceLocal = "local_docs"
docsIndexSourceAdditional = "additional_docs"
)
type DocsIndexStore struct {
mu sync.Mutex
root string
path string
loaded bool
state docsIndexState
}
type docsIndexState struct {
SchemaVersion int `json:"schema_version"`
Docs map[string]DocsIndexRecord `json:"docs"`
}
type DocsIndexRecord struct {
ID string `json:"id"`
Identifier string `json:"identifier"`
Title string `json:"title"`
URL string `json:"url,omitempty"`
Content string `json:"content,omitempty"`
GitOrigin string `json:"git_origin,omitempty"`
Status string `json:"status"`
Source string `json:"source"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
LastCrawlAt time.Time `json:"last_crawl_at,omitempty"`
LastIndexAt time.Time `json:"last_index_at,omitempty"`
}
func NewDocsIndexStore(root string) *DocsIndexStore {
trimmed := strings.TrimSpace(root)
if trimmed == "" {
trimmed = "docs-index"
}
return &DocsIndexStore{
root: trimmed,
path: filepath.Join(trimmed, "index.json"),
state: docsIndexState{
SchemaVersion: 1,
Docs: make(map[string]DocsIndexRecord),
},
}
}
func (store *DocsIndexStore) Upsert(record DocsIndexRecord) (DocsIndexRecord, error) {
if store == nil {
return DocsIndexRecord{}, fmt.Errorf("docs index store is nil")
}
store.mu.Lock()
defer store.mu.Unlock()
if err := store.loadLocked(); err != nil {
return DocsIndexRecord{}, err
}
now := time.Now().UTC()
identifier := strings.TrimSpace(record.Identifier)
if identifier == "" {
identifier = stableDocsIdentifier(record.Title, record.URL, record.Content, record.GitOrigin)
}
existing, ok := store.state.Docs[identifier]
if ok {
existing.Title = firstNonEmptyDocs(record.Title, existing.Title, identifier)
existing.URL = firstNonEmptyDocs(record.URL, existing.URL)
existing.Content = firstNonEmptyDocs(record.Content, existing.Content)
existing.GitOrigin = firstNonEmptyDocs(record.GitOrigin, existing.GitOrigin)
existing.Status = docsIndexStatusIndexed
existing.Source = firstNonEmptyDocs(record.Source, existing.Source, docsIndexSourceLocal)
existing.UpdatedAt = now
existing.LastCrawlAt = now
existing.LastIndexAt = now
store.state.Docs[identifier] = existing
return existing, store.saveLocked()
}
record.ID = firstNonEmptyDocs(record.ID, identifier)
record.Identifier = identifier
record.Title = firstNonEmptyDocs(record.Title, identifier)
record.Status = docsIndexStatusIndexed
record.Source = firstNonEmptyDocs(record.Source, docsIndexSourceLocal)
record.CreatedAt = now
record.UpdatedAt = now
record.LastCrawlAt = now
record.LastIndexAt = now
store.state.Docs[identifier] = record
return record, store.saveLocked()
}
func (store *DocsIndexStore) List(gitOrigin string, limit int32) ([]DocsIndexRecord, error) {
if store == nil {
return nil, fmt.Errorf("docs index store is nil")
}
store.mu.Lock()
defer store.mu.Unlock()
if err := store.loadLocked(); err != nil {
return nil, err
}
origin := strings.TrimSpace(gitOrigin)
records := make([]DocsIndexRecord, 0, len(store.state.Docs))
for _, record := range store.state.Docs {
if origin != "" && strings.TrimSpace(record.GitOrigin) != origin {
continue
}
records = append(records, record)
}
sort.SliceStable(records, func(i, j int) bool {
if records[i].UpdatedAt.Equal(records[j].UpdatedAt) {
return records[i].Identifier < records[j].Identifier
}
return records[i].UpdatedAt.After(records[j].UpdatedAt)
})
if limit > 0 && len(records) > int(limit) {
records = records[:limit]
}
return records, nil
}
func (store *DocsIndexStore) Get(identifier string) (DocsIndexRecord, bool, error) {
if store == nil {
return DocsIndexRecord{}, false, fmt.Errorf("docs index store is nil")
}
store.mu.Lock()
defer store.mu.Unlock()
if err := store.loadLocked(); err != nil {
return DocsIndexRecord{}, false, err
}
record, ok := store.state.Docs[strings.TrimSpace(identifier)]
return record, ok, nil
}
func (store *DocsIndexStore) Remove(identifier string) error {
if store == nil {
return fmt.Errorf("docs index store is nil")
}
store.mu.Lock()
defer store.mu.Unlock()
if err := store.loadLocked(); err != nil {
return err
}
delete(store.state.Docs, strings.TrimSpace(identifier))
return store.saveLocked()
}
func (store *DocsIndexStore) loadLocked() error {
if store.loaded {
return nil
}
state := docsIndexState{SchemaVersion: 1, Docs: make(map[string]DocsIndexRecord)}
data, err := os.ReadFile(store.path)
if err != nil {
if os.IsNotExist(err) {
store.state = state
store.loaded = true
return nil
}
return err
}
if len(data) != 0 {
if err := json.Unmarshal(data, &state); err != nil {
return err
}
}
if state.SchemaVersion <= 0 {
state.SchemaVersion = 1
}
if state.Docs == nil {
state.Docs = make(map[string]DocsIndexRecord)
}
store.state = state
store.loaded = true
return nil
}
func (store *DocsIndexStore) saveLocked() error {
if err := os.MkdirAll(store.root, 0o755); err != nil {
return err
}
data, err := json.MarshalIndent(store.state, "", " ")
if err != nil {
return err
}
tmp, err := os.CreateTemp(store.root, ".index-*.json.tmp")
if err != nil {
return err
}
tmpPath := tmp.Name()
defer func() { _ = os.Remove(tmpPath) }()
if _, err := tmp.Write(data); err != nil {
_ = tmp.Close()
return err
}
if err := tmp.Close(); err != nil {
return err
}
if err := os.Chmod(tmpPath, 0o600); err != nil {
return err
}
return os.Rename(tmpPath, store.path)
}
func stableDocsIdentifier(values ...string) string {
parts := make([]string, 0, len(values))
for _, value := range values {
trimmed := strings.TrimSpace(value)
if trimmed != "" {
parts = append(parts, trimmed)
}
}
if len(parts) == 0 {
return "doc_local"
}
sum := sha256.Sum256([]byte(strings.Join(parts, "\x00")))
return "doc_" + hex.EncodeToString(sum[:])[:24]
}
func firstNonEmptyDocs(values ...string) string {
for _, value := range values {
trimmed := strings.TrimSpace(value)
if trimmed != "" {
return trimmed
}
}
return ""
}
@@ -0,0 +1,491 @@
package forwarder
import (
"fmt"
"strings"
"unicode/utf8"
"github.com/sergi/go-diff/diffmatchpatch"
"google.golang.org/protobuf/encoding/protojson"
"cursor/gen/agentv1"
runtimecore "cursor/internal/backend/agent/core"
)
const editReadBinaryContentLimit = 32 * 1024
type editComputation struct {
BeforeContent string
AfterContent string
DiffString string
LinesAdded int32
LinesRemoved int32
Message string
}
type editResultPayload struct {
BeforeContent string
AfterContent string
DiffString string
LinesAdded int32
LinesRemoved int32
Message string
}
type editMatch struct {
Start int
End int
}
type normalizedEditContent struct {
Text string
OriginalOffsets []int
}
func computeEditDiff(beforeContent string, afterContent string) (string, int32, int32) {
dmp := diffmatchpatch.New()
charsBefore, charsAfter, lineArray := dmp.DiffLinesToChars(beforeContent, afterContent)
diffs := dmp.DiffMain(charsBefore, charsAfter, false)
diffs = dmp.DiffCharsToLines(diffs, lineArray)
linesAdded := int32(0)
linesRemoved := int32(0)
for _, diff := range diffs {
switch diff.Type {
case diffmatchpatch.DiffInsert:
linesAdded += countChangedLines(diff.Text)
case diffmatchpatch.DiffDelete:
linesRemoved += countChangedLines(diff.Text)
}
}
return dmp.PatchToText(dmp.PatchMake(beforeContent, diffs)), linesAdded, linesRemoved
}
func normalizeEditContentLineEndings(content string) normalizedEditContent {
var builder strings.Builder
builder.Grow(len(content))
offsets := make([]int, 1, len(content)+1)
offsets[0] = 0
for index := 0; index < len(content); {
if content[index] == '\r' && index+1 < len(content) && content[index+1] == '\n' {
builder.WriteByte('\n')
index += 2
offsets = append(offsets, index)
continue
}
if content[index] == '\r' {
builder.WriteByte('\n')
index++
offsets = append(offsets, index)
continue
}
builder.WriteByte(content[index])
index++
offsets = append(offsets, index)
}
return normalizedEditContent{
Text: builder.String(),
OriginalOffsets: offsets,
}
}
func normalizeLineEndingsToLF(value string) string {
if !strings.ContainsAny(value, "\r\n") {
return value
}
normalized := strings.ReplaceAll(value, "\r\n", "\n")
return strings.ReplaceAll(normalized, "\r", "\n")
}
func detectPreferredFileLineEnding(content string) string {
return dominantLineEnding(content, "\n")
}
func preferredMatchLineEnding(content string, match editMatch, defaultLineEnding string) string {
return dominantLineEnding(content[match.Start:match.End], defaultLineEnding)
}
func normalizeReplacementLineEndings(newString string, lineEnding string) string {
if !strings.ContainsAny(newString, "\r\n") {
return newString
}
normalized := normalizeLineEndingsToLF(newString)
switch lineEnding {
case "\r\n":
return strings.ReplaceAll(normalized, "\n", "\r\n")
case "\r":
return strings.ReplaceAll(normalized, "\n", "\r")
default:
return normalized
}
}
func dominantLineEnding(content string, fallback string) string {
if fallback == "" {
fallback = "\n"
}
counts := map[string]int{
"\n": 0,
"\r\n": 0,
"\r": 0,
}
first := ""
for index := 0; index < len(content); {
switch {
case content[index] == '\r' && index+1 < len(content) && content[index+1] == '\n':
counts["\r\n"]++
if first == "" {
first = "\r\n"
}
index += 2
case content[index] == '\r':
counts["\r"]++
if first == "" {
first = "\r"
}
index++
case content[index] == '\n':
counts["\n"]++
if first == "" {
first = "\n"
}
index++
default:
index++
}
}
best := ""
bestCount := 0
for _, candidate := range []string{"\r\n", "\n", "\r"} {
if counts[candidate] > bestCount {
best = candidate
bestCount = counts[candidate]
}
}
if bestCount == 0 {
return fallback
}
for _, candidate := range []string{"\r\n", "\n", "\r"} {
if counts[candidate] == bestCount && candidate == first {
return candidate
}
}
return best
}
func countChangedLines(text string) int32 {
if text == "" {
return 0
}
count := int32(strings.Count(text, "\n"))
if !strings.HasSuffix(text, "\n") {
count++
}
return count
}
func buildCompletedEditToolCall(path string, result *agentv1.EditResult) *agentv1.ToolCall {
args := &agentv1.EditArgs{Path: strings.TrimSpace(path)}
return &agentv1.ToolCall{
Tool: &agentv1.ToolCall_EditToolCall{
EditToolCall: &agentv1.EditToolCall{
Args: args,
Result: result,
},
},
}
}
func buildSuccessfulEditResult(path string, beforeContent string, afterContent string, diffString string, linesAdded int32, linesRemoved int32, message string) *agentv1.EditResult {
success := &agentv1.EditSuccess{
Path: strings.TrimSpace(path),
AfterFullFileContent: afterContent,
}
if beforeContent != "" {
success.BeforeFullFileContent = stringPtr(beforeContent)
}
if diffString != "" {
success.DiffString = stringPtr(diffString)
}
if linesAdded != 0 {
success.LinesAdded = int32Ptr(linesAdded)
}
if linesRemoved != 0 {
success.LinesRemoved = int32Ptr(linesRemoved)
}
if message != "" {
success.Message = stringPtr(message)
}
return &agentv1.EditResult{
Result: &agentv1.EditResult_Success{Success: success},
}
}
func buildFinalEditSuccessResult(path string, afterContent string, payload editResultPayload) *agentv1.EditResult {
diffString := payload.DiffString
linesAdded := payload.LinesAdded
linesRemoved := payload.LinesRemoved
if afterContent != payload.AfterContent {
diffString, linesAdded, linesRemoved = computeEditDiff(payload.BeforeContent, afterContent)
}
return buildSuccessfulEditResult(path, payload.BeforeContent, afterContent, diffString, linesAdded, linesRemoved, payload.Message)
}
func buildEditErrorResult(path string, errorMessage string) *agentv1.EditResult {
trimmedPath := strings.TrimSpace(path)
trimmedError := strings.TrimSpace(errorMessage)
return &agentv1.EditResult{
Result: &agentv1.EditResult_Error{
Error: &agentv1.EditError{
Path: trimmedPath,
Error: trimmedError,
ModelVisibleError: stringPtr(trimmedError),
},
},
}
}
func decodeReadDataAsEditableText(data []byte) (string, bool, string) {
if len(data) > editReadBinaryContentLimit {
return "", false, fmt.Sprintf("Read returned binary data larger than %d bytes", editReadBinaryContentLimit)
}
if !utf8.Valid(data) {
return "", false, "Read returned binary data that is not valid UTF-8 text"
}
return string(data), true, ""
}
func buildEditResultFromReadResult(path string, result *agentv1.ReadResult) *agentv1.EditResult {
if result == nil {
return buildEditErrorResult(path, "read result missing")
}
switch item := result.GetResult().(type) {
case *agentv1.ReadResult_FileNotFound:
return &agentv1.EditResult{
Result: &agentv1.EditResult_FileNotFound{
FileNotFound: &agentv1.EditFileNotFound{Path: firstNonEmpty(item.FileNotFound.GetPath(), path)},
},
}
case *agentv1.ReadResult_PermissionDenied:
return &agentv1.EditResult{
Result: &agentv1.EditResult_ReadPermissionDenied{
ReadPermissionDenied: &agentv1.EditReadPermissionDenied{Path: firstNonEmpty(item.PermissionDenied.GetPath(), path)},
},
}
case *agentv1.ReadResult_Rejected:
return &agentv1.EditResult{
Result: &agentv1.EditResult_Rejected{
Rejected: &agentv1.EditRejected{
Path: firstNonEmpty(item.Rejected.GetPath(), path),
Reason: item.Rejected.GetReason(),
},
},
}
case *agentv1.ReadResult_InvalidFile:
return buildEditErrorResult(firstNonEmpty(item.InvalidFile.GetPath(), path), item.InvalidFile.GetReason())
case *agentv1.ReadResult_Error:
return buildEditErrorResult(firstNonEmpty(item.Error.GetPath(), path), item.Error.GetError())
case *agentv1.ReadResult_Success:
return buildEditErrorResult(firstNonEmpty(item.Success.GetPath(), path), "read result did not include editable text content")
default:
return buildEditErrorResult(path, "unknown read result")
}
}
func buildEditResultFromWriteResult(path string, result *agentv1.WriteResult) *agentv1.EditResult {
if result == nil {
return buildEditErrorResult(path, "write result missing")
}
switch item := result.GetResult().(type) {
case *agentv1.WriteResult_PermissionDenied:
return &agentv1.EditResult{
Result: &agentv1.EditResult_WritePermissionDenied{
WritePermissionDenied: &agentv1.EditWritePermissionDenied{
Path: firstNonEmpty(item.PermissionDenied.GetPath(), path),
Error: item.PermissionDenied.GetError(),
},
},
}
case *agentv1.WriteResult_Rejected:
return &agentv1.EditResult{
Result: &agentv1.EditResult_Rejected{
Rejected: &agentv1.EditRejected{
Path: firstNonEmpty(item.Rejected.GetPath(), path),
Reason: item.Rejected.GetReason(),
},
},
}
case *agentv1.WriteResult_NoSpace:
return buildEditErrorResult(firstNonEmpty(item.NoSpace.GetPath(), path), "no space left")
case *agentv1.WriteResult_Error:
return buildEditErrorResult(firstNonEmpty(item.Error.GetPath(), path), item.Error.GetError())
case *agentv1.WriteResult_Success:
afterContent := item.Success.GetFileContentAfterWrite()
return buildSuccessfulEditResult(firstNonEmpty(item.Success.GetPath(), path), "", afterContent, "", 0, 0, "")
default:
return buildEditErrorResult(path, "unknown write result")
}
}
func extractReadContentForEdit(result *agentv1.ReadResult) (string, bool) {
success := result.GetSuccess()
if success == nil {
return "", false
}
switch output := success.GetOutput().(type) {
case *agentv1.ReadSuccess_Content:
return normalizeLineEndingsToLF(output.Content), true
case *agentv1.ReadSuccess_Data:
text, ok, _ := decodeReadDataAsEditableText(output.Data)
return normalizeLineEndingsToLF(text), ok
}
if success.GetOutputBlobId() != nil {
return "", false
}
return "", true
}
func summarizeEditResult(result *agentv1.EditResult) string {
if result == nil {
return "edit result missing"
}
switch item := result.GetResult().(type) {
case *agentv1.EditResult_Success:
return item.Success.GetAfterFullFileContent()
case *agentv1.EditResult_FileNotFound:
return fmt.Sprintf("file not found: %s", item.FileNotFound.GetPath())
case *agentv1.EditResult_ReadPermissionDenied:
return fmt.Sprintf("permission denied: %s", item.ReadPermissionDenied.GetPath())
case *agentv1.EditResult_WritePermissionDenied:
return item.WritePermissionDenied.GetError()
case *agentv1.EditResult_Rejected:
return item.Rejected.GetReason()
case *agentv1.EditResult_Error:
return item.Error.GetError()
default:
return "unknown edit result"
}
}
func summarizeEditResultWithoutPath(result *agentv1.EditResult) string {
if result == nil {
return "edit result missing"
}
switch item := result.GetResult().(type) {
case *agentv1.EditResult_Success:
return item.Success.GetAfterFullFileContent()
case *agentv1.EditResult_FileNotFound:
return "file not found"
case *agentv1.EditResult_ReadPermissionDenied:
return "permission denied"
case *agentv1.EditResult_WritePermissionDenied:
return item.WritePermissionDenied.GetError()
case *agentv1.EditResult_Rejected:
return item.Rejected.GetReason()
case *agentv1.EditResult_Error:
return item.Error.GetError()
default:
return "unknown edit result"
}
}
func editResultWithoutPath(result *agentv1.EditResult) *agentv1.EditResult {
if result == nil {
return nil
}
switch item := result.GetResult().(type) {
case *agentv1.EditResult_Success:
return &agentv1.EditResult{Result: &agentv1.EditResult_Success{Success: &agentv1.EditSuccess{
AfterFullFileContent: item.Success.GetAfterFullFileContent(),
BeforeFullFileContent: item.Success.BeforeFullFileContent,
DiffString: item.Success.DiffString,
LinesAdded: item.Success.LinesAdded,
LinesRemoved: item.Success.LinesRemoved,
Message: item.Success.Message,
}}}
case *agentv1.EditResult_FileNotFound:
return &agentv1.EditResult{Result: &agentv1.EditResult_FileNotFound{FileNotFound: &agentv1.EditFileNotFound{}}}
case *agentv1.EditResult_ReadPermissionDenied:
return &agentv1.EditResult{Result: &agentv1.EditResult_ReadPermissionDenied{ReadPermissionDenied: &agentv1.EditReadPermissionDenied{}}}
case *agentv1.EditResult_WritePermissionDenied:
return &agentv1.EditResult{Result: &agentv1.EditResult_WritePermissionDenied{WritePermissionDenied: &agentv1.EditWritePermissionDenied{Error: item.WritePermissionDenied.GetError()}}}
case *agentv1.EditResult_Rejected:
return &agentv1.EditResult{Result: &agentv1.EditResult_Rejected{Rejected: &agentv1.EditRejected{Reason: item.Rejected.GetReason()}}}
case *agentv1.EditResult_Error:
return &agentv1.EditResult{Result: &agentv1.EditResult_Error{Error: &agentv1.EditError{
Error: item.Error.GetError(),
ModelVisibleError: item.Error.ModelVisibleError,
}}}
default:
return result
}
}
func summarizeStructuredEditResult(result *agentv1.EditResult) string {
if result == nil {
return "edit result missing"
}
if _, ok := result.GetResult().(*agentv1.EditResult_Success); ok {
encoded, err := protojson.Marshal(result)
if err == nil {
return string(encoded)
}
}
return summarizeEditResult(result)
}
func hiddenWriteControlError(message *agentv1.ExecClientControlMessage) string {
if message == nil {
return "write operation failed"
}
switch item := message.GetMessage().(type) {
case *agentv1.ExecClientControlMessage_Throw:
return firstNonEmpty(strings.TrimSpace(item.Throw.GetError()), "write operation failed")
case *agentv1.ExecClientControlMessage_StreamClose:
return "write operation closed unexpectedly"
default:
return "write operation failed"
}
}
func readJSONStringAny(args map[string]any, keys ...string) (string, bool, bool) {
for _, key := range keys {
value, ok := args[key]
if !ok {
continue
}
typed, ok := value.(string)
if !ok {
return "", true, false
}
return typed, true, true
}
return "", false, false
}
func readJSONIntAny(args map[string]any, keys ...string) (int, bool, bool) {
value, found, err := runtimecore.ReadIntArg(args, keys...)
if err != nil {
return 0, true, false
}
return value, found, found
}
func readJSONBoolAny(args map[string]any, keys ...string) (bool, bool, bool) {
for _, key := range keys {
value, ok := args[key]
if !ok {
continue
}
typed, ok := value.(bool)
if !ok {
return false, true, false
}
return typed, true, true
}
return false, false, false
}
func int32Ptr(value int32) *int32 {
return &value
}
+612
View File
@@ -0,0 +1,612 @@
// events.go 负责构造对外兼容的 legacy RunSSE 消息。
package forwarder
import (
"encoding/json"
"strings"
"google.golang.org/protobuf/proto"
"google.golang.org/protobuf/types/known/structpb"
"cursor/gen/agentv1"
execbridge "cursor/internal/backend/agent/bridge/exec"
runtimecore "cursor/internal/backend/agent/core"
)
// buildHeartbeatMessage 构造一个服务端心跳消息。
func buildHeartbeatMessage() *agentv1.AgentServerMessage {
return &agentv1.AgentServerMessage{
Message: &agentv1.AgentServerMessage_InteractionUpdate{
InteractionUpdate: &agentv1.InteractionUpdate{
Message: &agentv1.InteractionUpdate_Heartbeat{
Heartbeat: &agentv1.HeartbeatUpdate{},
},
},
},
}
}
// buildTextDeltaMessage 构造文本增量消息。
func buildTextDeltaMessage(text string) *agentv1.AgentServerMessage {
return &agentv1.AgentServerMessage{
Message: &agentv1.AgentServerMessage_InteractionUpdate{
InteractionUpdate: &agentv1.InteractionUpdate{
Message: &agentv1.InteractionUpdate_TextDelta{
TextDelta: &agentv1.TextDeltaUpdate{Text: text},
},
},
},
}
}
// buildThinkingDeltaMessage 构造思考文本增量消息。
func buildThinkingDeltaMessage(text string, style agentv1.ThinkingStyle) *agentv1.AgentServerMessage {
styleCopy := style
return &agentv1.AgentServerMessage{
Message: &agentv1.AgentServerMessage_InteractionUpdate{
InteractionUpdate: &agentv1.InteractionUpdate{
Message: &agentv1.InteractionUpdate_ThinkingDelta{
ThinkingDelta: &agentv1.ThinkingDeltaUpdate{
Text: text,
ThinkingStyle: &styleCopy,
},
},
},
},
}
}
// buildThinkingCompletedMessage 构造思考阶段结束消息。
func buildThinkingCompletedMessage(durationMS int32) *agentv1.AgentServerMessage {
return &agentv1.AgentServerMessage{
Message: &agentv1.AgentServerMessage_InteractionUpdate{
InteractionUpdate: &agentv1.InteractionUpdate{
Message: &agentv1.InteractionUpdate_ThinkingCompleted{
ThinkingCompleted: &agentv1.ThinkingCompletedUpdate{
ThinkingDurationMs: durationMS,
},
},
},
},
}
}
// buildSummaryStartedMessage 构造摘要开始消息。
func buildSummaryStartedMessage() *agentv1.AgentServerMessage {
return &agentv1.AgentServerMessage{
Message: &agentv1.AgentServerMessage_InteractionUpdate{
InteractionUpdate: &agentv1.InteractionUpdate{
Message: &agentv1.InteractionUpdate_SummaryStarted{
SummaryStarted: &agentv1.SummaryStartedUpdate{},
},
},
},
}
}
// buildSummaryMessage 构造摘要增量/完成内容消息。
func buildSummaryMessage(summary string) *agentv1.AgentServerMessage {
return &agentv1.AgentServerMessage{
Message: &agentv1.AgentServerMessage_InteractionUpdate{
InteractionUpdate: &agentv1.InteractionUpdate{
Message: &agentv1.InteractionUpdate_Summary{
Summary: &agentv1.SummaryUpdate{Summary: strings.TrimSpace(summary)},
},
},
},
}
}
// buildSummaryCompletedMessage 构造摘要完成消息。
// field 11 在 legacy aiserver 协议里会被客户端按 GetThoughtAnnotationRequest 解码,
// 所以这里需要把 request_id 写入 field 1,驱动客户端继续查询 "Chat context summarized" 注解。
func buildSummaryCompletedMessage(requestID string) *agentv1.AgentServerMessage {
var message *string
if strings.TrimSpace(requestID) != "" {
value := strings.TrimSpace(requestID)
message = &value
}
return &agentv1.AgentServerMessage{
Message: &agentv1.AgentServerMessage_InteractionUpdate{
InteractionUpdate: &agentv1.InteractionUpdate{
Message: &agentv1.InteractionUpdate_SummaryCompleted{
SummaryCompleted: &agentv1.SummaryCompletedUpdate{
HookMessage: message,
},
},
},
},
}
}
// buildToolCallStartedMessage 构造工具调用开始消息。
func buildToolCallStartedMessage(callID string, modelCallID string, toolCall *agentv1.ToolCall) *agentv1.AgentServerMessage {
return &agentv1.AgentServerMessage{
Message: &agentv1.AgentServerMessage_InteractionUpdate{
InteractionUpdate: &agentv1.InteractionUpdate{
Message: &agentv1.InteractionUpdate_ToolCallStarted{
ToolCallStarted: &agentv1.ToolCallStartedUpdate{
CallId: callID,
ToolCall: toolCall,
ModelCallId: modelCallID,
},
},
},
},
}
}
// buildPartialToolCallMessage 构造工具调用参数流式生成中的兼容消息。
func buildPartialToolCallMessage(callID string, modelCallID string, toolCall *agentv1.ToolCall, argsTextDelta string) *agentv1.AgentServerMessage {
return &agentv1.AgentServerMessage{
Message: &agentv1.AgentServerMessage_InteractionUpdate{
InteractionUpdate: &agentv1.InteractionUpdate{
Message: &agentv1.InteractionUpdate_PartialToolCall{
PartialToolCall: &agentv1.PartialToolCallUpdate{
CallId: callID,
ToolCall: toolCall,
ArgsTextDelta: argsTextDelta,
ModelCallId: modelCallID,
},
},
},
},
}
}
// buildToolCallDeltaMessage 构造工具调用流式增量消息。
func buildToolCallDeltaMessage(callID string, modelCallID string, delta *agentv1.ToolCallDelta) *agentv1.AgentServerMessage {
return &agentv1.AgentServerMessage{
Message: &agentv1.AgentServerMessage_InteractionUpdate{
InteractionUpdate: &agentv1.InteractionUpdate{
Message: &agentv1.InteractionUpdate_ToolCallDelta{
ToolCallDelta: &agentv1.ToolCallDeltaUpdate{
CallId: callID,
ToolCallDelta: delta,
ModelCallId: modelCallID,
},
},
},
},
}
}
// buildToolCallCompletedMessage 构造工具调用完成消息。
func buildToolCallCompletedMessage(callID string, modelCallID string, toolCall *agentv1.ToolCall) *agentv1.AgentServerMessage {
return &agentv1.AgentServerMessage{
Message: &agentv1.AgentServerMessage_InteractionUpdate{
InteractionUpdate: &agentv1.InteractionUpdate{
Message: &agentv1.InteractionUpdate_ToolCallCompleted{
ToolCallCompleted: &agentv1.ToolCallCompletedUpdate{
CallId: callID,
ToolCall: toolCall,
ModelCallId: modelCallID,
},
},
},
},
}
}
// buildShellOutputDeltaMessage 把 shell 流输出包装成兼容消息。
func buildShellOutputDeltaMessage(delta *agentv1.ShellOutputDeltaUpdate) *agentv1.AgentServerMessage {
return &agentv1.AgentServerMessage{
Message: &agentv1.AgentServerMessage_InteractionUpdate{
InteractionUpdate: &agentv1.InteractionUpdate{
Message: &agentv1.InteractionUpdate_ShellOutputDelta{
ShellOutputDelta: delta,
},
},
},
}
}
// buildTurnEndedMessage 构造 turn 结束消息,并携带标准化后的 token 统计。
func buildTurnEndedMessage(inputTokens int64, outputTokens int64, cacheReadTokens int64, cacheWriteTokens int64) *agentv1.AgentServerMessage {
inputTokensValue := inputTokens
outputTokensValue := outputTokens
cacheReadTokensValue := cacheReadTokens
cacheWriteTokensValue := cacheWriteTokens
return &agentv1.AgentServerMessage{
Message: &agentv1.AgentServerMessage_InteractionUpdate{
InteractionUpdate: &agentv1.InteractionUpdate{
Message: &agentv1.InteractionUpdate_TurnEnded{
TurnEnded: &agentv1.TurnEndedUpdate{
InputTokens: &inputTokensValue,
OutputTokens: &outputTokensValue,
CacheReadTokens: &cacheReadTokensValue,
CacheWriteTokens: &cacheWriteTokensValue,
},
},
},
},
}
}
// buildCheckpointMessage 根据投影出的状态生成 legacy checkpoint 消息。
func buildCheckpointMessage(state *agentv1.ConversationStateStructure) *agentv1.AgentServerMessage {
cloned := &agentv1.ConversationStateStructure{}
if state != nil {
if next, ok := proto.Clone(state).(*agentv1.ConversationStateStructure); ok && next != nil {
cloned = next
}
}
if cloned.TokenDetails == nil {
cloned.TokenDetails = &agentv1.ConversationTokenDetails{}
}
if cloned.TokenDetails.MaxTokens == 0 {
cloned.TokenDetails.MaxTokens = projectedConversationMaxTokens
}
cloned.TurnTimings = nil
return &agentv1.AgentServerMessage{
Message: &agentv1.AgentServerMessage_ConversationCheckpointUpdate{
ConversationCheckpointUpdate: cloned,
},
}
}
// buildExecAbortMessage 构造对客户端执行桥的 abort 控制消息。
func buildExecAbortMessage(pending runtimecore.PendingExec) *agentv1.AgentServerMessage {
return &agentv1.AgentServerMessage{
Message: &agentv1.AgentServerMessage_ExecServerControlMessage{
ExecServerControlMessage: &agentv1.ExecServerControlMessage{
Message: &agentv1.ExecServerControlMessage_Abort{
Abort: &agentv1.ExecServerAbort{
Id: pending.MessageID,
},
},
},
},
}
}
// buildStartedToolCall 把工具意图映射为可发送给客户端的 started ToolCall 结构。
func buildStartedToolCall(invocation runtimecore.ToolInvocation) *agentv1.ToolCall {
switch strings.TrimSpace(invocation.ToolName) {
case "Glob":
args, err := execbridge.DecodeGlobToolArgs(invocation.ArgsJSON)
if err != nil {
args = &agentv1.GlobToolArgs{}
}
return &agentv1.ToolCall{
Tool: &agentv1.ToolCall_GlobToolCall{
GlobToolCall: &agentv1.GlobToolCall{Args: args},
},
}
case "Grep":
args, err := execbridge.DecodeGrepToolArgs(invocation.ArgsJSON, invocation.CallID)
if err != nil && args == nil {
args = &agentv1.GrepArgs{ToolCallId: invocation.CallID}
}
return &agentv1.ToolCall{
Tool: &agentv1.ToolCall_GrepToolCall{
GrepToolCall: &agentv1.GrepToolCall{Args: args},
},
}
case "Read":
args, err := execbridge.DecodeReadToolArgs(invocation.ArgsJSON)
if err != nil && args == nil {
args = &agentv1.ReadToolArgs{}
}
return &agentv1.ToolCall{
Tool: &agentv1.ToolCall_ReadToolCall{
ReadToolCall: &agentv1.ReadToolCall{
Args: args,
},
},
}
case "TodoWrite":
args, _ := decodeUpdateTodosArgsJSON(invocation.ArgsJSON)
return &agentv1.ToolCall{
Tool: &agentv1.ToolCall_UpdateTodosToolCall{
UpdateTodosToolCall: &agentv1.UpdateTodosToolCall{Args: args},
},
}
case "ReadTodos":
args, _ := decodeReadTodosArgsJSON(invocation.ArgsJSON)
return &agentv1.ToolCall{
Tool: &agentv1.ToolCall_ReadTodosToolCall{
ReadTodosToolCall: &agentv1.ReadTodosToolCall{Args: args},
},
}
case "Delete":
var input struct {
Path string `json:"path"`
}
_ = json.Unmarshal(invocation.ArgsJSON, &input)
return &agentv1.ToolCall{
Tool: &agentv1.ToolCall_DeleteToolCall{
DeleteToolCall: &agentv1.DeleteToolCall{
Args: &agentv1.DeleteArgs{Path: strings.TrimSpace(input.Path), ToolCallId: invocation.CallID},
},
},
}
case "Shell":
var input struct {
Command string `json:"command"`
Description string `json:"description,omitempty"`
WorkingDirectory string `json:"working_directory,omitempty"`
NotifyOnOutput *struct {
Pattern string `json:"pattern"`
Reason string `json:"reason"`
DebounceMS *float64 `json:"debounce_ms,omitempty"`
NotificationLimit *int32 `json:"notification_limit,omitempty"`
} `json:"notify_on_output,omitempty"`
}
_ = json.Unmarshal(invocation.ArgsJSON, &input)
return &agentv1.ToolCall{
Tool: &agentv1.ToolCall_ShellToolCall{
ShellToolCall: &agentv1.ShellToolCall{
Args: &agentv1.ShellArgs{
Command: strings.TrimSpace(input.Command),
WorkingDirectory: strings.TrimSpace(input.WorkingDirectory),
ToolCallId: invocation.CallID,
Description: stringPtr(strings.TrimSpace(input.Description)),
OutputNotification: buildShellOutputNotificationConfig(input.NotifyOnOutput),
},
},
},
}
case "AwaitShell":
args, err := decodeAwaitShellArgs(invocation.ArgsJSON)
if err != nil {
args = awaitShellArgs{}
}
return buildAwaitShellToolCall(buildAwaitArgsFromAwaitShellArgs(args), nil)
case "WriteShellStdin":
var input struct {
ShellID uint32 `json:"shell_id"`
Chars string `json:"chars"`
}
_ = json.Unmarshal(invocation.ArgsJSON, &input)
return &agentv1.ToolCall{
Tool: &agentv1.ToolCall_WriteShellStdinToolCall{
WriteShellStdinToolCall: &agentv1.WriteShellStdinToolCall{
Args: &agentv1.WriteShellStdinArgs{
ShellId: input.ShellID,
Chars: input.Chars,
},
},
},
}
case "Task":
return &agentv1.ToolCall{
Tool: &agentv1.ToolCall_TaskToolCall{
TaskToolCall: &agentv1.TaskToolCall{
Args: buildTaskArgsFromJSON(invocation.ArgsJSON),
},
},
}
case "Ls":
var input struct {
Path string `json:"path"`
Ignore []string `json:"ignore,omitempty"`
}
_ = json.Unmarshal(invocation.ArgsJSON, &input)
return &agentv1.ToolCall{
Tool: &agentv1.ToolCall_LsToolCall{
LsToolCall: &agentv1.LsToolCall{
Args: &agentv1.LsArgs{
Path: strings.TrimSpace(input.Path),
Ignore: append([]string(nil), input.Ignore...),
ToolCallId: invocation.CallID,
},
},
},
}
case "Write":
var input struct {
Path string `json:"path"`
Contents string `json:"contents"`
}
_ = json.Unmarshal(invocation.ArgsJSON, &input)
return &agentv1.ToolCall{
Tool: &agentv1.ToolCall_EditToolCall{
EditToolCall: &agentv1.EditToolCall{
Args: &agentv1.EditArgs{
Path: strings.TrimSpace(input.Path),
StreamContent: literalStringPtr(input.Contents),
},
},
},
}
case "PatchEdit":
var input struct {
FilePath string `json:"file_path"`
Path string `json:"path,omitempty"`
}
_ = json.Unmarshal(invocation.ArgsJSON, &input)
path := strings.TrimSpace(input.FilePath)
if path == "" {
path = strings.TrimSpace(input.Path)
}
return &agentv1.ToolCall{
Tool: &agentv1.ToolCall_EditToolCall{
EditToolCall: &agentv1.EditToolCall{
Args: &agentv1.EditArgs{
Path: path,
},
},
},
}
case "CallMcpTool":
var input struct {
Server string `json:"server"`
ToolName string `json:"toolName"`
Arguments map[string]any `json:"arguments,omitempty"`
}
_ = json.Unmarshal(invocation.ArgsJSON, &input)
return &agentv1.ToolCall{
Tool: &agentv1.ToolCall_McpToolCall{
McpToolCall: &agentv1.McpToolCall{
Args: &agentv1.McpArgs{
Name: forwarderCanonicalMCPToolLookupName(input.Server, input.ToolName),
Args: forwarderStructValueMap(input.Arguments),
ToolCallId: invocation.CallID,
ProviderIdentifier: strings.TrimSpace(input.Server),
ToolName: strings.TrimSpace(input.ToolName),
},
},
},
}
case "ListMcpResources":
var input struct {
Server string `json:"server,omitempty"`
}
_ = json.Unmarshal(invocation.ArgsJSON, &input)
return &agentv1.ToolCall{
Tool: &agentv1.ToolCall_ListMcpResourcesToolCall{
ListMcpResourcesToolCall: &agentv1.ListMcpResourcesToolCall{
Args: &agentv1.ListMcpResourcesExecArgs{
Server: stringPtr(strings.TrimSpace(input.Server)),
},
},
},
}
case "AskQuestion":
var args agentv1.AskQuestionArgs
_ = json.Unmarshal(invocation.ArgsJSON, &args)
return &agentv1.ToolCall{
Tool: &agentv1.ToolCall_AskQuestionToolCall{
AskQuestionToolCall: &agentv1.AskQuestionToolCall{Args: &args},
},
}
case "WebSearch":
var args agentv1.WebSearchArgs
_ = json.Unmarshal(invocation.ArgsJSON, &args)
return &agentv1.ToolCall{
Tool: &agentv1.ToolCall_WebSearchToolCall{
WebSearchToolCall: &agentv1.WebSearchToolCall{Args: &args},
},
}
case "WebFetch":
var args agentv1.WebFetchArgs
_ = json.Unmarshal(invocation.ArgsJSON, &args)
return &agentv1.ToolCall{
Tool: &agentv1.ToolCall_WebFetchToolCall{
WebFetchToolCall: &agentv1.WebFetchToolCall{Args: &args},
},
}
case "CreatePlan":
args, err := runtimecore.DecodeCreatePlanArgsJSON(invocation.ArgsJSON)
if err != nil {
args = &agentv1.CreatePlanArgs{}
}
return &agentv1.ToolCall{
Tool: &agentv1.ToolCall_CreatePlanToolCall{
CreatePlanToolCall: &agentv1.CreatePlanToolCall{Args: args},
},
}
case "GenerateImage":
carrier, _ := decodeGenerateImageToolCarrier(invocation.ArgsJSON)
args := buildGenerateImageArgsFromCarrier(carrier)
return &agentv1.ToolCall{
Tool: &agentv1.ToolCall_GenerateImageToolCall{
GenerateImageToolCall: &agentv1.GenerateImageToolCall{Args: args},
},
}
case "SwitchMode":
var args agentv1.SwitchModeArgs
_ = json.Unmarshal(invocation.ArgsJSON, &args)
args.ToolCallId = invocation.CallID
return &agentv1.ToolCall{
Tool: &agentv1.ToolCall_SwitchModeToolCall{
SwitchModeToolCall: &agentv1.SwitchModeToolCall{Args: &args},
},
}
case "FetchMcpResource":
var input struct {
Server string `json:"server"`
URI string `json:"uri"`
DownloadPath string `json:"downloadPath,omitempty"`
}
_ = json.Unmarshal(invocation.ArgsJSON, &input)
return &agentv1.ToolCall{
Tool: &agentv1.ToolCall_ReadMcpResourceToolCall{
ReadMcpResourceToolCall: &agentv1.ReadMcpResourceToolCall{
Args: &agentv1.ReadMcpResourceExecArgs{
Server: strings.TrimSpace(input.Server),
Uri: strings.TrimSpace(input.URI),
DownloadPath: stringPtr(strings.TrimSpace(input.DownloadPath)),
},
},
},
}
default:
return nil
}
}
func forwarderCanonicalMCPToolLookupName(server string, toolName string) string {
trimmedServer := strings.TrimSpace(server)
trimmedToolName := strings.TrimSpace(toolName)
if trimmedToolName == "" {
return ""
}
if trimmedServer == "" {
return trimmedToolName
}
return trimmedServer + "-" + trimmedToolName
}
func forwarderStructValueMap(items map[string]any) map[string]*structpb.Value {
if len(items) == 0 {
return map[string]*structpb.Value{}
}
result := make(map[string]*structpb.Value, len(items))
for key, value := range items {
item, err := structpb.NewValue(value)
if err != nil {
continue
}
result[key] = item
}
return result
}
// stringPtr 在字符串非空时返回指针,避免把空串写入 optional 字段。
func buildShellOutputNotificationConfig(input *struct {
Pattern string `json:"pattern"`
Reason string `json:"reason"`
DebounceMS *float64 `json:"debounce_ms,omitempty"`
NotificationLimit *int32 `json:"notification_limit,omitempty"`
}) *agentv1.ShellOutputNotificationConfig {
if input == nil {
return nil
}
pattern := strings.TrimSpace(input.Pattern)
reason := strings.TrimSpace(input.Reason)
if pattern == "" || reason == "" {
return nil
}
var debounce *float64
if input.DebounceMS != nil {
value := *input.DebounceMS / 1000
if value < 5 {
value = 5
}
debounce = &value
}
return &agentv1.ShellOutputNotificationConfig{
Pattern: pattern,
Reason: reason,
Debounce: debounce,
NotificationLimit: input.NotificationLimit,
}
}
func stringPtr(value string) *string {
if strings.TrimSpace(value) == "" {
return nil
}
next := strings.TrimSpace(value)
return &next
}
func literalStringPtr(value string) *string {
if value == "" {
return nil
}
next := value
return &next
}
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,172 @@
package forwarder
import (
"errors"
"log"
"os"
"path/filepath"
"strconv"
"strings"
"time"
)
const historyMaintenanceLockStaleAfter = 30 * time.Minute
const historyMaintenanceTempStaleAfter = 5 * time.Minute
func (service *Service) startHistoryMaintenance() {
if service == nil || service.store == nil {
return
}
go func() {
if err := service.runHistoryMaintenance(); err != nil {
log.Printf("forwarder history maintenance failed: %v", err)
}
}()
}
func (service *Service) runHistoryMaintenance() error {
historyRoot := strings.TrimSpace(service.store.HistoryDir())
if historyRoot == "" {
return nil
}
release, ok, err := acquireHistoryMaintenanceLock(historyRoot)
if err != nil || !ok {
return err
}
defer release()
entries, err := os.ReadDir(historyRoot)
if err != nil {
if errors.Is(err, os.ErrNotExist) {
return nil
}
return err
}
for _, entry := range entries {
if !entry.IsDir() {
cleanupRootLegacyHistoryArtifact(historyRoot, entry.Name())
continue
}
service.cleanupConversationLegacyArtifacts(filepath.Join(historyRoot, entry.Name()))
}
return nil
}
func cleanupRootLegacyHistoryArtifact(historyRoot string, name string) {
trimmedName := strings.TrimSpace(name)
if trimmedName == "" || trimmedName == ".history-maintenance.lock" {
return
}
if !isRootLegacyHistoryArtifact(trimmedName) {
return
}
path := filepath.Join(historyRoot, trimmedName)
if strings.Contains(trimmedName, ".tmp-") {
cleanupStaleHistoryTempArtifact(path)
return
}
if err := os.RemoveAll(path); err != nil && !errors.Is(err, os.ErrNotExist) {
log.Printf("forwarder root legacy cleanup failed path=%s err=%v", path, err)
}
}
func isRootLegacyHistoryArtifact(name string) bool {
if strings.TrimSpace(name) == usageFileName || strings.TrimSpace(name) == usageFileName+".lock" {
return false
}
if strings.HasSuffix(name, ".json") || strings.HasSuffix(name, ".jsonl") {
return true
}
if strings.Contains(name, ".tmp-") {
return true
}
if strings.HasSuffix(name, ".lock") {
return true
}
return false
}
func (service *Service) cleanupConversationLegacyArtifacts(conversationDir string) {
for _, name := range []string{
"turns",
"active",
"latest.json",
"summary.json",
"replay.json",
"runtime.json",
"request.json",
"recovery.json",
"conversation.json",
"entries.jsonl",
} {
path := filepath.Join(conversationDir, name)
if err := os.RemoveAll(path); err != nil && !errors.Is(err, os.ErrNotExist) {
log.Printf("forwarder legacy cleanup failed path=%s err=%v", path, err)
}
}
entries, err := os.ReadDir(conversationDir)
if err != nil {
return
}
for _, entry := range entries {
if !entry.IsDir() {
if strings.Contains(entry.Name(), ".tmp-") {
cleanupStaleHistoryTempArtifact(filepath.Join(conversationDir, entry.Name()))
}
continue
}
if _, err := strconv.Atoi(strings.TrimSpace(entry.Name())); err != nil {
continue
}
path := filepath.Join(conversationDir, entry.Name())
if err := os.RemoveAll(path); err != nil && !errors.Is(err, os.ErrNotExist) {
log.Printf("forwarder numeric legacy cleanup failed path=%s err=%v", path, err)
}
}
}
func cleanupStaleHistoryTempArtifact(path string) {
info, err := os.Stat(path)
if err != nil {
return
}
if time.Since(info.ModTime()) < historyMaintenanceTempStaleAfter {
return
}
if err := os.Remove(path); err != nil && !errors.Is(err, os.ErrNotExist) {
log.Printf("forwarder temp cleanup failed path=%s err=%v", path, err)
}
}
func acquireHistoryMaintenanceLock(historyRoot string) (func(), bool, error) {
if err := os.MkdirAll(historyRoot, 0o755); err != nil {
return nil, false, err
}
lockPath := filepath.Join(historyRoot, ".history-maintenance.lock")
for attempt := 0; attempt < 2; attempt++ {
file, err := os.OpenFile(lockPath, os.O_CREATE|os.O_EXCL|os.O_WRONLY, 0o600)
if err == nil {
_, _ = file.WriteString(time.Now().UTC().Format(time.RFC3339Nano))
_ = file.Close()
return func() {
_ = os.Remove(lockPath)
}, true, nil
}
if !errors.Is(err, os.ErrExist) {
return nil, false, err
}
info, statErr := os.Stat(lockPath)
if statErr != nil {
if errors.Is(statErr, os.ErrNotExist) {
continue
}
return nil, false, statErr
}
if time.Since(info.ModTime()) > historyMaintenanceLockStaleAfter {
_ = os.Remove(lockPath)
continue
}
return nil, false, nil
}
return nil, false, nil
}
@@ -0,0 +1,533 @@
package forwarder
import (
"encoding/base64"
"encoding/json"
"fmt"
"strings"
"time"
"cursor/gen/agentv1"
runtimecore "cursor/internal/backend/agent/core"
)
func (service *Service) handleInteractionToolInvocation(stream *ActiveStream, invocation runtimecore.ToolInvocation) error {
if service == nil || stream == nil {
return fmt.Errorf("active stream is required")
}
serverMessage, pendingInteraction, err := service.interactionBridge.OpenQuery(invocation)
if err != nil {
return newRecoverableToolInvocationError(err)
}
pendingInteraction.ModelCallID = invocation.ModelCallID
pendingInteraction.ReasoningContent = invocation.ReasoningContent
pendingInteraction.ReasoningSignature = invocation.ReasoningSignature
pendingInteraction.ReasoningSignatureSource = invocation.ReasoningSignatureSource
if pendingInteraction.OpenedAt.IsZero() {
pendingInteraction.OpenedAt = time.Now().UTC()
}
stream.mu.Lock()
if stream.PendingInteractions == nil {
stream.PendingInteractions = make(map[string]runtimecore.PendingInteraction)
}
pendingInteraction.ProviderPass = stream.ProviderPassCount
stream.PendingInteractions[pendingInteraction.InteractionID] = pendingInteraction
stream.UpdatedAt = time.Now().UTC()
stream.mu.Unlock()
removePending := func() {
stream.mu.Lock()
delete(stream.PendingInteractions, pendingInteraction.InteractionID)
stream.UpdatedAt = time.Now().UTC()
stream.mu.Unlock()
}
if err := service.publishCheckpoint(stream.RequestID, stream.ConversationID); err != nil {
removePending()
return err
}
if err := service.broker.Publish(stream.RequestID, StreamEvent{Message: serverMessage}); err != nil {
removePending()
return err
}
return nil
}
func (service *Service) handleInteractionResult(intent InboundIntent) error {
stream, ok := service.broker.Get(intent.RequestID)
if !ok || stream == nil {
return fmt.Errorf("request is not active: %s", intent.RequestID)
}
if intent.InteractionResponse == nil {
return fmt.Errorf("interaction response is required")
}
pending, found := selectPendingInteraction(intent.InteractionResponse, stream)
if !found {
return fmt.Errorf("pending interaction not found")
}
result, err := service.interactionBridge.ApplyInteractionResponse(intent.InteractionResponse, pending)
if err != nil {
return err
}
markInteractionCompleted(stream, pending)
toolName := strings.TrimSpace(deriveToolNameFromPendingInteraction(pending))
if result.ToolCall != nil {
applySwitchModeMetadata(stream, result.ToolCall)
if err := service.appendToolResult(stream, result.ToolCallID, toolName, pending.ArgsJSON, result.ToolResultPayload, pending.ReasoningContent, result.ToolCall); err != nil {
return err
}
} else if strings.TrimSpace(result.ToolResultPayload) != "" {
if err := service.appendToolResult(stream, pending.ToolCallID, toolName, pending.ArgsJSON, result.ToolResultPayload, pending.ReasoningContent, nil); err != nil {
return err
}
}
if err := service.publishToolCallCompleted(intent.RequestID, result.ToolCallID, pending.ModelCallID, result.ToolCall); err != nil {
return err
}
if err := service.applyApprovedSwitchMode(stream, stream.ConversationID, result.ToolCall); err != nil {
return err
}
if switchToolCall := result.ToolCall.GetSwitchModeToolCall(); switchToolCall != nil && switchToolCall.GetResult().GetSuccess() != nil {
targetMode, err := switchModeTarget(switchToolCall.GetArgs())
if err != nil {
return err
}
modeEntry, err := newModeMetadataEntry(stream.TurnSeq, intent.RequestID, targetMode, true, ModeSourceSwitchModeTool)
if err != nil {
return err
}
if _, err := service.appendConversationEntries(stream, stream.ConversationID, []HistoryEntry{
modeEntry,
newModeChangePromptContextEntry(stream.TurnSeq, intent.RequestID, targetMode),
}); err != nil {
return err
}
}
if err := service.syncSummaryCarryForward(stream.ConversationID, intent.RequestID, pending.ModelCallID); err != nil {
return err
}
if err := service.publishCheckpoint(intent.RequestID, stream.ConversationID); err != nil {
return err
}
if !shouldAutoResumeAfterInteraction(pending) {
rememberPendingProviderCompletion(stream, pendingTurnCompletion{
ConversationID: stream.ConversationID,
RequestID: stream.RequestID,
TurnSeq: stream.TurnSeq,
ModelCallID: pending.ModelCallID,
ProviderPass: pending.ProviderPass,
Disposition: completionDispositionCompleteAfterExternal,
})
}
return service.reconcileStream(stream)
}
func (service *Service) applyApprovedSwitchMode(stream *ActiveStream, conversationID string, toolCall *agentv1.ToolCall) error {
switchToolCall := toolCall.GetSwitchModeToolCall()
if switchToolCall == nil || switchToolCall.GetResult().GetSuccess() == nil {
return nil
}
targetMode, err := switchModeTarget(switchToolCall.GetArgs())
if err != nil {
return err
}
targetAlias, err := modeAlias(targetMode)
if err != nil {
return err
}
stream.mu.Lock()
stream.Mode = targetMode
stream.UpdatedAt = time.Now().UTC()
stream.mu.Unlock()
_, err = service.updateConversationMetaAndCheckpoint(stream, conversationID, func(item *ConversationFile) error {
if item == nil {
return nil
}
item.Mode = targetAlias
return nil
})
return err
}
func applySwitchModeMetadata(stream *ActiveStream, toolCall *agentv1.ToolCall) {
switchToolCall := toolCall.GetSwitchModeToolCall()
if switchToolCall == nil {
return
}
success := switchToolCall.GetResult().GetSuccess()
if success == nil {
return
}
if fromAlias, err := modeAlias(currentStreamMode(stream)); err == nil {
success.FromModeId = fromAlias
}
}
func switchModeTarget(args *agentv1.SwitchModeArgs) (agentv1.AgentMode, error) {
if args == nil {
return agentv1.AgentMode_AGENT_MODE_UNSPECIFIED, fmt.Errorf("switch mode args are required")
}
return parseTargetModeID(args.GetTargetModeId())
}
func markInteractionCompleted(stream *ActiveStream, pending runtimecore.PendingInteraction) {
if stream == nil {
return
}
stream.mu.Lock()
delete(stream.PendingInteractions, pending.InteractionID)
stream.UpdatedAt = time.Now().UTC()
stream.mu.Unlock()
}
func deriveToolNameFromPendingInteraction(pending runtimecore.PendingInteraction) string {
switch strings.TrimSpace(pending.InteractionKind) {
case "ask_question":
return "AskQuestion"
case "create_plan":
return "CreatePlan"
case "web_search":
return "WebSearch"
case "web_fetch":
return "WebFetch"
case "switch_mode":
return "SwitchMode"
default:
return ""
}
}
func isInteractionTool(name string) bool {
switch strings.TrimSpace(name) {
case "AskQuestion", "CreatePlan", "WebSearch", "WebFetch", "SwitchMode":
return true
default:
return false
}
}
func shouldAutoResumeAfterInteraction(pending runtimecore.PendingInteraction) bool {
switch strings.TrimSpace(pending.InteractionKind) {
case "create_plan":
return false
default:
return true
}
}
type generateImageToolCarrier struct {
Description string `json:"description,omitempty"`
FilePath string `json:"file_path,omitempty"`
ReferenceImagePaths []string `json:"reference_image_paths,omitempty"`
ImageData string `json:"image_data,omitempty"`
}
func isImmediateNativeTool(name string) bool {
switch strings.TrimSpace(name) {
case "GenerateImage", "AwaitShell":
return true
default:
return false
}
}
func (service *Service) handleImmediateNativeToolInvocation(stream *ActiveStream, invocation runtimecore.ToolInvocation) error {
switch strings.TrimSpace(invocation.ToolName) {
case "GenerateImage":
return service.handleGenerateImageToolInvocation(stream, invocation)
case "AwaitShell":
return service.handleAwaitShellToolInvocation(stream, invocation)
default:
return fmt.Errorf("unsupported immediate native tool: %s", invocation.ToolName)
}
}
func (service *Service) handleGenerateImageToolInvocation(stream *ActiveStream, invocation runtimecore.ToolInvocation) error {
carrier, decodeErr := decodeGenerateImageToolCarrier(invocation.ArgsJSON)
args := buildGenerateImageArgsFromCarrier(carrier)
sanitizedInvocation := invocation
sanitizedInvocation.ArgsJSON = encodeGenerateImageArgsForHistory(args)
if decodeErr != nil {
result, payload := buildGenerateImageErrorResult(decodeErr.Error())
return service.completeImmediateToolResult(stream, sanitizedInvocation, payload, buildGenerateImageToolCall(args, result))
}
imageData, err := normalizeGenerateImageBase64(carrier.ImageData)
if err != nil {
result, payload := buildGenerateImageErrorResult(err.Error())
return service.completeImmediateToolResult(stream, sanitizedInvocation, payload, buildGenerateImageToolCall(args, result))
}
result, payload := buildGenerateImageSuccessResult(args.GetFilePath(), imageData)
if err := service.completeImmediateToolResult(stream, sanitizedInvocation, payload, buildGenerateImageToolCall(args, result)); err != nil {
return err
}
markProviderTerminalToolInvocation(stream)
return nil
}
func markProviderTerminalToolInvocation(stream *ActiveStream) {
if stream == nil {
return
}
stream.mu.Lock()
stream.ProviderTerminalToolInvocation = true
stream.UpdatedAt = time.Now().UTC()
stream.mu.Unlock()
}
func decodeGenerateImageToolCarrier(raw []byte) (generateImageToolCarrier, error) {
var carrier generateImageToolCarrier
if len(raw) == 0 {
return carrier, nil
}
if err := json.Unmarshal(raw, &carrier); err != nil {
return carrier, fmt.Errorf("decode GenerateImage args failed: %w", err)
}
carrier.Description = strings.TrimSpace(carrier.Description)
carrier.FilePath = strings.TrimSpace(carrier.FilePath)
carrier.ImageData = strings.TrimSpace(carrier.ImageData)
carrier.ReferenceImagePaths = compactTrimmedStrings(carrier.ReferenceImagePaths)
return carrier, nil
}
func compactTrimmedStrings(values []string) []string {
if len(values) == 0 {
return nil
}
items := make([]string, 0, len(values))
for _, value := range values {
trimmed := strings.TrimSpace(value)
if trimmed != "" {
items = append(items, trimmed)
}
}
return items
}
func buildGenerateImageArgsFromCarrier(carrier generateImageToolCarrier) *agentv1.GenerateImageArgs {
args := &agentv1.GenerateImageArgs{
Description: strings.TrimSpace(carrier.Description),
ReferenceImagePaths: append([]string(nil), carrier.ReferenceImagePaths...),
}
if filePath := strings.TrimSpace(carrier.FilePath); filePath != "" {
args.FilePath = &filePath
}
return args
}
func encodeGenerateImageArgsForHistory(args *agentv1.GenerateImageArgs) []byte {
payload := map[string]any{}
if args != nil {
if description := strings.TrimSpace(args.GetDescription()); description != "" {
payload["description"] = description
}
if filePath := strings.TrimSpace(args.GetFilePath()); filePath != "" {
payload["file_path"] = filePath
}
if referenceImagePaths := compactTrimmedStrings(args.GetReferenceImagePaths()); len(referenceImagePaths) > 0 {
payload["reference_image_paths"] = referenceImagePaths
}
}
encoded, err := json.Marshal(payload)
if err != nil {
return []byte("{}")
}
return encoded
}
func normalizeGenerateImageBase64(raw string) (string, error) {
trimmed := strings.TrimSpace(raw)
if trimmed == "" {
return "", fmt.Errorf("GenerateImage image_data is required")
}
if strings.HasPrefix(strings.ToLower(trimmed), "data:image/") {
return "", fmt.Errorf("GenerateImage image_data must be raw base64 without a data URL prefix")
}
compact := strings.Join(strings.Fields(trimmed), "")
if compact == "" {
return "", fmt.Errorf("GenerateImage image_data is required")
}
if _, err := base64.StdEncoding.DecodeString(compact); err != nil {
return "", fmt.Errorf("GenerateImage image_data must be valid base64: %w", err)
}
return compact, nil
}
func buildGenerateImageToolCall(args *agentv1.GenerateImageArgs, result *agentv1.GenerateImageResult) *agentv1.ToolCall {
return &agentv1.ToolCall{
Tool: &agentv1.ToolCall_GenerateImageToolCall{
GenerateImageToolCall: &agentv1.GenerateImageToolCall{
Args: args,
Result: result,
},
},
}
}
func buildGenerateImageSuccessResult(filePath string, imageData string) (*agentv1.GenerateImageResult, string) {
trimmedFilePath := strings.TrimSpace(filePath)
payload := fmt.Sprintf("generate image success image_data_bytes=%d", len(imageData))
if trimmedFilePath != "" {
payload = fmt.Sprintf("generate image success file_path=%s image_data_bytes=%d", trimmedFilePath, len(imageData))
}
return &agentv1.GenerateImageResult{
Result: &agentv1.GenerateImageResult_Success{
Success: &agentv1.GenerateImageSuccess{
FilePath: trimmedFilePath,
ImageData: imageData,
},
},
}, payload
}
func buildGenerateImageErrorResult(message string) (*agentv1.GenerateImageResult, string) {
errorText := strings.TrimSpace(message)
if errorText == "" {
errorText = "image generation failed"
}
return &agentv1.GenerateImageResult{
Result: &agentv1.GenerateImageResult_Error{
Error: &agentv1.GenerateImageError{Error: errorText},
},
}, errorText
}
func isLocalStateTool(name string) bool {
switch strings.TrimSpace(name) {
case "TodoWrite":
return true
default:
return false
}
}
func (service *Service) handleLocalStateToolInvocation(stream *ActiveStream, invocation runtimecore.ToolInvocation) error {
switch strings.TrimSpace(invocation.ToolName) {
case "TodoWrite":
return service.handleTodoWriteToolInvocation(stream, invocation)
default:
return fmt.Errorf("unsupported local state tool: %s", invocation.ToolName)
}
}
func (service *Service) handleTodoWriteToolInvocation(stream *ActiveStream, invocation runtimecore.ToolInvocation) error {
decodedArgs, err := decodeUpdateTodosArgsJSONWithPresence(invocation.ArgsJSON)
if err != nil {
return service.completeImmediateToolResult(stream, invocation, summarizeUpdateTodosResult(&agentv1.UpdateTodosResult{
Result: &agentv1.UpdateTodosResult_Error{
Error: &agentv1.UpdateTodosError{Error: err.Error()},
},
}), buildUpdateTodosToolCall(&agentv1.UpdateTodosArgs{}, &agentv1.UpdateTodosResult{
Result: &agentv1.UpdateTodosResult_Error{
Error: &agentv1.UpdateTodosError{Error: err.Error()},
},
}))
}
args := decodedArgs.Args
conversation, _, _, err := service.snapshotCheckpointConversation(stream)
if err != nil {
return err
}
structuredState, err := projectConversationStructuredState(conversation)
if err != nil {
return err
}
now := time.Now().UTC()
applyMerge := shouldMergeTodoUpdate(args, decodedArgs.MergeSet, structuredState.Todos)
args.Merge = applyMerge
var nextTodos []*agentv1.TodoItem
if applyMerge {
nextTodos, err = mergeTodoItems(structuredState.Todos, args.GetTodos(), now)
} else {
nextTodos, err = normalizeTodoItems(args.GetTodos(), now, true)
}
if err != nil {
result := buildUpdateTodosErrorResult(args, err.Error())
return service.completeImmediateToolResult(stream, invocation, summarizeUpdateTodosResult(result), buildUpdateTodosToolCall(args, result))
}
if !applyMerge {
if missingIDs := missingActiveTodoReplacementIDs(structuredState.Todos, nextTodos); len(missingIDs) > 0 {
result := buildUpdateTodosErrorResult(args, unsafeTodoReplaceError(missingIDs))
return service.completeImmediateToolResult(stream, invocation, summarizeUpdateTodosResult(result), buildUpdateTodosToolCall(args, result))
}
}
result := &agentv1.UpdateTodosResult{
Result: &agentv1.UpdateTodosResult_Success{
Success: &agentv1.UpdateTodosSuccess{
Todos: cloneTodoItems(nextTodos),
TotalCount: int32(len(nextTodos)),
WasMerge: applyMerge,
},
},
}
return service.completeImmediateToolResult(stream, invocation, summarizeUpdateTodosResult(result), buildUpdateTodosToolCall(args, result))
}
func shouldMergeTodoUpdate(args *agentv1.UpdateTodosArgs, mergeSet bool, existing []*agentv1.TodoItem) bool {
if args.GetMerge() {
return true
}
if mergeSet || len(existing) == 0 {
return false
}
if len(missingActiveTodoReplacementIDs(existing, args.GetTodos())) > 0 {
return true
}
for _, item := range args.GetTodos() {
if item == nil {
continue
}
if strings.TrimSpace(item.GetContent()) == "" {
return true
}
}
return false
}
func (service *Service) handleReadTodosToolInvocation(stream *ActiveStream, invocation runtimecore.ToolInvocation) error {
args, err := decodeReadTodosArgsJSON(invocation.ArgsJSON)
if err != nil {
return service.completeImmediateToolResult(stream, invocation, summarizeReadTodosResult(&agentv1.ReadTodosResult{
Result: &agentv1.ReadTodosResult_Error{
Error: &agentv1.ReadTodosError{Error: err.Error()},
},
}), buildReadTodosToolCall(&agentv1.ReadTodosArgs{}, &agentv1.ReadTodosResult{
Result: &agentv1.ReadTodosResult_Error{
Error: &agentv1.ReadTodosError{Error: err.Error()},
},
}))
}
conversation, _, _, err := service.snapshotCheckpointConversation(stream)
if err != nil {
return err
}
structuredState, err := projectConversationStructuredState(conversation)
if err != nil {
return err
}
filtered := filterTodoItems(structuredState.Todos, args.GetStatusFilter(), args.GetIdFilter())
result := &agentv1.ReadTodosResult{
Result: &agentv1.ReadTodosResult_Success{
Success: &agentv1.ReadTodosSuccess{
Todos: filtered,
TotalCount: int32(len(filtered)),
},
},
}
return service.completeImmediateToolResult(stream, invocation, summarizeReadTodosResult(result), buildReadTodosToolCall(args, result))
}
func (service *Service) completeImmediateToolResult(stream *ActiveStream, invocation runtimecore.ToolInvocation, resultText string, toolCall *agentv1.ToolCall) error {
if err := service.appendToolResult(stream, invocation.CallID, strings.TrimSpace(invocation.ToolName), invocation.ArgsJSON, resultText, invocation.ReasoningContent, toolCall); err != nil {
return err
}
if err := service.publishToolCallCompleted(stream.RequestID, invocation.CallID, invocation.ModelCallID, toolCall); err != nil {
return err
}
if err := service.syncSummaryCarryForward(stream.ConversationID, stream.RequestID, invocation.ModelCallID); err != nil {
return err
}
if err := service.publishCheckpoint(stream.RequestID, stream.ConversationID); err != nil {
return err
}
return service.reconcileStream(stream)
}
+42
View File
@@ -0,0 +1,42 @@
package forwarder
import (
"context"
"errors"
"strings"
)
var errProviderLoopInterrupted = errors.New("provider loop interrupted")
func isTerminalStreamStatus(status StreamStatus) bool {
switch status {
case StreamStatusCanceled, StreamStatusCompleted, StreamStatusFailed:
return true
default:
return false
}
}
func providerLoopInterruptErr(ctx context.Context, stream *ActiveStream, modelCallID string) error {
if ctx != nil && ctx.Err() != nil {
return errProviderLoopInterrupted
}
if stream == nil {
return nil
}
stream.mu.Lock()
defer stream.mu.Unlock()
if isTerminalStreamStatus(stream.Status) {
return errProviderLoopInterrupted
}
switch stream.Phase {
case TurnPhaseCanceled, TurnPhaseCompleted, TurnPhaseFailed:
return errProviderLoopInterrupted
}
expectedModelCallID := strings.TrimSpace(modelCallID)
currentModelCallID := strings.TrimSpace(stream.CurrentModelCallID)
if expectedModelCallID != "" && currentModelCallID != "" && currentModelCallID != expectedModelCallID {
return errProviderLoopInterrupted
}
return nil
}
@@ -0,0 +1,72 @@
package forwarder
import (
"crypto/sha256"
"encoding/hex"
"encoding/json"
"strings"
"time"
modeladapter "cursor/internal/backend/agent/model"
)
func cloneStringAnyMap(value map[string]any) map[string]any {
if len(value) == 0 {
return nil
}
payload, err := json.Marshal(value)
if err != nil {
return nil
}
var decoded map[string]any
if err := json.Unmarshal(payload, &decoded); err != nil {
return nil
}
return decoded
}
func parseRFC3339Time(value string) time.Time {
trimmed := strings.TrimSpace(value)
if trimmed == "" {
return time.Time{}
}
parsed, err := time.Parse(time.RFC3339Nano, trimmed)
if err != nil {
return time.Time{}
}
return parsed.UTC()
}
func marshalCarryForwardMessage(message modeladapter.Message) ([]byte, error) {
payload := map[string]any{
"role": message.Role,
"content": message.Content,
}
if len(message.ContentParts) > 0 {
payload["content_parts"] = message.ContentParts
}
if strings.TrimSpace(message.ReasoningContent) != "" || (strings.TrimSpace(message.Role) == "assistant" && len(message.ToolCalls) > 0) {
payload["reasoning_content"] = message.ReasoningContent
}
if strings.TrimSpace(message.ReasoningSignature) != "" {
payload["reasoning_signature"] = message.ReasoningSignature
}
if strings.TrimSpace(message.ReasoningSignatureSource) != "" {
payload["reasoning_signature_source"] = strings.TrimSpace(message.ReasoningSignatureSource)
}
if len(message.ToolCalls) > 0 {
payload["tool_calls"] = message.ToolCalls
}
if strings.TrimSpace(message.ToolCallID) != "" {
payload["tool_call_id"] = message.ToolCallID
}
if strings.TrimSpace(message.Name) != "" {
payload["name"] = message.Name
}
return json.Marshal(payload)
}
func shortSHA256(value string) string {
sum := sha256.Sum256([]byte(value))
return hex.EncodeToString(sum[:])[:12]
}
@@ -0,0 +1,73 @@
// legacy_stream.go 提供兼容 Cursor legacy RunSSE 响应头的 HTTP 包装器。
package forwarder
import (
"context"
"net/http"
"cursor/gen/agentv1"
"cursor/gen/aiserverv1"
"connectrpc.com/connect"
)
const legacyRunSSEContentType = "text/event-stream"
// NewLegacyRunSSEHandler 构造一个对外表现为 text/event-stream 的 Connect ServerStream 处理器。
func NewLegacyRunSSEHandler(
procedure string,
implementation func(context.Context, *connect.Request[aiserverv1.BidiRequestId], *connect.ServerStream[agentv1.AgentServerMessage]) error,
options ...connect.HandlerOption,
) http.Handler {
inner := connect.NewServerStreamHandler(procedure, implementation, options...)
return http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {
inner.ServeHTTP(newLegacyRunSSEHeaderWriter(writer), request)
})
}
// legacyRunSSEHeaderWriter 在真正写出 header 时覆盖 Content-Type。
type legacyRunSSEHeaderWriter struct {
http.ResponseWriter
wroteHeader bool
}
// newLegacyRunSSEHeaderWriter 创建一个响应头包装器。
func newLegacyRunSSEHeaderWriter(writer http.ResponseWriter) *legacyRunSSEHeaderWriter {
return &legacyRunSSEHeaderWriter{ResponseWriter: writer}
}
// WriteHeader 在输出状态码前补齐 legacy 所需响应头。
func (writer *legacyRunSSEHeaderWriter) WriteHeader(statusCode int) {
writer.applyLegacyHeaders()
writer.wroteHeader = true
writer.ResponseWriter.WriteHeader(statusCode)
}
// Write 在首次写 body 时懒加载 header。
func (writer *legacyRunSSEHeaderWriter) Write(payload []byte) (int, error) {
if !writer.wroteHeader {
writer.WriteHeader(http.StatusOK)
}
return writer.ResponseWriter.Write(payload)
}
// Flush 尝试把底层缓冲区立即刷新给客户端。
func (writer *legacyRunSSEHeaderWriter) Flush() {
if flusher, ok := writer.ResponseWriter.(http.Flusher); ok {
flusher.Flush()
}
}
// Unwrap 返回底层 ResponseWriter。
func (writer *legacyRunSSEHeaderWriter) Unwrap() http.ResponseWriter {
return writer.ResponseWriter
}
// applyLegacyHeaders 设置 Cursor legacy RunSSE 需要的响应头。
func (writer *legacyRunSSEHeaderWriter) applyLegacyHeaders() {
header := writer.ResponseWriter.Header()
header.Set("Content-Type", legacyRunSSEContentType)
if header.Get("Cache-Control") == "" {
header.Set("Cache-Control", "no-cache")
}
}
+34
View File
@@ -0,0 +1,34 @@
// module.go 负责把 forwarder service 装配成 legacy HTTP/Connect handler。
package forwarder
import (
"net/http"
"connectrpc.com/connect"
modeladapter "cursor/internal/backend/agent/model"
)
type Module struct {
Service *Service
LocalBidiHandler http.Handler
LocalRunSSE http.Handler
AiHandler http.Handler
RepositoryServiceHandler http.Handler
UploadServiceHandler http.Handler
}
// NewModule 创建 forwarder 模块,并导出本地 Bidi / RunSSE 处理器。
func NewModule(historyRoot string, channelService modeladapter.ChannelResolver) *Module {
service := NewService(historyRoot, channelService)
legacyBidiAppendProcedure := "/aiserver.v1.BidiService/BidiAppend"
legacyRunSSEProcedure := "/agent.v1.AgentService/RunSSE"
return &Module{
Service: service,
LocalBidiHandler: connect.NewUnaryHandler(legacyBidiAppendProcedure, service.BidiAppend),
LocalRunSSE: NewLegacyRunSSEHandler(legacyRunSSEProcedure, service.RunSSE),
AiHandler: newAIHandler(service),
RepositoryServiceHandler: newRepositoryServiceHandler(service),
UploadServiceHandler: newUploadServiceHandler(service),
}
}
@@ -0,0 +1,745 @@
package forwarder
import (
"encoding/json"
"fmt"
"log"
"strings"
"time"
"google.golang.org/protobuf/encoding/protojson"
"cursor/gen/agentv1"
execbridge "cursor/internal/backend/agent/bridge/exec"
runtimecore "cursor/internal/backend/agent/core"
)
const (
patchEditToolName = "PatchEdit"
patchEditReadExecKindName = "patch_edit_read"
patchEditWriteExecKindName = "patch_edit_write"
patchEditPostReadExecKindName = "patch_edit_post_read"
)
type patchEditArgs struct {
Path string
OldString string
NewString string
NewStringSet bool
ReplaceAll bool
ReplaceAllSet bool
}
type pendingPatchEditPayload struct {
ToolName string `json:"tool_name"`
RawArgsJSON json.RawMessage `json:"raw_args_json,omitempty"`
ArgsDecodeFailed bool `json:"args_decode_failed,omitempty"`
Args patchEditArgs `json:"args,omitempty"`
ResolvedPath string `json:"resolved_path"`
BeforeContent string `json:"before_content,omitempty"`
AfterContent string `json:"after_content,omitempty"`
DiffString string `json:"diff_string,omitempty"`
LinesAdded int32 `json:"lines_added,omitempty"`
LinesRemoved int32 `json:"lines_removed,omitempty"`
Message string `json:"message,omitempty"`
}
type queuedPatchEditOperation struct {
ToolCallID string
ModelCallID string
ProviderPass int
ReasoningContent string
ReasoningSignature string
ReasoningSignatureSource string
Payload pendingPatchEditPayload
}
func isPatchEditToolName(name string) bool {
return strings.TrimSpace(name) == patchEditToolName
}
func isHiddenPatchEditExecKind(kind string) bool {
switch strings.TrimSpace(kind) {
case patchEditReadExecKindName, patchEditWriteExecKindName, patchEditPostReadExecKindName:
return true
default:
return false
}
}
func (service *Service) handlePatchEditToolInvocation(stream *ActiveStream, invocation runtimecore.ToolInvocation) error {
if stream == nil {
return fmt.Errorf("patch edit stream is required")
}
toolName := strings.TrimSpace(invocation.ToolName)
if !isPatchEditToolName(toolName) {
return fmt.Errorf("unsupported patch edit tool: %s", invocation.ToolName)
}
payload := pendingPatchEditPayload{
ToolName: patchEditToolName,
RawArgsJSON: append(json.RawMessage(nil), invocation.ArgsJSON...),
}
args, decodeErr := decodePatchEditArgs(invocation.ArgsJSON)
payload.Args = args
payload.ArgsDecodeFailed = decodeErr != nil
resolvedPath := strings.TrimSpace(args.Path)
payload.ResolvedPath = resolvedPath
if decodeErr == nil {
payload.Args.Path = resolvedPath
visibleArgsJSON, err := payload.Args.MarshalJSON()
if err != nil {
return err
}
invocation.ArgsJSON = visibleArgsJSON
}
startedToolCall := buildStartedToolCall(invocation)
if startedToolCall != nil {
toolCallPayload, err := protojson.Marshal(startedToolCall)
if err != nil {
return err
}
_, err = service.appendConversationEntries(stream, stream.ConversationID, []HistoryEntry{
newToolCallEntryWithProviderMetadata(stream.TurnSeq, stream.RequestID, invocation.CallID, patchEditToolName, invocation.ReasoningContent, invocation.ReasoningSignature, invocation.ReasoningSignatureSource, invocation.ReasoningProviderItemID, invocation.ReasoningProviderStatus, invocation.ReasoningProviderSummary, invocation.ProviderItemID, invocation.ProviderCallID, invocation.ProviderStatus, toolCallPayload),
})
if err != nil {
return err
}
}
if err := service.broker.Publish(stream.RequestID, StreamEvent{
Message: buildToolCallStartedMessage(invocation.CallID, invocation.ModelCallID, startedToolCall),
}); err != nil {
return err
}
providerPass := currentProviderPass(stream)
if decodeErr != nil {
return service.finishPatchEditOperation(stream, invocation.CallID, invocation.ModelCallID, providerPass, invocation.ReasoningContent, payload, buildEditErrorResult(resolvedPath, decodeErr.Error()))
}
return service.startOrQueuePatchEditOperation(stream, queuedPatchEditOperation{
ToolCallID: invocation.CallID,
ModelCallID: invocation.ModelCallID,
ProviderPass: providerPass,
ReasoningContent: invocation.ReasoningContent,
ReasoningSignature: invocation.ReasoningSignature,
ReasoningSignatureSource: invocation.ReasoningSignatureSource,
Payload: payload,
})
}
func (service *Service) startOrQueuePatchEditOperation(stream *ActiveStream, operation queuedPatchEditOperation) error {
key := patchEditQueueKey(operation.Payload.ResolvedPath)
if key != "" && patchEditPathHasInFlightOperation(stream, key) {
enqueuePatchEditOperation(stream, key, operation)
log.Printf(
"forwarder patch edit queued behind in-flight edit conversation_id=%s request_id=%s tool_call_id=%s path=%s",
strings.TrimSpace(stream.ConversationID),
strings.TrimSpace(stream.RequestID),
strings.TrimSpace(operation.ToolCallID),
key,
)
return service.publishCheckpoint(stream.RequestID, stream.ConversationID)
}
return service.startQueuedPatchEditOperation(stream, operation)
}
func (service *Service) startQueuedPatchEditOperation(stream *ActiveStream, operation queuedPatchEditOperation) error {
return service.startHiddenPatchEditRead(
stream,
operation.ToolCallID,
operation.ModelCallID,
operation.ProviderPass,
operation.ReasoningContent,
operation.ReasoningSignature,
operation.ReasoningSignatureSource,
operation.Payload,
)
}
func (service *Service) startHiddenPatchEditRead(stream *ActiveStream, toolCallID string, modelCallID string, providerPass int, reasoningContent string, reasoningSignature string, reasoningSignatureSource string, payload pendingPatchEditPayload) error {
readArgsJSON, err := json.Marshal(map[string]any{"path": strings.TrimSpace(payload.ResolvedPath)})
if err != nil {
return err
}
serverMessage, pendingExec, err := service.execBridge.OpenExec(execbridge.OpenExecContext{
ConversationID: stream.ConversationID,
}, runtimecore.ToolInvocation{
CallID: strings.TrimSpace(toolCallID),
ToolName: "Read",
ArgsJSON: readArgsJSON,
ModelCallID: strings.TrimSpace(modelCallID),
})
if err != nil {
return err
}
pendingArgsJSON, err := payload.MarshalJSON()
if err != nil {
return err
}
pendingExec.ModelCallID = strings.TrimSpace(modelCallID)
pendingExec.ProviderPass = providerPass
pendingExec.ToolCallID = strings.TrimSpace(toolCallID)
pendingExec.ReasoningContent = reasoningContent
pendingExec.ReasoningSignature = strings.TrimSpace(reasoningSignature)
pendingExec.ReasoningSignatureSource = strings.TrimSpace(reasoningSignatureSource)
pendingExec.ExecKind = patchEditReadExecKindName
pendingExec.ArgsJSON = pendingArgsJSON
stream.mu.Lock()
stream.PendingExecs[pendingExec.ExecID] = pendingExec
stream.mu.Unlock()
if err := service.publishCheckpoint(stream.RequestID, stream.ConversationID); err != nil {
return err
}
return service.broker.Publish(stream.RequestID, StreamEvent{Message: serverMessage})
}
func (service *Service) startHiddenPatchEditWrite(stream *ActiveStream, toolCallID string, modelCallID string, providerPass int, reasoningContent string, reasoningSignature string, reasoningSignatureSource string, payload pendingPatchEditPayload, beforeContent string) error {
computation, err := computePatchEdit(payload, beforeContent)
if err != nil {
return service.finishPatchEditOperation(stream, strings.TrimSpace(toolCallID), strings.TrimSpace(modelCallID), providerPass, reasoningContent, payload, buildEditErrorResult(payload.ResolvedPath, err.Error()))
}
transportContent := preparePatchEditWriteContentsForClient(computation.AfterContent)
writeArgsJSON, err := json.Marshal(map[string]any{
"path": strings.TrimSpace(payload.ResolvedPath),
"contents": transportContent,
})
if err != nil {
return err
}
serverMessage, pendingExec, err := service.execBridge.OpenExec(execbridge.OpenExecContext{
ConversationID: stream.ConversationID,
}, runtimecore.ToolInvocation{
CallID: strings.TrimSpace(toolCallID),
ToolName: "Write",
ArgsJSON: writeArgsJSON,
ModelCallID: strings.TrimSpace(modelCallID),
})
if err != nil {
return err
}
payload.BeforeContent = computation.BeforeContent
payload.AfterContent = computation.AfterContent
payload.DiffString = computation.DiffString
payload.LinesAdded = computation.LinesAdded
payload.LinesRemoved = computation.LinesRemoved
payload.Message = computation.Message
pendingArgsJSON, err := payload.MarshalJSON()
if err != nil {
return err
}
pendingExec.ModelCallID = strings.TrimSpace(modelCallID)
pendingExec.ProviderPass = providerPass
pendingExec.ToolCallID = strings.TrimSpace(toolCallID)
pendingExec.ReasoningContent = reasoningContent
pendingExec.ReasoningSignature = strings.TrimSpace(reasoningSignature)
pendingExec.ReasoningSignatureSource = strings.TrimSpace(reasoningSignatureSource)
pendingExec.ExecKind = patchEditWriteExecKindName
pendingExec.ArgsJSON = pendingArgsJSON
stream.mu.Lock()
stream.PendingExecs[pendingExec.ExecID] = pendingExec
stream.mu.Unlock()
if err := service.publishCheckpoint(stream.RequestID, stream.ConversationID); err != nil {
return err
}
return service.broker.Publish(stream.RequestID, StreamEvent{Message: serverMessage})
}
func (service *Service) startHiddenPatchEditPostRead(stream *ActiveStream, toolCallID string, modelCallID string, providerPass int, reasoningContent string, reasoningSignature string, reasoningSignatureSource string, payload pendingPatchEditPayload) error {
readArgsJSON, err := json.Marshal(map[string]any{"path": strings.TrimSpace(payload.ResolvedPath)})
if err != nil {
return err
}
serverMessage, pendingExec, err := service.execBridge.OpenExec(execbridge.OpenExecContext{
ConversationID: stream.ConversationID,
}, runtimecore.ToolInvocation{
CallID: strings.TrimSpace(toolCallID),
ToolName: "Read",
ArgsJSON: readArgsJSON,
ModelCallID: strings.TrimSpace(modelCallID),
})
if err != nil {
return err
}
pendingArgsJSON, err := payload.MarshalJSON()
if err != nil {
return err
}
pendingExec.ModelCallID = strings.TrimSpace(modelCallID)
pendingExec.ProviderPass = providerPass
pendingExec.ToolCallID = strings.TrimSpace(toolCallID)
pendingExec.ReasoningContent = reasoningContent
pendingExec.ReasoningSignature = strings.TrimSpace(reasoningSignature)
pendingExec.ReasoningSignatureSource = strings.TrimSpace(reasoningSignatureSource)
pendingExec.ExecKind = patchEditPostReadExecKindName
pendingExec.ArgsJSON = pendingArgsJSON
stream.mu.Lock()
stream.PendingExecs[pendingExec.ExecID] = pendingExec
stream.mu.Unlock()
if err := service.publishCheckpoint(stream.RequestID, stream.ConversationID); err != nil {
return err
}
return service.broker.Publish(stream.RequestID, StreamEvent{Message: serverMessage})
}
func (service *Service) handleHiddenPatchEditExecResult(stream *ActiveStream, pending runtimecore.PendingExec, message *agentv1.ExecClientMessage) error {
if stream == nil || message == nil {
return fmt.Errorf("patch edit exec result requires stream and message")
}
payload, err := decodePendingPatchEditPayload(pending.ArgsJSON)
if err != nil {
return err
}
switch strings.TrimSpace(pending.ExecKind) {
case patchEditReadExecKindName:
markExecCompleted(stream, pending)
readResult := message.GetReadResult()
if readResult == nil {
return service.finishPatchEditOperation(stream, pending.ToolCallID, pending.ModelCallID, pending.ProviderPass, pending.ReasoningContent, payload, buildEditErrorResult(payload.ResolvedPath, "patch edit read result missing"))
}
content, ok, readErr := extractCompleteReadContentForPatchEdit(readResult)
if !ok {
if readErr != "" {
return service.finishPatchEditOperation(stream, pending.ToolCallID, pending.ModelCallID, pending.ProviderPass, pending.ReasoningContent, payload, buildEditErrorResult(payload.ResolvedPath, readErr))
}
return service.finishPatchEditOperation(stream, pending.ToolCallID, pending.ModelCallID, pending.ProviderPass, pending.ReasoningContent, payload, buildEditResultFromReadResult(payload.ResolvedPath, readResult))
}
return service.startHiddenPatchEditWrite(stream, pending.ToolCallID, pending.ModelCallID, pending.ProviderPass, pending.ReasoningContent, pending.ReasoningSignature, pending.ReasoningSignatureSource, payload, content)
case patchEditWriteExecKindName:
markExecCompleted(stream, pending)
writeResult := message.GetWriteResult()
if writeResult == nil {
return service.finishPatchEditOperation(stream, pending.ToolCallID, pending.ModelCallID, pending.ProviderPass, pending.ReasoningContent, payload, buildEditErrorResult(payload.ResolvedPath, "patch edit write result missing"))
}
switch item := writeResult.GetResult().(type) {
case *agentv1.WriteResult_Success:
payload.ResolvedPath = firstNonEmpty(strings.TrimSpace(item.Success.GetPath()), strings.TrimSpace(payload.ResolvedPath), strings.TrimSpace(patchEditPayloadPath(payload)))
if item.Success.GetFileContentAfterWrite() != "" {
finalAfterContent, reconciled, ok := service.reconcilePatchEditObservedContent(stream, pending.ToolCallID, payload.ResolvedPath, payload.AfterContent, item.Success.GetFileContentAfterWrite())
if !ok {
return service.finishPatchEditOperation(stream, pending.ToolCallID, pending.ModelCallID, pending.ProviderPass, pending.ReasoningContent, payload, buildEditErrorResult(payload.ResolvedPath, "patch edit write verification failed: observed file content differs from expected content"))
}
if finalAfterContent != payload.AfterContent {
payload.DiffString, payload.LinesAdded, payload.LinesRemoved = computeEditDiff(payload.BeforeContent, finalAfterContent)
}
payload.AfterContent = finalAfterContent
if reconciled {
payload.Message = appendPatchEditMessage(payload.Message, "write result matched after client line-ending normalization")
}
}
setPatchEditPayloadPath(&payload, payload.ResolvedPath)
return service.startHiddenPatchEditPostRead(stream, pending.ToolCallID, pending.ModelCallID, pending.ProviderPass, pending.ReasoningContent, pending.ReasoningSignature, pending.ReasoningSignatureSource, payload)
default:
return service.finishPatchEditOperation(stream, pending.ToolCallID, pending.ModelCallID, pending.ProviderPass, pending.ReasoningContent, payload, buildEditResultFromWriteResult(payload.ResolvedPath, writeResult))
}
case patchEditPostReadExecKindName:
markExecCompleted(stream, pending)
readResult := message.GetReadResult()
if readResult == nil {
return service.finishPatchEditOperation(stream, pending.ToolCallID, pending.ModelCallID, pending.ProviderPass, pending.ReasoningContent, payload, buildFinalEditSuccessResult(payload.ResolvedPath, payload.AfterContent, patchEditPayloadAsEditPayload(payload)))
}
if content, ok := extractReadContentForEdit(readResult); ok {
if success := readResult.GetSuccess(); success != nil {
payload.ResolvedPath = firstNonEmpty(strings.TrimSpace(success.GetPath()), payload.ResolvedPath)
setPatchEditPayloadPath(&payload, payload.ResolvedPath)
}
finalAfterContent, reconciled, ok := service.reconcilePatchEditObservedContent(stream, pending.ToolCallID, payload.ResolvedPath, payload.AfterContent, content)
if !ok {
return service.finishPatchEditOperation(stream, pending.ToolCallID, pending.ModelCallID, pending.ProviderPass, pending.ReasoningContent, payload, buildEditErrorResult(payload.ResolvedPath, "patch edit post-read verification failed: observed file content differs from expected content"))
}
if reconciled {
payload.Message = appendPatchEditMessage(payload.Message, "post-read matched after client line-ending normalization")
}
return service.finishPatchEditOperation(stream, pending.ToolCallID, pending.ModelCallID, pending.ProviderPass, pending.ReasoningContent, payload, buildFinalEditSuccessResult(payload.ResolvedPath, finalAfterContent, patchEditPayloadAsEditPayload(payload)))
}
return service.finishPatchEditOperation(stream, pending.ToolCallID, pending.ModelCallID, pending.ProviderPass, pending.ReasoningContent, payload, buildFinalEditSuccessResult(payload.ResolvedPath, payload.AfterContent, patchEditPayloadAsEditPayload(payload)))
default:
return fmt.Errorf("unsupported hidden patch edit exec kind: %s", pending.ExecKind)
}
}
func (service *Service) handleHiddenPatchEditExecControl(stream *ActiveStream, pending runtimecore.PendingExec, message *agentv1.ExecClientControlMessage) error {
if stream == nil || message == nil {
return fmt.Errorf("patch edit exec control requires stream and message")
}
if _, ok := message.GetMessage().(*agentv1.ExecClientControlMessage_Heartbeat); ok {
return nil
}
if _, ok := message.GetMessage().(*agentv1.ExecClientControlMessage_StreamClose); ok {
return nil
}
payload, err := decodePendingPatchEditPayload(pending.ArgsJSON)
if err != nil {
return err
}
markExecCompleted(stream, pending)
if strings.TrimSpace(pending.ExecKind) == patchEditPostReadExecKindName {
return service.finishPatchEditOperation(stream, pending.ToolCallID, pending.ModelCallID, pending.ProviderPass, pending.ReasoningContent, payload, buildFinalEditSuccessResult(payload.ResolvedPath, payload.AfterContent, patchEditPayloadAsEditPayload(payload)))
}
return service.finishPatchEditOperation(stream, pending.ToolCallID, pending.ModelCallID, pending.ProviderPass, pending.ReasoningContent, payload, buildEditErrorResult(payload.ResolvedPath, hiddenPatchEditControlError(message)))
}
func (service *Service) finishPatchEditOperation(stream *ActiveStream, toolCallID string, modelCallID string, providerPass int, reasoningContent string, payload pendingPatchEditPayload, result *agentv1.EditResult) error {
if stream == nil {
return nil
}
if result == nil {
result = buildEditErrorResult(payload.ResolvedPath, "patch edit result missing")
}
path := firstNonEmpty(resultPath(result), payload.ResolvedPath, patchEditPayloadPath(payload))
setPatchEditPayloadPath(&payload, path)
argsJSON, err := patchEditPayloadArgsJSON(payload)
if err != nil {
return err
}
toolCall := buildCompletedEditToolCall(path, result)
historyToolCall := buildCompletedEditToolCall(path, compactPatchEditHistoryEditResult(path, result))
if err := service.appendToolResult(stream, strings.TrimSpace(toolCallID), patchEditToolName, argsJSON, summarizePatchEditResult(path, result), reasoningContent, historyToolCall); err != nil {
return err
}
if err := service.publishToolCallCompleted(stream.RequestID, strings.TrimSpace(toolCallID), strings.TrimSpace(modelCallID), toolCall); err != nil {
return err
}
if err := service.syncSummaryCarryForward(stream.ConversationID, stream.RequestID, strings.TrimSpace(modelCallID)); err != nil {
return err
}
if err := service.publishCheckpoint(stream.RequestID, stream.ConversationID); err != nil {
return err
}
if next, ok := takeNextQueuedPatchEditOperation(stream, firstNonEmpty(payload.ResolvedPath, path)); ok {
if err := service.startQueuedPatchEditOperation(stream, next); err != nil {
return err
}
}
return service.reconcileStream(stream)
}
func patchEditQueueKey(path string) string {
return strings.TrimSpace(path)
}
func patchEditPathHasInFlightOperation(stream *ActiveStream, key string) bool {
if stream == nil || strings.TrimSpace(key) == "" {
return false
}
stream.mu.Lock()
defer stream.mu.Unlock()
for _, pending := range stream.PendingExecs {
if !isHiddenPatchEditExecKind(pending.ExecKind) {
continue
}
payload, err := decodePendingPatchEditPayload(pending.ArgsJSON)
if err != nil {
continue
}
if patchEditQueueKey(firstNonEmpty(payload.ResolvedPath, patchEditPayloadPath(payload))) == key {
return true
}
}
return false
}
func enqueuePatchEditOperation(stream *ActiveStream, key string, operation queuedPatchEditOperation) {
if stream == nil || strings.TrimSpace(key) == "" {
return
}
if len(operation.Payload.RawArgsJSON) > 0 {
operation.Payload.RawArgsJSON = append(json.RawMessage(nil), operation.Payload.RawArgsJSON...)
}
stream.mu.Lock()
if stream.PatchEditQueues == nil {
stream.PatchEditQueues = make(map[string][]queuedPatchEditOperation)
}
stream.PatchEditQueues[key] = append(stream.PatchEditQueues[key], operation)
stream.UpdatedAt = time.Now().UTC()
stream.mu.Unlock()
}
func takeNextQueuedPatchEditOperation(stream *ActiveStream, path string) (queuedPatchEditOperation, bool) {
key := patchEditQueueKey(path)
if stream == nil || key == "" {
return queuedPatchEditOperation{}, false
}
stream.mu.Lock()
defer stream.mu.Unlock()
if stream.PatchEditQueues == nil {
return queuedPatchEditOperation{}, false
}
queue := stream.PatchEditQueues[key]
if len(queue) == 0 {
delete(stream.PatchEditQueues, key)
return queuedPatchEditOperation{}, false
}
next := queue[0]
if len(queue) == 1 {
delete(stream.PatchEditQueues, key)
} else {
remaining := make([]queuedPatchEditOperation, len(queue)-1)
copy(remaining, queue[1:])
stream.PatchEditQueues[key] = remaining
}
stream.UpdatedAt = time.Now().UTC()
return next, true
}
func decodePatchEditArgs(raw []byte) (patchEditArgs, error) {
var rawArgs map[string]any
if err := json.Unmarshal(raw, &rawArgs); err != nil {
return patchEditArgs{}, fmt.Errorf("decode PatchEdit args failed: %w", err)
}
path, pathFound, pathValid := readJSONStringAny(rawArgs, "path", "file_path", "Path", "FilePath")
oldString, oldFound, oldValid := readJSONStringAny(rawArgs, "old_string", "oldString", "OldString")
newString, newFound, newValid := readJSONStringAny(rawArgs, "new_string", "newString", "NewString")
replaceAll, replaceAllFound, replaceAllValid := readJSONBoolAny(rawArgs, "replace_all", "replaceAll", "ReplaceAll")
result := patchEditArgs{
Path: path,
OldString: oldString,
NewString: newString,
NewStringSet: newFound,
ReplaceAll: replaceAll,
ReplaceAllSet: replaceAllFound,
}
switch {
case pathFound && !pathValid:
return result, fmt.Errorf("PatchEdit path must be a string")
case strings.TrimSpace(result.Path) == "":
return result, fmt.Errorf("PatchEdit path is required")
case !isAbsoluteToolPath(result.Path):
return result, fmt.Errorf("PatchEdit path must be absolute")
case oldFound && !oldValid:
return result, fmt.Errorf("PatchEdit old_string must be a string")
case !oldFound || oldString == "":
return result, fmt.Errorf("PatchEdit old_string is required")
case newFound && !newValid:
return result, fmt.Errorf("PatchEdit new_string must be a string")
case !newFound:
return result, fmt.Errorf("PatchEdit new_string is required")
case replaceAllFound && !replaceAllValid:
return result, fmt.Errorf("PatchEdit replace_all must be a boolean")
}
return result, nil
}
func computePatchEdit(payload pendingPatchEditPayload, beforeContent string) (editComputation, error) {
args := payload.Args
occurrences := strings.Count(beforeContent, args.OldString)
switch {
case occurrences == 0:
return editComputation{}, fmt.Errorf("PatchEdit old_string not found")
case occurrences > 1 && !args.ReplaceAll:
return editComputation{}, fmt.Errorf("PatchEdit old_string is not unique; found %d occurrences", occurrences)
}
replaceCount := 1
if args.ReplaceAll {
replaceCount = -1
}
afterContent := strings.Replace(beforeContent, args.OldString, args.NewString, replaceCount)
diffString, linesAdded, linesRemoved := computeEditDiff(beforeContent, afterContent)
return editComputation{
BeforeContent: beforeContent,
AfterContent: afterContent,
DiffString: diffString,
LinesAdded: linesAdded,
LinesRemoved: linesRemoved,
Message: "PatchEdit applied",
}, nil
}
func extractCompleteReadContentForPatchEdit(result *agentv1.ReadResult) (string, bool, string) {
success := result.GetSuccess()
if success == nil {
return "", false, ""
}
if success.GetTruncated() {
return "", false, "patch edit requires a complete Read result, but the file content was truncated"
}
switch output := success.GetOutput().(type) {
case *agentv1.ReadSuccess_Content:
return normalizeLineEndingsToLF(output.Content), true, ""
case *agentv1.ReadSuccess_Data:
text, ok, readErr := decodeReadDataAsEditableText(output.Data)
if !ok {
return "", false, "patch edit requires editable text content, but " + readErr
}
return normalizeLineEndingsToLF(text), true, ""
}
if success.GetOutputBlobId() != nil {
return "", false, "patch edit requires inline file content, but Read returned only an output blob"
}
return "", false, "patch edit requires inline file content, but Read did not include content"
}
func (service *Service) reconcilePatchEditObservedContent(stream *ActiveStream, toolCallID string, path string, expected string, observed string) (string, bool, bool) {
if observed == expected {
return observed, false, true
}
expectedNormalized := normalizeLineEndingsToLF(expected)
observedNormalized := normalizeLineEndingsToLF(observed)
lineEndingEquivalent := expectedNormalized == observedNormalized
log.Printf(
"forwarder patch edit observed content mismatch conversation_id=%s request_id=%s tool_call_id=%s path=%s expected_bytes=%d observed_bytes=%d expected_sha256=%s observed_sha256=%s line_ending_equivalent=%t",
strings.TrimSpace(stream.ConversationID),
strings.TrimSpace(stream.RequestID),
strings.TrimSpace(toolCallID),
strings.TrimSpace(path),
len(expected),
len(observed),
shortSHA256(expected),
shortSHA256(observed),
lineEndingEquivalent,
)
if lineEndingEquivalent {
return observed, true, true
}
return observed, false, false
}
func appendPatchEditMessage(existing string, addition string) string {
existing = strings.TrimSpace(existing)
addition = strings.TrimSpace(addition)
if existing == "" {
return addition
}
if addition == "" {
return existing
}
return existing + "; " + addition
}
func patchEditPayloadPath(payload pendingPatchEditPayload) string {
return strings.TrimSpace(payload.Args.Path)
}
func setPatchEditPayloadPath(payload *pendingPatchEditPayload, path string) {
if payload == nil {
return
}
payload.Args.Path = strings.TrimSpace(path)
}
func patchEditPayloadArgsJSON(payload pendingPatchEditPayload) ([]byte, error) {
if payload.ArgsDecodeFailed && len(payload.RawArgsJSON) > 0 {
return append([]byte(nil), payload.RawArgsJSON...), nil
}
return payload.Args.MarshalJSON()
}
func patchEditPayloadAsEditPayload(payload pendingPatchEditPayload) editResultPayload {
return editResultPayload{
BeforeContent: payload.BeforeContent,
AfterContent: payload.AfterContent,
DiffString: payload.DiffString,
LinesAdded: payload.LinesAdded,
LinesRemoved: payload.LinesRemoved,
Message: payload.Message,
}
}
func summarizePatchEditResult(path string, result *agentv1.EditResult) string {
success := result.GetSuccess()
if success == nil {
return summarizeEditResultWithoutPath(result)
}
encoded, err := json.Marshal(map[string]any{
"success": map[string]any{
"path": firstNonEmpty(success.GetPath(), path),
},
})
if err != nil {
return fmt.Sprintf(`{"success":{"path":%q}}`, firstNonEmpty(success.GetPath(), path))
}
return string(encoded)
}
func compactPatchEditHistoryEditResult(path string, result *agentv1.EditResult) *agentv1.EditResult {
success := result.GetSuccess()
if success == nil {
return editResultWithoutPath(result)
}
return &agentv1.EditResult{
Result: &agentv1.EditResult_Success{
Success: &agentv1.EditSuccess{
Path: firstNonEmpty(success.GetPath(), path),
},
},
}
}
func boundedPatchEditDiffString(diffString string) string {
return truncateProjectedReplayText(patchEditToolName, diffString, projectedPatchEditReplayLimit)
}
func (args patchEditArgs) MarshalJSON() ([]byte, error) {
payload := map[string]any{
"path": strings.TrimSpace(args.Path),
"old_string": args.OldString,
"new_string": args.NewString,
}
if args.ReplaceAllSet || args.ReplaceAll {
payload["replace_all"] = args.ReplaceAll
}
return json.Marshal(payload)
}
func (args *patchEditArgs) UnmarshalJSON(raw []byte) error {
var rawArgs map[string]any
if err := json.Unmarshal(raw, &rawArgs); err != nil {
return err
}
path, pathFound, pathValid := readJSONStringAny(rawArgs, "path", "file_path", "Path", "FilePath")
if pathFound && !pathValid {
return fmt.Errorf("PatchEdit path must be a string")
}
oldString, oldFound, oldValid := readJSONStringAny(rawArgs, "old_string", "oldString", "OldString")
if oldFound && !oldValid {
return fmt.Errorf("PatchEdit old_string must be a string")
}
newString, newFound, newValid := readJSONStringAny(rawArgs, "new_string", "newString", "NewString")
if newFound && !newValid {
return fmt.Errorf("PatchEdit new_string must be a string")
}
replaceAll, replaceAllFound, replaceAllValid := readJSONBoolAny(rawArgs, "replace_all", "replaceAll", "ReplaceAll")
if replaceAllFound && !replaceAllValid {
return fmt.Errorf("PatchEdit replace_all must be a boolean")
}
*args = patchEditArgs{
Path: path,
OldString: oldString,
NewString: newString,
NewStringSet: newFound,
ReplaceAll: replaceAll,
ReplaceAllSet: replaceAllFound,
}
return nil
}
func (payload pendingPatchEditPayload) MarshalJSON() ([]byte, error) {
type alias pendingPatchEditPayload
return json.Marshal(alias(payload))
}
func decodePendingPatchEditPayload(raw []byte) (pendingPatchEditPayload, error) {
var payload pendingPatchEditPayload
if err := json.Unmarshal(raw, &payload); err != nil {
return pendingPatchEditPayload{}, fmt.Errorf("decode patch edit pending payload failed: %w", err)
}
return payload, nil
}
func hiddenPatchEditControlError(message *agentv1.ExecClientControlMessage) string {
if message == nil {
return "patch edit operation failed"
}
switch item := message.GetMessage().(type) {
case *agentv1.ExecClientControlMessage_Throw:
return firstNonEmpty(strings.TrimSpace(item.Throw.GetError()), "patch edit operation failed")
case *agentv1.ExecClientControlMessage_StreamClose:
return "patch edit operation closed unexpectedly"
default:
return "patch edit operation failed"
}
}
@@ -0,0 +1,296 @@
package forwarder
import (
"os"
"path/filepath"
"strings"
"time"
"cursor/gen/agentv1"
)
type streamPathContext struct {
workspacePaths []string
terminalsFolder string
requestFileContents map[string]string
}
func updateStreamRequestContextData(stream *ActiveStream, requestContext *agentv1.RequestContext) {
if stream == nil {
return
}
var workspacePaths []string
var terminalsFolder string
var fileContents map[string]string
if requestContext != nil {
env := requestContext.GetEnv()
workspacePaths = compactWorkspacePaths(env.GetWorkspacePaths(), env.GetProjectFolder())
terminalsFolder = strings.TrimSpace(env.GetTerminalsFolder())
fileContents = cloneStringMap(requestContext.GetFileContents())
}
stream.mu.Lock()
stream.WorkspacePaths = workspacePaths
stream.TerminalsFolder = terminalsFolder
stream.RequestFileContents = fileContents
stream.UpdatedAt = time.Now().UTC()
stream.mu.Unlock()
}
func snapshotStreamPathContext(stream *ActiveStream) streamPathContext {
if stream == nil {
return streamPathContext{}
}
stream.mu.Lock()
defer stream.mu.Unlock()
return streamPathContext{
workspacePaths: append([]string(nil), stream.WorkspacePaths...),
terminalsFolder: strings.TrimSpace(stream.TerminalsFolder),
requestFileContents: cloneStringMap(stream.RequestFileContents),
}
}
func compactWorkspacePaths(workspacePaths []string, projectFolder string) []string {
items := append([]string(nil), workspacePaths...)
if trimmedProjectFolder := strings.TrimSpace(projectFolder); trimmedProjectFolder != "" {
items = append(items, trimmedProjectFolder)
}
seen := make(map[string]struct{}, len(items))
result := make([]string, 0, len(items))
for _, item := range items {
trimmed := strings.TrimSpace(item)
if trimmed == "" {
continue
}
cleaned := filepath.Clean(trimmed)
if _, ok := seen[cleaned]; ok {
continue
}
seen[cleaned] = struct{}{}
result = append(result, cleaned)
}
return result
}
func cloneStringMap(input map[string]string) map[string]string {
if len(input) == 0 {
return nil
}
cloned := make(map[string]string, len(input))
for key, value := range input {
trimmedKey := strings.TrimSpace(key)
if trimmedKey == "" {
continue
}
cloned[trimmedKey] = value
}
if len(cloned) == 0 {
return nil
}
return cloned
}
func resolveWorkspacePath(path string, workspaceRoots []string, requireExisting bool) (string, bool) {
trimmedPath := strings.TrimSpace(path)
if trimmedPath == "" {
return "", false
}
cleanedPath := filepath.Clean(trimmedPath)
if pathCandidateUsable(cleanedPath, requireExisting) {
return cleanedPath, true
}
if len(workspaceRoots) == 0 {
return "", false
}
if !filepath.IsAbs(cleanedPath) {
for _, workspaceRoot := range workspaceRoots {
candidate := filepath.Join(workspaceRoot, cleanedPath)
if pathCandidateUsable(candidate, requireExisting) {
return filepath.Clean(candidate), true
}
}
return "", false
}
pathParts := splitPathParts(cleanedPath)
if len(pathParts) == 0 {
return "", false
}
for _, workspaceRoot := range workspaceRoots {
rootBase := strings.TrimSpace(filepath.Base(workspaceRoot))
if rootBase == "" || rootBase == "." || rootBase == string(filepath.Separator) {
continue
}
if index := lastIndexFold(pathParts, rootBase); index >= 0 {
candidate := workspaceRoot
if index+1 < len(pathParts) {
candidate = filepath.Join(append([]string{workspaceRoot}, pathParts[index+1:]...)...)
}
if pathCandidateUsable(candidate, requireExisting) {
return filepath.Clean(candidate), true
}
}
}
for suffixLen := len(pathParts); suffixLen >= 1; suffixLen-- {
suffixParts := append([]string(nil), pathParts[len(pathParts)-suffixLen:]...)
for _, workspaceRoot := range workspaceRoots {
candidate := joinWorkspaceRootWithSuffix(workspaceRoot, suffixParts)
if pathCandidateUsable(candidate, requireExisting) {
return filepath.Clean(candidate), true
}
}
}
return "", false
}
func isAbsoluteToolPath(path string) bool {
trimmed := strings.TrimSpace(path)
if trimmed == "" {
return false
}
if strings.HasPrefix(trimmed, "/") {
return true
}
if strings.HasPrefix(trimmed, `\\`) {
return true
}
return len(trimmed) >= 3 && isASCIIAlpha(trimmed[0]) && trimmed[1] == ':' && isPathSeparator(trimmed[2])
}
func isASCIIAlpha(value byte) bool {
return (value >= 'A' && value <= 'Z') || (value >= 'a' && value <= 'z')
}
func isPathSeparator(value byte) bool {
return value == '\\' || value == '/'
}
func lookupRequestFileContents(context streamPathContext, originalPath string, resolvedPath string) (string, bool) {
if len(context.requestFileContents) == 0 {
return "", false
}
aliases := buildPathAliases(originalPath, resolvedPath, context.workspacePaths)
for _, alias := range aliases {
if value, ok := context.requestFileContents[alias]; ok {
return value, true
}
}
return "", false
}
func buildPathAliases(originalPath string, resolvedPath string, workspaceRoots []string) []string {
seen := make(map[string]struct{}, 16)
aliases := make([]string, 0, 16)
addAlias := func(value string) {
trimmed := strings.TrimSpace(value)
if trimmed == "" {
return
}
if _, ok := seen[trimmed]; ok {
return
}
seen[trimmed] = struct{}{}
aliases = append(aliases, trimmed)
}
for _, candidate := range []string{originalPath, resolvedPath} {
trimmed := strings.TrimSpace(candidate)
if trimmed == "" {
continue
}
addAlias(trimmed)
addAlias(filepath.Clean(trimmed))
addAlias(filepath.ToSlash(filepath.Clean(trimmed)))
}
for _, candidate := range []string{originalPath, resolvedPath} {
cleaned := filepath.Clean(strings.TrimSpace(candidate))
if cleaned == "" {
continue
}
for _, workspaceRoot := range workspaceRoots {
if rel, err := filepath.Rel(workspaceRoot, cleaned); err == nil && rel != "." && !strings.HasPrefix(rel, ".."+string(filepath.Separator)) && rel != ".." {
addAlias(rel)
addAlias(filepath.ToSlash(rel))
}
}
}
return aliases
}
func splitPathParts(path string) []string {
cleaned := filepath.Clean(strings.TrimSpace(path))
if cleaned == "" {
return nil
}
volumeName := filepath.VolumeName(cleaned)
if volumeName != "" {
cleaned = strings.TrimPrefix(cleaned, volumeName)
}
cleaned = strings.TrimPrefix(cleaned, string(filepath.Separator))
if cleaned == "" {
return nil
}
parts := strings.Split(cleaned, string(filepath.Separator))
result := make([]string, 0, len(parts))
for _, part := range parts {
trimmed := strings.TrimSpace(part)
if trimmed != "" {
result = append(result, trimmed)
}
}
return result
}
func lastIndexFold(items []string, target string) int {
for index := len(items) - 1; index >= 0; index-- {
if strings.EqualFold(strings.TrimSpace(items[index]), strings.TrimSpace(target)) {
return index
}
}
return -1
}
func joinWorkspaceRootWithSuffix(workspaceRoot string, suffixParts []string) string {
root := filepath.Clean(strings.TrimSpace(workspaceRoot))
if root == "" {
return ""
}
parts := append([]string(nil), suffixParts...)
if len(parts) > 0 && strings.EqualFold(parts[0], filepath.Base(root)) {
parts = parts[1:]
}
if len(parts) == 0 {
return root
}
return filepath.Join(append([]string{root}, parts...)...)
}
func pathCandidateUsable(path string, requireExisting bool) bool {
cleaned := filepath.Clean(strings.TrimSpace(path))
if cleaned == "" {
return false
}
info, err := os.Stat(cleaned)
if err == nil {
return !info.IsDir() || requireExisting
}
if requireExisting {
return false
}
parent := filepath.Dir(cleaned)
if parent == "" || parent == "." || parent == cleaned {
return false
}
parentInfo, parentErr := os.Stat(parent)
return parentErr == nil && parentInfo.IsDir()
}
@@ -0,0 +1,37 @@
package forwarder
import "testing"
func TestIsAbsoluteToolPath(t *testing.T) {
tests := []struct {
name string
path string
want bool
}{
{name: "empty", path: "", want: false},
{name: "spaces", path: " ", want: false},
{name: "relative", path: "internal/backend/forwarder/path_resolution.go", want: false},
{name: "dot relative", path: "./path_resolution.go", want: false},
{name: "parent relative", path: "../path_resolution.go", want: false},
{name: "windows drive relative", path: `C:path_resolution.go`, want: false},
{name: "windows drive only", path: `C:`, want: false},
{name: "invalid drive prefix", path: `1:/path_resolution.go`, want: false},
{name: "posix absolute", path: "/Users/example/path_resolution.go", want: true},
{name: "posix absolute with spaces", path: " /tmp/path_resolution.go ", want: true},
{name: "windows backslash absolute", path: `C:\Users\example\path_resolution.go`, want: true},
{name: "windows slash absolute", path: `C:/Users/example/path_resolution.go`, want: true},
{name: "windows unc backslash", path: `\\server\share\path_resolution.go`, want: true},
{name: "windows unc slash", path: `//server/share/path_resolution.go`, want: true},
{name: "windows extended drive", path: `\\?\C:\Users\example\path_resolution.go`, want: true},
{name: "windows extended unc", path: `\\?\UNC\server\share\path_resolution.go`, want: true},
{name: "windows device path", path: `\\.\C:\Users\example\path_resolution.go`, want: true},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
if got := isAbsoluteToolPath(test.path); got != test.want {
t.Fatalf("isAbsoluteToolPath(%q) = %v, want %v", test.path, got, test.want)
}
})
}
}
@@ -0,0 +1,98 @@
package forwarder
import "strings"
// prepareWriteContentsForClient applies a narrow Windows compatibility shim for
// whole-file Write. Some Windows clients write CRLF payloads as CRCRLF on disk,
// so we send LF-only content when the target path is clearly a Windows path and
// the logical contents already use CRLF. The client's text-mode write then
// expands LF back to the intended CRLF on disk.
func prepareWriteContentsForClient(path string, contents string) string {
if !looksLikeWindowsPath(path) || !strings.Contains(contents, "\r\n") {
return contents
}
return normalizeCRLFToLF(contents)
}
// preparePatchEditWriteContentsForClient sends LF-only transport for PatchEdit.
// The client write path may expand LF to the platform newline, so sending CRLF
// can become CRCRLF on disk.
func preparePatchEditWriteContentsForClient(contents string) string {
if !strings.Contains(contents, "\r") {
return contents
}
return normalizeObservedLineEndingsToLF(contents)
}
// reconcilePostWriteObservedContent only corrects one known Windows anomaly:
// every line ending in the expected content was duplicated during write-back
// verification, which makes the file appear to contain a blank line between
// each real line. Some read paths normalize line endings to LF, so the
// comparison also checks the normalized representation.
func reconcilePostWriteObservedContent(expected string, observed string) (string, bool) {
if observed == expected {
return observed, false
}
if observed == systematicallyDoubledLineEndings(expected) {
return expected, true
}
expectedNormalized := normalizeObservedLineEndingsToLF(expected)
observedNormalized := normalizeObservedLineEndingsToLF(observed)
if observedNormalized == expectedNormalized {
return expected, true
}
if observedNormalized == systematicallyDoubledLineEndings(expectedNormalized) {
return expected, true
}
return observed, false
}
func systematicallyDoubledLineEndings(content string) string {
if content == "" {
return ""
}
var builder strings.Builder
builder.Grow(len(content) * 2)
for index := 0; index < len(content); {
switch {
case content[index] == '\r' && index+1 < len(content) && content[index+1] == '\n':
builder.WriteString("\r\n\r\n")
index += 2
case content[index] == '\n':
builder.WriteString("\n\n")
index++
default:
builder.WriteByte(content[index])
index++
}
}
return builder.String()
}
func normalizeObservedLineEndingsToLF(content string) string {
if content == "" {
return ""
}
normalized := normalizeCRLFToLF(content)
return strings.ReplaceAll(normalized, "\r", "\n")
}
func normalizeCRLFToLF(content string) string {
if content == "" || !strings.Contains(content, "\r\n") {
return content
}
return strings.ReplaceAll(content, "\r\n", "\n")
}
func looksLikeWindowsPath(path string) bool {
trimmed := strings.TrimSpace(path)
if len(trimmed) >= 3 && isASCIILetter(trimmed[0]) && trimmed[1] == ':' {
return trimmed[2] == '\\' || trimmed[2] == '/'
}
return strings.HasPrefix(trimmed, `\\`)
}
func isASCIILetter(value byte) bool {
return (value >= 'a' && value <= 'z') || (value >= 'A' && value <= 'Z')
}
@@ -0,0 +1,18 @@
//go:build !windows
package forwarder
import (
"errors"
"os"
"syscall"
)
func processExists(pid int) bool {
process, err := os.FindProcess(pid)
if err != nil || process == nil {
return false
}
err = process.Signal(syscall.Signal(0))
return err == nil || errors.Is(err, syscall.EPERM)
}
@@ -0,0 +1,18 @@
//go:build windows
package forwarder
import (
"errors"
"golang.org/x/sys/windows"
)
func processExists(pid int) bool {
handle, err := windows.OpenProcess(windows.PROCESS_QUERY_LIMITED_INFORMATION, false, uint32(pid))
if err == nil {
_ = windows.CloseHandle(handle)
return true
}
return errors.Is(err, windows.ERROR_ACCESS_DENIED)
}
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,130 @@
package forwarder
import (
"crypto/sha256"
"encoding/hex"
"encoding/json"
"strings"
modeladapter "cursor/internal/backend/agent/model"
)
func newPromptContextMessage(source string, message modeladapter.Message, persist bool) PromptContextMessage {
context := PromptContextMessage{
Source: strings.TrimSpace(source),
Message: message,
Persist: persist,
}
context.ContentHash = promptContextContentHash(context.Message)
return context
}
func normalizePromptContextMessage(context PromptContextMessage) PromptContextMessage {
context.Source = strings.TrimSpace(context.Source)
context.Message.Role = strings.TrimSpace(context.Message.Role)
context.Message.Content = strings.TrimSpace(context.Message.Content)
context.ContentHash = strings.TrimSpace(context.ContentHash)
if context.ContentHash == "" {
context.ContentHash = promptContextContentHash(context.Message)
}
return context
}
func promptContextContentHash(message modeladapter.Message) string {
sum := sha256.Sum256([]byte(strings.TrimSpace(message.Role) + "\x00" + strings.TrimSpace(message.Content)))
return hex.EncodeToString(sum[:])
}
func filterCurrentTurnPromptContexts(conversation *ConversationFile, contexts []PromptContextMessage) []PromptContextMessage {
if len(contexts) == 0 {
return nil
}
seen := collectCurrentTurnPromptContextKeys(conversation)
filtered := make([]PromptContextMessage, 0, len(contexts))
for _, context := range contexts {
context = normalizePromptContextMessage(context)
if !isReplayablePromptContext(context) {
continue
}
if context.Persist {
if _, ok := seen[promptContextKey(context)]; ok {
continue
}
}
filtered = append(filtered, context)
}
return filtered
}
func collectCurrentTurnPromptContextKeys(conversation *ConversationFile) map[string]struct{} {
keys := make(map[string]struct{})
if conversation == nil {
return keys
}
currentTurnSeq := conversation.NextTurnSeq - 1
if currentTurnSeq <= 0 {
return keys
}
for _, entry := range conversation.Entries {
if entry.TurnSeq != currentTurnSeq || strings.TrimSpace(entry.Kind) != "prompt_context" {
continue
}
var payload promptContextEntryPayload
if err := json.Unmarshal(entry.Payload, &payload); err != nil {
continue
}
context := normalizePromptContextMessage(PromptContextMessage{
Source: payload.Source,
ContentHash: payload.ContentHash,
Message: modeladapter.Message{
Role: payload.Role,
Content: payload.Content,
},
Persist: true,
})
if !isReplayablePromptContext(context) {
continue
}
keys[promptContextKey(context)] = struct{}{}
}
return keys
}
func promptContextKey(context PromptContextMessage) string {
context = normalizePromptContextMessage(context)
return context.Source + "\x00" + context.ContentHash
}
func isReplayablePromptContext(context PromptContextMessage) bool {
context = normalizePromptContextMessage(context)
if context.Source == "" || strings.TrimSpace(context.Message.Content) == "" {
return false
}
if context.Message.Role == "" {
context.Message.Role = "user"
}
if context.Message.Role != "user" && context.Message.Role != "system" {
return false
}
if strings.TrimSpace(context.Message.Name) != "" || strings.TrimSpace(context.Message.ToolCallID) != "" {
return false
}
return len(context.Message.ToolCalls) == 0 && len(context.Message.ContentParts) == 0
}
func newPromptContextEntry(turnSeq int64, requestID string, context PromptContextMessage) HistoryEntry {
context = normalizePromptContextMessage(context)
payload, _ := json.Marshal(promptContextEntryPayload{
Source: context.Source,
Role: firstNonEmpty(strings.TrimSpace(context.Message.Role), "user"),
Content: strings.TrimSpace(context.Message.Content),
ContentHash: context.ContentHash,
})
return HistoryEntry{
TurnSeq: turnSeq,
RequestID: strings.TrimSpace(requestID),
Role: "system",
Kind: "prompt_context",
Payload: payload,
}
}
+323
View File
@@ -0,0 +1,323 @@
package forwarder
import (
"fmt"
"sort"
"strings"
"unicode/utf8"
"google.golang.org/protobuf/proto"
"cursor/gen/agentv1"
)
const (
promptGuardUserTextChars = 32000
promptGuardUserRichTextChars = 32000
promptGuardSubagentReminderChars = 12000
promptGuardSelectedFileChars = 16000
promptGuardSelectedFilesTotalChars = 64000
promptGuardSelectedFilesMaxCount = 12
promptGuardRequestFileChars = 16000
promptGuardRequestFilesTotalChars = 64000
promptGuardRequestFilesMaxCount = 12
promptGuardRuleChars = 6000
promptGuardRulesTotalChars = 24000
promptGuardRulesMaxCount = 40
promptGuardSkillDescriptionChars = 800
promptGuardSkillDescriptionsTotal = 16000
promptGuardSkillDescriptorsMaxCount = 32
promptGuardAgentSkillContentChars = 6000
promptGuardAgentSkillsMaxCount = 16
promptGuardRealtimeTextChars = 12000
promptGuardCompiledMessageChars = 120000
)
func normalizeUserMessageForStorage(userMessage *agentv1.UserMessage) *agentv1.UserMessage {
if userMessage == nil {
return nil
}
cloned, ok := proto.Clone(userMessage).(*agentv1.UserMessage)
if !ok || cloned == nil {
return userMessage
}
cloned.Text = truncatePromptGuardText("user_message.text", cloned.GetText(), promptGuardUserTextChars)
if richText := strings.TrimSpace(cloned.GetRichText()); richText != "" {
cloned.RichText = stringPtr(truncatePromptGuardText("user_message.rich_text", richText, promptGuardUserRichTextChars))
}
if reminder := strings.TrimSpace(cloned.GetSubagentSystemReminder()); reminder != "" {
cloned.SubagentSystemReminder = stringPtr(truncatePromptGuardText("user_message.subagent_system_reminder", reminder, promptGuardSubagentReminderChars))
}
cloned.SelectedContext = guardSelectedContext(cloned.GetSelectedContext())
return cloned
}
func guardRequestContextForStorage(requestContext *agentv1.RequestContext) {
if requestContext == nil {
return
}
requestContext.Rules = guardCursorRules(requestContext.GetRules())
requestContext.FileContents = guardStringMap(
requestContext.GetFileContents(),
"request_context.file_contents",
promptGuardRequestFileChars,
promptGuardRequestFilesTotalChars,
promptGuardRequestFilesMaxCount,
)
if summary := strings.TrimSpace(requestContext.GetUserIntentSummary()); summary != "" {
requestContext.UserIntentSummary = stringPtr(truncatePromptGuardText("request_context.user_intent_summary", summary, promptGuardRealtimeTextChars))
}
if hooks := strings.TrimSpace(requestContext.GetHooksAdditionalContext()); hooks != "" {
requestContext.HooksAdditionalContext = stringPtr(truncatePromptGuardText("request_context.hooks_additional_context", hooks, promptGuardRealtimeTextChars))
}
if commit := strings.TrimSpace(requestContext.GetCommitAttributionMessage()); commit != "" {
requestContext.CommitAttributionMessage = stringPtr(truncatePromptGuardText("request_context.commit_attribution_message", commit, promptGuardRealtimeTextChars))
}
if pr := strings.TrimSpace(requestContext.GetPrAttributionMessage()); pr != "" {
requestContext.PrAttributionMessage = stringPtr(truncatePromptGuardText("request_context.pr_attribution_message", pr, promptGuardRealtimeTextChars))
}
}
func guardCompiledConversationForProvider(compiled CompiledConversation) CompiledConversation {
for index := range compiled.Messages {
message := &compiled.Messages[index]
if strings.TrimSpace(message.Role) == "system" {
continue
}
if strings.TrimSpace(message.Content) != "" {
message.Content = truncatePromptGuardText("compiled."+firstNonEmpty(strings.TrimSpace(message.Role), "message"), message.Content, promptGuardCompiledMessageChars)
}
for partIndex := range message.ContentParts {
if strings.TrimSpace(message.ContentParts[partIndex].Text) == "" {
continue
}
message.ContentParts[partIndex].Text = truncatePromptGuardText("compiled.content_part", message.ContentParts[partIndex].Text, promptGuardCompiledMessageChars)
}
}
return compiled
}
func guardSelectedContext(selectedContext *agentv1.SelectedContext) *agentv1.SelectedContext {
if selectedContext == nil {
return nil
}
cloned, ok := proto.Clone(selectedContext).(*agentv1.SelectedContext)
if !ok || cloned == nil {
return selectedContext
}
cloned.Files = guardSelectedFiles(cloned.GetFiles())
cloned.SelectedSkills = guardAgentSkills(cloned.GetSelectedSkills())
cloned.ExtraContext = guardStringSlice(cloned.GetExtraContext(), "selected_context.extra_context", promptGuardRealtimeTextChars, promptGuardRealtimeTextChars, promptGuardAgentSkillsMaxCount)
return cloned
}
func guardSelectedFiles(files []*agentv1.SelectedFile) []*agentv1.SelectedFile {
if len(files) == 0 {
return nil
}
result := make([]*agentv1.SelectedFile, 0, minInt(len(files), promptGuardSelectedFilesMaxCount))
remaining := promptGuardSelectedFilesTotalChars
for _, file := range files {
if file == nil || len(result) >= promptGuardSelectedFilesMaxCount {
continue
}
content := strings.TrimSpace(file.GetContent())
if content == "" {
continue
}
limit := minInt(promptGuardSelectedFileChars, remaining)
if limit <= 0 {
break
}
cloned, ok := proto.Clone(file).(*agentv1.SelectedFile)
if !ok || cloned == nil {
continue
}
cloned.Content = truncatePromptGuardText("selected_context.files.content", content, limit)
remaining -= promptGuardRuneCount(cloned.GetContent())
result = append(result, cloned)
}
return result
}
func guardCursorRules(rules []*agentv1.CursorRule) []*agentv1.CursorRule {
if len(rules) == 0 {
return nil
}
result := make([]*agentv1.CursorRule, 0, minInt(len(rules), promptGuardRulesMaxCount))
remaining := promptGuardRulesTotalChars
for _, rule := range rules {
if rule == nil || len(result) >= promptGuardRulesMaxCount {
continue
}
content := strings.TrimSpace(rule.GetContent())
if content == "" {
continue
}
limit := minInt(promptGuardRuleChars, remaining)
if limit <= 0 {
break
}
cloned, ok := proto.Clone(rule).(*agentv1.CursorRule)
if !ok || cloned == nil {
continue
}
cloned.Content = truncatePromptGuardText("request_context.rules.content", content, limit)
remaining -= promptGuardRuneCount(cloned.GetContent())
result = append(result, cloned)
}
return result
}
func guardSkillDescriptors(descriptors []*agentv1.SkillDescriptor) []*agentv1.SkillDescriptor {
if len(descriptors) == 0 {
return nil
}
result := make([]*agentv1.SkillDescriptor, 0, minInt(len(descriptors), promptGuardSkillDescriptorsMaxCount))
remaining := promptGuardSkillDescriptionsTotal
for _, descriptor := range descriptors {
if descriptor == nil || len(result) >= promptGuardSkillDescriptorsMaxCount {
continue
}
description := strings.TrimSpace(descriptor.GetDescription())
if description == "" {
continue
}
limit := minInt(promptGuardSkillDescriptionChars, remaining)
if limit <= 0 {
break
}
cloned, ok := proto.Clone(descriptor).(*agentv1.SkillDescriptor)
if !ok || cloned == nil {
continue
}
cloned.Description = truncatePromptGuardText("skill.description", description, limit)
remaining -= promptGuardRuneCount(cloned.GetDescription())
result = append(result, cloned)
}
return result
}
func guardAgentSkills(skills []*agentv1.AgentSkill) []*agentv1.AgentSkill {
if len(skills) == 0 {
return nil
}
result := make([]*agentv1.AgentSkill, 0, minInt(len(skills), promptGuardAgentSkillsMaxCount))
for _, skill := range skills {
if skill == nil || len(result) >= promptGuardAgentSkillsMaxCount {
continue
}
cloned, ok := proto.Clone(skill).(*agentv1.AgentSkill)
if !ok || cloned == nil {
continue
}
if content := strings.TrimSpace(cloned.GetContent()); content != "" {
cloned.Content = truncatePromptGuardText("agent_skill.content", content, promptGuardAgentSkillContentChars)
}
if description := strings.TrimSpace(cloned.GetDescription()); description != "" {
cloned.Description = truncatePromptGuardText("agent_skill.description", description, promptGuardSkillDescriptionChars)
}
result = append(result, cloned)
}
return result
}
func guardStringMap(input map[string]string, label string, itemLimit int, totalLimit int, maxItems int) map[string]string {
if len(input) == 0 {
return nil
}
keys := make([]string, 0, len(input))
values := make(map[string]string, len(input))
for key, value := range input {
trimmedKey := strings.TrimSpace(key)
trimmedValue := strings.TrimSpace(value)
if trimmedKey == "" || trimmedValue == "" {
continue
}
if _, exists := values[trimmedKey]; exists {
continue
}
keys = append(keys, trimmedKey)
values[trimmedKey] = trimmedValue
}
sort.Strings(keys)
result := make(map[string]string, minInt(len(keys), maxItems))
remaining := totalLimit
for _, key := range keys {
if len(result) >= maxItems {
break
}
content := values[key]
if content == "" {
continue
}
limit := minInt(itemLimit, remaining)
if limit <= 0 {
break
}
result[key] = truncatePromptGuardText(label, content, limit)
remaining -= promptGuardRuneCount(result[key])
}
if len(result) == 0 {
return nil
}
return result
}
func guardStringSlice(input []string, label string, itemLimit int, totalLimit int, maxItems int) []string {
if len(input) == 0 {
return nil
}
result := make([]string, 0, minInt(len(input), maxItems))
remaining := totalLimit
for _, item := range input {
if len(result) >= maxItems {
break
}
content := strings.TrimSpace(item)
if content == "" {
continue
}
limit := minInt(itemLimit, remaining)
if limit <= 0 {
break
}
truncated := truncatePromptGuardText(label, content, limit)
remaining -= promptGuardRuneCount(truncated)
result = append(result, truncated)
}
return result
}
func truncatePromptGuardText(label string, text string, limit int) string {
if limit <= 0 {
return ""
}
if promptGuardRuneCount(text) <= limit {
return text
}
runes := []rune(text)
notice := fmt.Sprintf("\n\n[truncated: %s exceeded %d chars; kept head and tail from %d chars]\n\n", strings.TrimSpace(label), limit, len(runes))
noticeRunes := []rune(notice)
keep := limit - len(noticeRunes)
if keep <= 0 {
return string(runes[:limit])
}
head := keep * 2 / 3
tail := keep - head
if tail <= 0 {
return string(runes[:head]) + notice
}
return string(runes[:head]) + notice + string(runes[len(runes)-tail:])
}
func promptGuardRuneCount(text string) int {
return utf8.RuneCountInString(text)
}
func minInt(a int, b int) int {
if a < b {
return a
}
return b
}
+62
View File
@@ -0,0 +1,62 @@
// provider.go 把 forwarder 的 canonical 请求转交给现有的 provider adapter 层。
package forwarder
import (
"context"
"encoding/json"
"strings"
modeladapter "cursor/internal/backend/agent/model"
)
type DefaultProviderGateway struct {
router modeladapter.ModelAdapterRouter
}
// NewProviderGateway 创建默认 provider 网关。
func NewProviderGateway(resolver modeladapter.ChannelResolver) *DefaultProviderGateway {
return &DefaultProviderGateway{
router: modeladapter.NewRouter(resolver),
}
}
// StartStream 把 forwarder 的 provider 请求翻译成 modeladapter.StreamRequest 并发起流式调用。
func (gateway *DefaultProviderGateway) StartStream(ctx context.Context, req ProviderRequest, sink func(modeladapter.ModelEvent) error) error {
if ctx == nil {
ctx = context.Background()
}
requestKnobs := make(map[string]any, len(req.RequestKnobs)+2)
for key, value := range req.RequestKnobs {
requestKnobs[key] = value
}
requestKnobs["stream"] = true
if req.MaxTokens > 0 {
requestKnobs["max_tokens"] = req.MaxTokens
}
if strings.TrimSpace(req.ThinkingEffort) != "" {
requestKnobs["runtime_thinking_effort"] = strings.TrimSpace(req.ThinkingEffort)
}
err := gateway.router.Stream(ctx, modeladapter.StreamRequest{
RequestID: req.RequestID,
RunID: req.RunID,
ModelCallID: req.ModelCallID,
ConversationID: req.ConversationID,
Mode: req.Mode,
ModelID: req.ModelID,
ThinkingEffort: req.ThinkingEffort,
Messages: req.Messages,
StableMessageCount: req.StableMessageCount,
Tools: append([]json.RawMessage(nil), req.Tools...),
MaxTokens: req.MaxTokens,
Stream: true,
RequestKnobs: requestKnobs,
CompileSummary: req.CompileSummary,
Observer: req.Observer,
ArtifactPaths: req.ArtifactPaths,
RequestBodyOverride: req.RequestBodyOverride,
}, sink)
if err != nil {
return providerTerminalError{cause: err}
}
return nil
}
@@ -0,0 +1,16 @@
// reasoning_metadata.go 保存 history/replay 共用的 reasoning metadata 判定。
package forwarder
import (
"strings"
modeladapter "cursor/internal/backend/agent/model"
)
func hasReplayableReasoningPayload(reasoningContent string, reasoningSignature string, reasoningSignatureSource string) bool {
if strings.TrimSpace(reasoningContent) != "" {
return true
}
return strings.TrimSpace(reasoningSignature) != "" &&
strings.TrimSpace(reasoningSignatureSource) == modeladapter.ReasoningSignatureSourceOpenAIResponses
}
+457
View File
@@ -0,0 +1,457 @@
// reminders.go 负责按 mode 和上下文生成最小的 system reminder 集合。
package forwarder
import (
"encoding/json"
"fmt"
"os"
"path/filepath"
"strings"
"cursor/gen/agentv1"
modeladapter "cursor/internal/backend/agent/model"
promptassets "cursor/prompt"
)
type DefaultReminderInjector struct {
}
const (
promptContextSourcePlanTurnContract = "plan_turn_contract"
promptContextSourceActiveModeContract = "active_mode_contract"
promptContextSourceLatestUserIntent = "latest_user_intent"
promptContextSourceCurrentUserRequest = "current_user_request"
promptContextSourceSubagentContract = "subagent_contract"
promptContextSourceSubagentEmptyStopRecovery = "subagent_empty_stop_recovery"
promptContextSourceDebugModeReminder = "debug_mode_reminder"
)
// NewReminderInjector 创建默认 reminder 注入器。
func NewReminderInjector() *DefaultReminderInjector {
return &DefaultReminderInjector{}
}
// Inject 根据 mode、最近用户输入和工具上下文生成本轮附加提醒。
func (injector *DefaultReminderInjector) Inject(mode agentv1.AgentMode, conversation *ConversationFile, replayMessages []modeladapter.Message, latestUserText string, toolNames []string) PromptReminders {
_ = toolNames
reminders := make([]string, 0, 6)
normalizedMode, err := validateSupportedActiveMode(mode)
if err != nil {
normalizedMode = agentv1.AgentMode_AGENT_MODE_AGENT
}
if conversation != nil && isChildConversationSubagentTypeName(conversation.SubagentTypeName) {
return appendCurrentTurnAttentionReminders(PromptReminders{
SystemParts: reminders,
PromptContexts: []PromptContextMessage{
newPromptContextReminder(promptContextSourceSubagentContract, subagentContractText()),
newPromptContextReminder(promptContextSourceActiveModeContract, currentModeContractText(normalizedMode, true)),
},
}, latestUserText)
}
if normalizedMode == agentv1.AgentMode_AGENT_MODE_DEBUG {
return debugModePromptReminders(conversation)
}
reminders = append(reminders, "If multiple <current_plan> or <todo_list> blocks appear in the conversation, treat the last block of each type as the current source of truth.")
switch normalizedMode {
case agentv1.AgentMode_AGENT_MODE_ASK:
reminders = append(reminders, "You are in ask mode. Prefer direct answers and only use tools when they are necessary to answer accurately.")
reminders = append(reminders, "Lead with the conclusion, keep the response concise, and avoid unsolicited example code or long bullet lists.")
case agentv1.AgentMode_AGENT_MODE_PLAN:
reminders = append(reminders, "You are currently working in plan mode for the user. Prioritize investigation, decomposition, tradeoff analysis, and producing or refining a concrete plan.")
reminders = append(reminders, "Do not directly modify files in plan mode. Avoid direct file-editing tools such as Write, Delete, and PatchEdit.")
reminders = append(reminders, "For non-trivial plan-mode work, first do a quick reconnaissance yourself, then launch 2-4 parallel Task subagents with subagent_type=\"explore\" to investigate distinct angles before CreatePlan. Avoid using exactly one subagent for broad tasks: either handle narrow tasks directly, or split broad tasks into multiple independent investigations. Synthesize the subagent results yourself; do not delegate the final plan.")
reminders = append(reminders, "For narrow, well-scoped tasks, you can investigate directly and then create a lean plan with only the essential stages, tradeoffs, and next steps.")
if hasCurrentPlan(conversation) {
reminders = append(reminders, "A current plan already exists. Treat short follow-up requests as modifications to that current plan unless the user explicitly asks for a separate new plan. When calling CreatePlan for an existing plan, send the complete revised plan, preserve relevant existing content, incorporate the user's requested changes, and omit the name field. The CreatePlan name field is only allowed on the first CreatePlan call; never use a later name to rename or create a separate plan.")
}
case agentv1.AgentMode_AGENT_MODE_MULTITASK:
reminders = append(reminders, "You are in multitask mode. Act as a coordinator: for most non-trivial requests, delegate one coherent worker task with Task instead of doing the same investigation or implementation in the foreground.")
reminders = append(reminders, "After delegating the only coherent worker task for a request, do not continue the same work in the foreground. Only do distinct coordination work, answer a new independent question, or synthesize after multiple workers return.")
reminders = append(reminders, "Do not wait, sleep, or poll just for a running worker to complete. End the response unless there is separate useful coordination to do.")
reminders = append(reminders, "Do not over-decompose small or medium tasks into many sibling workers. Use multiple sibling workers only for clearly independent top-level workstreams.")
default:
reminders = append(reminders, "You are in agent mode. Use the available tools when they materially improve correctness or efficiency.")
reminders = append(reminders, "When reporting progress or completion, lead with the result, mention only key changes or verification, and avoid long recaps, exhaustive lists, or unsolicited example code.")
}
result := PromptReminders{
SystemParts: reminders,
PromptContexts: []PromptContextMessage{
newPromptContextReminder(promptContextSourceActiveModeContract, currentModeContractText(normalizedMode, false)),
},
}
if normalizedMode == agentv1.AgentMode_AGENT_MODE_PLAN {
if reminder := strings.TrimSpace(promptassets.MustReadPlanSystemReminder()); reminder != "" {
result.PromptContexts = append([]PromptContextMessage{
newPromptContextReminder(promptContextSourcePlanTurnContract, reminder),
}, result.PromptContexts...)
}
return appendCurrentTurnAttentionReminders(result, latestUserText)
}
if normalizedMode != agentv1.AgentMode_AGENT_MODE_AGENT {
return appendCurrentTurnAttentionReminders(result, latestUserText)
}
candidate, ok := extractLatestSuccessfulEditReminder(replayMessages)
if !ok {
return appendCurrentTurnAttentionReminders(result, latestUserText)
}
result.PromptContexts = append(result.PromptContexts, newPromptContextMessage(
"latest_edit_reminder",
modeladapter.Message{
Role: "user",
Content: strings.TrimSpace(fmt.Sprintf(`<system_reminder>
You recently successfully edited %q.
For this file, the latest source of truth is the most recent successful %s, not earlier reads or memory.
When modifying this file:
- use PatchEdit with path, old_string, new_string, and optional replace_all
- copy old_string exactly from the latest file content; line endings are not normalized or treated equivalently during matching
- replace_all defaults to false, so old_string must match exactly one occurrence unless you intentionally set replace_all to true
- new_string may be empty to delete old_string
- preserve spaces, tabs, indentation, punctuation, and line endings exactly in old_string
- only read the file again yourself if you need exact current content or extra context to choose a unique old_string
</system_reminder>`, candidate.Path, candidate.SourceField)),
},
false,
))
return appendCurrentTurnAttentionReminders(result, latestUserText)
}
func appendCurrentTurnAttentionReminders(result PromptReminders, latestUserText string) PromptReminders {
if reminder := latestUserIntentReminderText(latestUserText); reminder != "" {
result.PromptContexts = append(result.PromptContexts, newPromptContextReminder(promptContextSourceLatestUserIntent, reminder))
}
return appendCurrentUserRequestReminder(result, latestUserText)
}
func appendCurrentUserRequestReminder(result PromptReminders, latestUserText string) PromptReminders {
if strings.TrimSpace(latestUserText) == "" {
return result
}
result.PromptContexts = append(result.PromptContexts, newCurrentUserRequestReminder(latestUserText))
return result
}
func latestUserIntentReminderText(latestUserText string) string {
normalizedText := strings.ToLower(strings.TrimSpace(latestUserText))
switch {
case strings.Contains(normalizedText, "review"), strings.Contains(latestUserText, "评审"), strings.Contains(latestUserText, "审查"):
return "When reviewing code, focus on bugs, regressions, behavioral risks, and missing tests."
case strings.Contains(normalizedText, "plan"), strings.Contains(latestUserText, "计划"):
return "Prefer clear staged plans with concrete checkpoints."
default:
return ""
}
}
func newCurrentUserRequestReminder(latestUserText string) PromptContextMessage {
return newPromptContextMessage(
promptContextSourceCurrentUserRequest,
modeladapter.Message{
Role: "user",
Content: strings.TrimSpace(fmt.Sprintf(`<current_user_request>
%s
</current_user_request>
Handle this request. Use any latest tool results as evidence when present, and provide the requested answer once you have enough information. Treat surrounding system reminders as constraints, not as the user's task.`, strings.TrimSpace(latestUserText))),
},
false,
)
}
func newPromptContextReminder(source string, content string) PromptContextMessage {
return newPromptContextMessage(
source,
modeladapter.Message{
Role: "user",
Content: wrapSystemReminder(content),
},
false,
)
}
func subagentContractText() string {
return strings.Join([]string{
"The turn that contains this reminder runs inside a subagent child conversation. Work as an investigator for the parent agent, not as the final user-facing assistant.",
"Return a short textual result: lead with the conclusion, keep only the key evidence, and do not produce a long response.",
"Use the available agent tools when they materially improve correctness or efficiency. Do not ask the user questions. If required information is missing, report the gap to the parent agent instead of asking the user directly.",
}, "\n\n")
}
func currentModeContractText(mode agentv1.AgentMode, childSubagent bool) string {
if childSubagent {
return "For the turn that contains this reminder, the active mode is a subagent child conversation. Use the available agent tools, but do not call AskQuestion. Return only a concise investigation result for the parent agent."
}
switch normalizeMode(mode) {
case agentv1.AgentMode_AGENT_MODE_PLAN:
return "For the turn that contains this reminder, the active mode is plan. Do not modify files or system state. Use CreatePlan when the plan is ready or needs updating."
case agentv1.AgentMode_AGENT_MODE_ASK:
return "For the turn that contains this reminder, the active mode is ask. Prefer a direct answer. Use tools only when they materially improve accuracy, and do not call CreatePlan."
case agentv1.AgentMode_AGENT_MODE_DEBUG:
return "For the turn that contains this reminder, the active mode is debug. Follow the Debug Mode workflow from the static debug prompt: inspect or reproduce before editing, keep 3-5 concrete hypotheses, use the injected debug session log path when temporary instrumentation is useful, and verify with runtime evidence. Do not call CreatePlan or SwitchMode."
case agentv1.AgentMode_AGENT_MODE_MULTITASK:
return "For the turn that contains this reminder, the active mode is multitask. Act as the foreground coordinator: delegate most non-trivial work to a coherent worker with Task, avoid duplicating delegated work in the foreground, and do not wait just for a worker to finish."
default:
return "For the turn that contains this reminder, the active mode is agent. CreatePlan is not available in this mode; do not call CreatePlan. If the user explicitly asks to create or revise a plan, call SwitchMode to return to plan mode first. If there is an accepted or current plan, execute or continue the implementation using the available agent-mode tools."
}
}
func debugModePromptReminders(conversation *ConversationFile) PromptReminders {
initial := !hasPreviousPromptContextSource(conversation, promptContextSourceDebugModeReminder)
content := renderDebugSystemReminder(promptassets.MustReadDebugSystemReminder(initial), conversation)
if strings.TrimSpace(content) == "" {
return PromptReminders{}
}
return PromptReminders{
PromptContexts: []PromptContextMessage{
newPromptContextMessage(promptContextSourceDebugModeReminder, modeladapter.Message{
Role: "user",
Content: strings.TrimSpace(content),
}, false),
},
}
}
func hasPreviousPromptContextSource(conversation *ConversationFile, source string) bool {
if conversation == nil {
return false
}
needle := strings.TrimSpace(source)
if needle == "" {
return false
}
for _, entry := range conversation.Entries {
if strings.TrimSpace(entry.Kind) != "prompt_context" {
continue
}
var payload promptContextEntryPayload
if err := json.Unmarshal(entry.Payload, &payload); err != nil {
continue
}
if strings.TrimSpace(payload.Source) == needle {
return true
}
}
return false
}
func renderDebugSystemReminder(template string, conversation *ConversationFile) string {
sessionID := firstNonEmpty(debugSessionID(conversation), "debug")
logPath := debugLogPath(sessionID)
serverEndpoint := debugServerEndpoint(sessionID)
result := strings.TrimSpace(template)
for _, replacement := range []struct {
placeholder string
value string
}{
{placeholder: "{{DEBUG_SERVER_ENDPOINT}}", value: serverEndpoint},
{placeholder: "{{DEBUG_LOG_PATH}}", value: logPath},
{placeholder: "{{DEBUG_SESSION_ID}}", value: sessionID},
} {
result = strings.ReplaceAll(result, replacement.placeholder, replacement.value)
}
return result
}
func debugLogPath(sessionID string) string {
filename := fmt.Sprintf("debug-%s.log", firstNonEmpty(strings.TrimSpace(sessionID), "debug"))
cwd, err := os.Getwd()
if err != nil || strings.TrimSpace(cwd) == "" {
return filepath.Join(".cursor", filename)
}
return filepath.Join(cwd, ".cursor", filename)
}
func debugServerEndpoint(sessionID string) string {
return "http://127.0.0.1:7337/ingest/" + firstNonEmpty(strings.TrimSpace(sessionID), "debug")
}
func debugSessionID(conversation *ConversationFile) string {
if conversation == nil {
return ""
}
for _, candidate := range []string{
conversation.CurrentRequestID,
conversation.CurrentLoopID,
conversation.ConversationID,
} {
if normalized := normalizeDebugSessionID(candidate); normalized != "" {
return normalized
}
}
return ""
}
func normalizeDebugSessionID(raw string) string {
trimmed := strings.TrimSpace(raw)
if trimmed == "" {
return ""
}
var builder strings.Builder
for _, value := range trimmed {
switch {
case value >= 'a' && value <= 'z':
builder.WriteRune(value)
case value >= 'A' && value <= 'Z':
builder.WriteRune(value)
case value >= '0' && value <= '9':
builder.WriteRune(value)
case value == '-', value == '_':
builder.WriteRune(value)
default:
builder.WriteRune('-')
}
}
return strings.Trim(builder.String(), "-")
}
func wrapSystemReminder(content string) string {
trimmed := strings.TrimSpace(content)
if trimmed == "" {
return ""
}
if strings.HasPrefix(trimmed, "<system_reminder>") && strings.HasSuffix(trimmed, "</system_reminder>") {
return trimmed
}
return "<system_reminder>\n" + trimmed + "\n</system_reminder>"
}
func hasCurrentPlan(conversation *ConversationFile) bool {
if conversation == nil {
return false
}
state, err := projectConversationStructuredState(conversation)
if err != nil {
return false
}
return state.HasPlan || len(state.Plans) > 0
}
type editReminderCandidate struct {
ToolCallID string
Path string
SourceField string
}
type replayEditSuccess struct {
Success *struct {
Path string `json:"path"`
AfterFullFileContent string `json:"afterFullFileContent"`
AfterFullFileContentSnake string `json:"after_full_file_content"`
DiffString string `json:"diffString"`
DiffStringSnake string `json:"diff_string"`
} `json:"success"`
}
func extractLatestSuccessfulEditReminder(messages []modeladapter.Message) (editReminderCandidate, bool) {
if len(messages) == 0 {
return editReminderCandidate{}, false
}
toolPaths := collectReplayEditToolPaths(messages)
for index := len(messages) - 1; index >= 0; index-- {
message := messages[index]
if strings.TrimSpace(message.Role) != "tool" {
continue
}
toolName := strings.TrimSpace(message.Name)
if !isEditReminderToolName(toolName) {
continue
}
toolCallID := strings.TrimSpace(message.ToolCallID)
if toolCallID == "" {
continue
}
var result replayEditSuccess
if err := json.Unmarshal([]byte(message.Content), &result); err != nil || result.Success == nil {
continue
}
path := strings.TrimSpace(result.Success.Path)
if path == "" {
path = toolPaths[toolCallID]
}
if path == "" {
continue
}
sourceField := editReminderSourceField(result)
if sourceField == "" {
continue
}
return editReminderCandidate{
ToolCallID: toolCallID,
Path: path,
SourceField: sourceField,
}, true
}
return editReminderCandidate{}, false
}
func editReminderSourceField(result replayEditSuccess) string {
if result.Success == nil {
return ""
}
diffString := strings.TrimSpace(firstNonEmpty(result.Success.DiffStringSnake, result.Success.DiffString))
if diffString != "" && !looksLikeProjectedReplayTruncation(diffString) {
return "`success.diff_string`"
}
afterContent := strings.TrimSpace(firstNonEmpty(result.Success.AfterFullFileContentSnake, result.Success.AfterFullFileContent))
if afterContent != "" && !looksLikeProjectedReplayTruncation(afterContent) {
return "`success.after_full_file_content`"
}
return ""
}
func looksLikeProjectedReplayTruncation(value string) bool {
trimmed := strings.TrimSpace(value)
return strings.Contains(trimmed, "[truncated:") || strings.Contains(trimmed, "[tool result replay truncated:")
}
func collectReplayEditToolPaths(messages []modeladapter.Message) map[string]string {
paths := make(map[string]string)
for _, message := range messages {
if strings.TrimSpace(message.Role) != "assistant" || len(message.ToolCalls) == 0 {
continue
}
for _, toolCall := range message.ToolCalls {
toolName := strings.TrimSpace(toolCall.Function.Name)
if !isEditReminderToolName(toolName) {
continue
}
toolCallID := strings.TrimSpace(toolCall.ID)
if toolCallID == "" {
continue
}
if path := extractPathFromToolArguments(toolCall.Function.Arguments); path != "" {
paths[toolCallID] = path
}
}
}
return paths
}
func isEditReminderToolName(toolName string) bool {
switch strings.TrimSpace(toolName) {
case "Write", "PatchEdit", "PatchEditLines", "PatchEditSpan":
return true
default:
return false
}
}
func extractPathFromToolArguments(arguments string) string {
trimmed := strings.TrimSpace(arguments)
if trimmed == "" || trimmed == "{}" || trimmed == "null" {
return ""
}
var payload map[string]any
if err := json.Unmarshal([]byte(trimmed), &payload); err != nil {
return ""
}
for _, key := range []string{"path", "file_path"} {
if value, ok := payload[key].(string); ok && strings.TrimSpace(value) != "" {
return strings.TrimSpace(value)
}
}
return ""
}
@@ -0,0 +1,358 @@
package forwarder
import (
"context"
"crypto/sha256"
"encoding/hex"
"fmt"
"net/http"
"path/filepath"
"strings"
"time"
"connectrpc.com/connect"
"cursor/gen/aiserverv1"
"cursor/internal/logger"
)
const (
RepositoryServiceFastRepoInitHandshakeV2Procedure = "/aiserver.v1.RepositoryService/FastRepoInitHandshakeV2"
RepositoryServiceFastRepoInitHandshakeProcedure = "/aiserver.v1.RepositoryService/FastRepoInitHandshake"
RepositoryServiceFastRepoSyncCompleteProcedure = "/aiserver.v1.RepositoryService/FastRepoSyncComplete"
RepositoryServiceSyncMerkleSubtreeV2Procedure = "/aiserver.v1.RepositoryService/SyncMerkleSubtreeV2"
RepositoryServiceSyncMerkleSubtreeProcedure = "/aiserver.v1.RepositoryService/SyncMerkleSubtree"
RepositoryServiceFastUpdateFileV2Procedure = "/aiserver.v1.RepositoryService/FastUpdateFileV2"
RepositoryServiceFastUpdateFileProcedure = "/aiserver.v1.RepositoryService/FastUpdateFile"
RepositoryServiceEnsureIndexCreatedProcedure = "/aiserver.v1.RepositoryService/EnsureIndexCreated"
RepositoryServiceGetCopyStatusProcedure = "/aiserver.v1.RepositoryService/GetCopyStatus"
RepositoryServiceGetUploadLimitsProcedure = "/aiserver.v1.RepositoryService/GetUploadLimits"
RepositoryServiceGetNumFilesToSendProcedure = "/aiserver.v1.RepositoryService/GetNumFilesToSend"
RepositoryServiceGetAvailableChunkingStrategiesProcedure = "/aiserver.v1.RepositoryService/GetAvailableChunkingStrategies"
RepositoryServiceGetHighLevelFolderDescriptionProcedure = "/aiserver.v1.RepositoryService/GetHighLevelFolderDescription"
RepositoryServiceRepositoryStatusProcedure = "/aiserver.v1.RepositoryService/RepositoryStatus"
RepositoryServiceBatchRepositoryStatusProcedure = "/aiserver.v1.RepositoryService/BatchRepositoryStatus"
)
func newRepositoryServiceHandler(service *Service) http.Handler {
mux := http.NewServeMux()
mux.Handle(
RepositoryServiceFastRepoInitHandshakeV2Procedure,
connect.NewUnaryHandler(RepositoryServiceFastRepoInitHandshakeV2Procedure, service.FastRepoInitHandshakeV2),
)
mux.Handle(
RepositoryServiceFastRepoInitHandshakeProcedure,
connect.NewUnaryHandler(RepositoryServiceFastRepoInitHandshakeProcedure, service.FastRepoInitHandshake),
)
mux.Handle(
RepositoryServiceFastRepoSyncCompleteProcedure,
connect.NewUnaryHandler(RepositoryServiceFastRepoSyncCompleteProcedure, service.FastRepoSyncComplete),
)
mux.Handle(
RepositoryServiceSyncMerkleSubtreeV2Procedure,
connect.NewUnaryHandler(RepositoryServiceSyncMerkleSubtreeV2Procedure, service.SyncMerkleSubtreeV2),
)
mux.Handle(
RepositoryServiceSyncMerkleSubtreeProcedure,
connect.NewUnaryHandler(RepositoryServiceSyncMerkleSubtreeProcedure, service.SyncMerkleSubtree),
)
mux.Handle(
RepositoryServiceFastUpdateFileV2Procedure,
connect.NewUnaryHandler(RepositoryServiceFastUpdateFileV2Procedure, service.FastUpdateFileV2),
)
mux.Handle(
RepositoryServiceFastUpdateFileProcedure,
connect.NewUnaryHandler(RepositoryServiceFastUpdateFileProcedure, service.FastUpdateFile),
)
mux.Handle(
RepositoryServiceEnsureIndexCreatedProcedure,
connect.NewUnaryHandler(RepositoryServiceEnsureIndexCreatedProcedure, service.EnsureIndexCreated),
)
mux.Handle(
RepositoryServiceGetCopyStatusProcedure,
connect.NewUnaryHandler(RepositoryServiceGetCopyStatusProcedure, service.GetCopyStatus),
)
mux.Handle(
RepositoryServiceGetUploadLimitsProcedure,
connect.NewUnaryHandler(RepositoryServiceGetUploadLimitsProcedure, service.GetUploadLimits),
)
mux.Handle(
RepositoryServiceGetNumFilesToSendProcedure,
connect.NewUnaryHandler(RepositoryServiceGetNumFilesToSendProcedure, service.GetNumFilesToSend),
)
mux.Handle(
RepositoryServiceGetAvailableChunkingStrategiesProcedure,
connect.NewUnaryHandler(RepositoryServiceGetAvailableChunkingStrategiesProcedure, service.GetAvailableChunkingStrategies),
)
mux.Handle(
RepositoryServiceGetHighLevelFolderDescriptionProcedure,
connect.NewUnaryHandler(RepositoryServiceGetHighLevelFolderDescriptionProcedure, service.GetHighLevelFolderDescription),
)
mux.Handle(
RepositoryServiceRepositoryStatusProcedure,
connect.NewUnaryHandler(RepositoryServiceRepositoryStatusProcedure, service.RepositoryStatus),
)
mux.Handle(
RepositoryServiceBatchRepositoryStatusProcedure,
connect.NewUnaryHandler(RepositoryServiceBatchRepositoryStatusProcedure, service.BatchRepositoryStatus),
)
mux.Handle("/", http.NotFoundHandler())
return mux
}
func (service *Service) FastRepoInitHandshakeV2(_ context.Context, req *connect.Request[aiserverv1.FastRepoInitHandshakeV2Request]) (*connect.Response[aiserverv1.FastRepoInitHandshakeV2Response], error) {
repo := req.Msg.GetRepository()
codebaseID := repositoryCodebaseID(repo)
if err := service.ensureRepositoryCodebase(codebaseID, repo); err != nil {
return nil, connect.NewError(connect.CodeInternal, err)
}
logger.Infof("RepositoryService FastRepoInitHandshakeV2 ready codebase_id=%s repo=%s root_hash_present=%t", codebaseID, repositoryLogName(repo), strings.TrimSpace(req.Msg.GetRootHash()) != "")
return connect.NewResponse(&aiserverv1.FastRepoInitHandshakeV2Response{
Status: aiserverv1.FastRepoInitHandshakeV2Response_STATUS_SUCCESS,
Codebases: []*aiserverv1.RepositoryCodebaseInfo{
{
CodebaseId: codebaseID,
Status: aiserverv1.RepositoryCodebaseInfo_STATUS_UP_TO_DATE,
},
},
}), nil
}
func (service *Service) FastRepoInitHandshake(_ context.Context, req *connect.Request[aiserverv1.FastRepoInitHandshakeRequest]) (*connect.Response[aiserverv1.FastRepoInitHandshakeResponse], error) {
repo := req.Msg.GetRepository()
codebaseID := repositoryCodebaseID(repo)
if err := service.ensureRepositoryCodebase(codebaseID, repo); err != nil {
return nil, connect.NewError(connect.CodeInternal, err)
}
repoName := repositoryRepoName(repo)
logger.Infof("RepositoryService FastRepoInitHandshake ready codebase_id=%s repo=%s root_hash_present=%t", codebaseID, repositoryLogName(repo), strings.TrimSpace(req.Msg.GetRootHash()) != "")
return connect.NewResponse(&aiserverv1.FastRepoInitHandshakeResponse{
Status: aiserverv1.FastRepoInitHandshakeResponse_STATUS_UP_TO_DATE,
RepoName: repoName,
}), nil
}
func (service *Service) FastRepoSyncComplete(_ context.Context, req *connect.Request[aiserverv1.FastRepoSyncCompleteRequest]) (*connect.Response[aiserverv1.FastRepoSyncCompleteResponse], error) {
store, err := service.requireCodebaseIndexStore()
if err != nil {
return nil, connect.NewError(connect.CodeInternal, err)
}
for _, codebase := range req.Msg.GetCodebases() {
codebaseID := strings.TrimSpace(codebase.GetCodebaseId())
if codebaseID == "" {
continue
}
if err := store.MarkRepositorySyncComplete(codebaseID, codebase.GetPathKeyHash()); err != nil {
return nil, connect.NewError(connect.CodeInternal, err)
}
}
logger.Infof("RepositoryService FastRepoSyncComplete acknowledged codebase_count=%d", len(req.Msg.GetCodebases()))
return connect.NewResponse(&aiserverv1.FastRepoSyncCompleteResponse{}), nil
}
func (service *Service) SyncMerkleSubtreeV2(_ context.Context, req *connect.Request[aiserverv1.SyncMerkleSubtreeV2Request]) (*connect.Response[aiserverv1.SyncMerkleSubtreeV2Response], error) {
logger.Infof("RepositoryService SyncMerkleSubtreeV2 matched codebase_id=%s partial_paths=%d", strings.TrimSpace(req.Msg.GetCodebaseId()), len(req.Msg.GetLocalPartialPaths()))
return connect.NewResponse(&aiserverv1.SyncMerkleSubtreeV2Response{
Result: &aiserverv1.SyncMerkleSubtreeV2Response_Match{Match: true},
Results: repositoryPartialPathMatchResults(len(req.Msg.GetLocalPartialPaths())),
}), nil
}
func (service *Service) SyncMerkleSubtree(_ context.Context, req *connect.Request[aiserverv1.SyncMerkleSubtreeRequest]) (*connect.Response[aiserverv1.SyncMerkleSubtreeResponse], error) {
logger.Infof("RepositoryService SyncMerkleSubtree matched repo=%s", repositoryLogName(req.Msg.GetRepository()))
return connect.NewResponse(&aiserverv1.SyncMerkleSubtreeResponse{
Result: &aiserverv1.SyncMerkleSubtreeResponse_Match{Match: true},
}), nil
}
func (service *Service) FastUpdateFileV2(_ context.Context, req *connect.Request[aiserverv1.FastUpdateFileV2Request]) (*connect.Response[aiserverv1.FastUpdateFileV2Response], error) {
logger.Infof("RepositoryService FastUpdateFileV2 accepted codebase_id=%s update_type=%s file_updates=%d", strings.TrimSpace(req.Msg.GetCodebaseId()), req.Msg.GetUpdateType().String(), len(req.Msg.GetFileUpdates()))
return connect.NewResponse(&aiserverv1.FastUpdateFileV2Response{
Status: aiserverv1.FastUpdateFileV2Response_STATUS_SUCCESS,
}), nil
}
func (service *Service) FastUpdateFile(_ context.Context, req *connect.Request[aiserverv1.FastUpdateFileRequest]) (*connect.Response[aiserverv1.FastUpdateFileResponse], error) {
logger.Infof("RepositoryService FastUpdateFile accepted repo=%s update_type=%s", repositoryLogName(req.Msg.GetRepository()), req.Msg.GetUpdateType().String())
return connect.NewResponse(&aiserverv1.FastUpdateFileResponse{
Status: aiserverv1.FastUpdateFileResponse_STATUS_SUCCESS,
}), nil
}
func (service *Service) EnsureIndexCreated(_ context.Context, req *connect.Request[aiserverv1.EnsureIndexCreatedRequest]) (*connect.Response[aiserverv1.EnsureIndexCreatedResponse], error) {
repo := req.Msg.GetRepository()
codebaseID := repositoryCodebaseID(repo)
if err := service.ensureRepositoryCodebase(codebaseID, repo); err != nil {
return nil, connect.NewError(connect.CodeInternal, err)
}
logger.Infof("RepositoryService EnsureIndexCreated ready codebase_id=%s repo=%s", codebaseID, repositoryLogName(repo))
return connect.NewResponse(&aiserverv1.EnsureIndexCreatedResponse{}), nil
}
func (service *Service) GetCopyStatus(_ context.Context, req *connect.Request[aiserverv1.GetCopyStatusRequest]) (*connect.Response[aiserverv1.GetCopyStatusResponse], error) {
store, err := service.requireCodebaseIndexStore()
if err != nil {
return nil, connect.NewError(connect.CodeInternal, err)
}
if codebaseID := strings.TrimSpace(req.Msg.GetCodebaseId()); codebaseID != "" {
if err := store.MarkRepositoryCopyComplete(codebaseID, req.Msg.GetCopyTaskHandle()); err != nil {
return nil, connect.NewError(connect.CodeInternal, err)
}
}
completed := aiserverv1.GetCopyStatusResponse_COMPLETED_STATUS_UP_TO_DATE
logger.Infof("RepositoryService GetCopyStatus completed codebase_id=%s copy_task_handle=%s", strings.TrimSpace(req.Msg.GetCodebaseId()), strings.TrimSpace(req.Msg.GetCopyTaskHandle()))
return connect.NewResponse(&aiserverv1.GetCopyStatusResponse{
Phase: aiserverv1.GetCopyStatusResponse_PHASE_COMPLETED,
PercentDone: 1,
CompletedStatus: &completed,
}), nil
}
func (service *Service) GetUploadLimits(_ context.Context, req *connect.Request[aiserverv1.GetUploadLimitsRequest]) (*connect.Response[aiserverv1.GetUploadLimitsResponse], error) {
logger.Infof("RepositoryService GetUploadLimits allowed repo=%s", repositoryLogName(req.Msg.GetRepository()))
return connect.NewResponse(&aiserverv1.GetUploadLimitsResponse{
SoftLimit: 1_000_000,
HardLimit: 1_000_000,
}), nil
}
func (service *Service) GetNumFilesToSend(_ context.Context, req *connect.Request[aiserverv1.GetNumFilesToSendRequest]) (*connect.Response[aiserverv1.GetNumFilesToSendResponse], error) {
logger.Infof("RepositoryService GetNumFilesToSend none repo=%s", repositoryLogName(req.Msg.GetRepository()))
return connect.NewResponse(&aiserverv1.GetNumFilesToSendResponse{NumFiles: 0}), nil
}
func (service *Service) GetAvailableChunkingStrategies(_ context.Context, req *connect.Request[aiserverv1.GetAvailableChunkingStrategiesRequest]) (*connect.Response[aiserverv1.GetAvailableChunkingStrategiesResponse], error) {
logger.Infof("RepositoryService GetAvailableChunkingStrategies default repo=%s", repositoryLogName(req.Msg.GetRepository()))
return connect.NewResponse(&aiserverv1.GetAvailableChunkingStrategiesResponse{
ChunkingStrategies: []aiserverv1.ChunkingStrategy{aiserverv1.ChunkingStrategy_CHUNKING_STRATEGY_DEFAULT},
}), nil
}
func (service *Service) GetHighLevelFolderDescription(_ context.Context, req *connect.Request[aiserverv1.GetHighLevelFolderDescriptionRequest]) (*connect.Response[aiserverv1.GetHighLevelFolderDescriptionResponse], error) {
logger.Infof("RepositoryService GetHighLevelFolderDescription empty workspace_root=%s top_level_count=%d", strings.TrimSpace(req.Msg.GetWorkspaceRootPath()), len(req.Msg.GetTopLevelRelativeWorkspacePaths()))
return connect.NewResponse(&aiserverv1.GetHighLevelFolderDescriptionResponse{}), nil
}
func (service *Service) RepositoryStatus(_ context.Context, req *connect.Request[aiserverv1.RepositoryStatusRequest]) (*connect.Response[aiserverv1.RepositoryStatusResponse], error) {
logger.Infof("RepositoryService RepositoryStatus synced repo=%s", repositoryLogName(req.Msg.GetRepository()))
return connect.NewResponse(repositoryStatusSyncedResponse()), nil
}
func (service *Service) BatchRepositoryStatus(_ context.Context, req *connect.Request[aiserverv1.BatchRepositoryStatusRequest]) (*connect.Response[aiserverv1.BatchRepositoryStatusResponse], error) {
requests := req.Msg.GetRequests()
responses := make([]*aiserverv1.RepositoryStatusResponse, 0, len(requests))
for _, item := range requests {
logger.Infof("RepositoryService BatchRepositoryStatus synced repo=%s", repositoryLogName(item.GetRepository()))
responses = append(responses, repositoryStatusSyncedResponse())
}
return connect.NewResponse(&aiserverv1.BatchRepositoryStatusResponse{Responses: responses}), nil
}
func repositoryStatusSyncedResponse() *aiserverv1.RepositoryStatusResponse {
owner := true
return &aiserverv1.RepositoryStatusResponse{
IsOwner: &owner,
Status: &aiserverv1.RepositoryStatusResponse_Synced_{
Synced: &aiserverv1.RepositoryStatusResponse_Synced{
Branch: "local",
Commit: "local",
},
},
}
}
func (service *Service) ensureRepositoryCodebase(codebaseID string, repo *aiserverv1.RepositoryInfo) error {
store, err := service.requireCodebaseIndexStore()
if err != nil {
return err
}
_, err = store.EnsureRepositoryIndexed(RepositoryIndexStateRecord{
CodebaseID: strings.TrimSpace(codebaseID),
Repo: repositoryRepoName(repo),
TargetDir: strings.TrimSpace(repo.GetRelativeWorkspacePath()),
Status: repositoryIndexStatusReady,
LastHandshakeAt: time.Now().UTC(),
})
return err
}
func repositoryPartialPathMatchResults(count int) []*aiserverv1.SyncMerkleSubtreeV2Response_PartialPathResult {
if count <= 0 {
return nil
}
results := make([]*aiserverv1.SyncMerkleSubtreeV2Response_PartialPathResult, 0, count)
for index := 0; index < count; index++ {
results = append(results, &aiserverv1.SyncMerkleSubtreeV2Response_PartialPathResult{
Result: &aiserverv1.SyncMerkleSubtreeV2Response_PartialPathResult_Match{Match: true},
})
}
return results
}
func repositoryCodebaseID(repo *aiserverv1.RepositoryInfo, extras ...string) string {
parts := repositoryIdentityParts(repo)
for _, extra := range extras {
trimmed := strings.TrimSpace(extra)
if trimmed != "" {
parts = append(parts, trimmed)
}
}
if len(parts) == 0 {
return "cb_local"
}
sum := sha256.Sum256([]byte(strings.Join(parts, "\x00")))
return "cb_" + hex.EncodeToString(sum[:])[:24]
}
func repositoryIdentityParts(repo *aiserverv1.RepositoryInfo) []string {
if repo == nil {
return nil
}
parts := make([]string, 0, 6+len(repo.GetRemoteUrls()))
for _, value := range []string{
repo.GetWorkspaceUri(),
repo.GetRelativeWorkspacePath(),
repo.GetRepoOwner(),
repo.GetRepoName(),
} {
trimmed := strings.TrimSpace(value)
if trimmed != "" {
parts = append(parts, trimmed)
}
}
for _, remoteURL := range repo.GetRemoteUrls() {
trimmed := strings.TrimSpace(remoteURL)
if trimmed != "" {
parts = append(parts, trimmed)
}
}
if repo.GetOrthogonalTransformSeed() != 0 {
parts = append(parts, fmt.Sprintf("seed:%f", repo.GetOrthogonalTransformSeed()))
}
return parts
}
func repositoryRepoName(repo *aiserverv1.RepositoryInfo) string {
if repo == nil {
return "local"
}
if owner := strings.TrimSpace(repo.GetRepoOwner()); owner != "" {
if name := strings.TrimSpace(repo.GetRepoName()); name != "" {
return owner + "/" + name
}
}
for _, value := range []string{repo.GetRepoName(), repo.GetRelativeWorkspacePath(), repo.GetWorkspaceUri()} {
trimmed := strings.TrimSpace(value)
if trimmed != "" {
return filepath.Base(trimmed)
}
}
return "local"
}
func repositoryLogName(repo *aiserverv1.RepositoryInfo) string {
if repo == nil {
return ""
}
return strings.TrimSpace(strings.Join(repositoryIdentityParts(repo), "|"))
}
@@ -0,0 +1,395 @@
package forwarder
import (
"fmt"
"path/filepath"
"strings"
"google.golang.org/protobuf/proto"
"cursor/gen/agentv1"
)
func normalizeRequestContextForStorage(requestContext *agentv1.RequestContext) *agentv1.RequestContext {
if requestContext == nil {
return nil
}
return normalizeRequestContextForStorageMode(requestContext, true)
}
func normalizeRequestContextForStorageMode(requestContext *agentv1.RequestContext, includeStatic bool) *agentv1.RequestContext {
if requestContext == nil {
return nil
}
if !includeStatic {
return normalizeRealtimeRequestContextForStorage(requestContext)
}
cloned, ok := proto.Clone(requestContext).(*agentv1.RequestContext)
if !ok || cloned == nil {
return requestContext
}
descriptors := collectSkillDescriptors(cloned)
cloned.Rules = filterNonSkillRules(cloned.GetRules())
cloned.AgentSkills = nil
descriptors = guardSkillDescriptors(descriptors)
if len(descriptors) > 0 {
cloned.SkillOptions = &agentv1.SkillOptions{SkillDescriptors: descriptors}
} else {
cloned.SkillOptions = nil
}
guardRequestContextForStorage(cloned)
return cloned
}
func normalizeRealtimeRequestContextForStorage(requestContext *agentv1.RequestContext) *agentv1.RequestContext {
if requestContext == nil {
return nil
}
normalized := &agentv1.RequestContext{}
normalized.Rules = guardCursorRules(filterNonSkillRules(requestContext.GetRules()))
if fileContents := normalizeRealtimeFileContents(requestContext.GetFileContents()); len(fileContents) > 0 {
normalized.FileContents = fileContents
}
if summary := strings.TrimSpace(requestContext.GetUserIntentSummary()); summary != "" {
normalized.UserIntentSummary = stringPtr(truncatePromptGuardText("request_context.user_intent_summary", summary, promptGuardRealtimeTextChars))
}
if hooks := strings.TrimSpace(requestContext.GetHooksAdditionalContext()); hooks != "" {
normalized.HooksAdditionalContext = stringPtr(truncatePromptGuardText("request_context.hooks_additional_context", hooks, promptGuardRealtimeTextChars))
}
if commit := strings.TrimSpace(requestContext.GetCommitAttributionMessage()); commit != "" {
normalized.CommitAttributionMessage = stringPtr(truncatePromptGuardText("request_context.commit_attribution_message", commit, promptGuardRealtimeTextChars))
}
if pr := strings.TrimSpace(requestContext.GetPrAttributionMessage()); pr != "" {
normalized.PrAttributionMessage = stringPtr(truncatePromptGuardText("request_context.pr_attribution_message", pr, promptGuardRealtimeTextChars))
}
if !hasRealtimeRequestContextContent(normalized) {
return nil
}
return normalized
}
func normalizeRealtimeFileContents(fileContents map[string]string) map[string]string {
if len(fileContents) == 0 {
return nil
}
normalized := make(map[string]string, len(fileContents))
for path, content := range fileContents {
trimmedPath := strings.TrimSpace(path)
if trimmedPath == "" || strings.TrimSpace(content) == "" {
continue
}
normalized[trimmedPath] = content
}
if len(normalized) == 0 {
return nil
}
return guardStringMap(normalized, "request_context.file_contents", promptGuardRequestFileChars, promptGuardRequestFilesTotalChars, promptGuardRequestFilesMaxCount)
}
func hasRealtimeRequestContextContent(requestContext *agentv1.RequestContext) bool {
if requestContext == nil {
return false
}
return len(requestContext.GetRules()) > 0 ||
len(requestContext.GetFileContents()) > 0 ||
strings.TrimSpace(requestContext.GetUserIntentSummary()) != "" ||
strings.TrimSpace(requestContext.GetHooksAdditionalContext()) != "" ||
strings.TrimSpace(requestContext.GetCommitAttributionMessage()) != "" ||
strings.TrimSpace(requestContext.GetPrAttributionMessage()) != ""
}
func collectSkillDescriptors(requestContext *agentv1.RequestContext) []*agentv1.SkillDescriptor {
if requestContext == nil {
return nil
}
orderedKeys := make([]string, 0, 16)
descriptorsByKey := make(map[string]*agentv1.SkillDescriptor)
addDescriptor := func(descriptor *agentv1.SkillDescriptor) {
normalized := normalizeSkillDescriptor(descriptor)
if normalized == nil {
return
}
key := skillDescriptorKey(normalized)
if key == "" {
return
}
if existing, ok := descriptorsByKey[key]; ok {
mergeSkillDescriptor(existing, normalized)
return
}
orderedKeys = append(orderedKeys, key)
descriptorsByKey[key] = normalized
}
for _, descriptor := range requestContext.GetSkillOptions().GetSkillDescriptors() {
addDescriptor(descriptor)
}
for _, skill := range requestContext.GetAgentSkills() {
addDescriptor(skillDescriptorFromAgentSkill(skill))
}
for _, rule := range requestContext.GetRules() {
if isSkillFilePath(rule.GetFullPath()) {
addDescriptor(skillDescriptorFromCursorRule(rule))
}
}
descriptors := make([]*agentv1.SkillDescriptor, 0, len(orderedKeys))
for _, key := range orderedKeys {
if descriptor := descriptorsByKey[key]; descriptor != nil {
descriptors = append(descriptors, descriptor)
}
}
return descriptors
}
func filterNonSkillRules(rules []*agentv1.CursorRule) []*agentv1.CursorRule {
if len(rules) == 0 {
return nil
}
filtered := make([]*agentv1.CursorRule, 0, len(rules))
for _, rule := range rules {
if rule == nil {
continue
}
if isSkillFilePath(rule.GetFullPath()) {
continue
}
filtered = append(filtered, rule)
}
return filtered
}
func skillDescriptorFromAgentSkill(skill *agentv1.AgentSkill) *agentv1.SkillDescriptor {
if skill == nil {
return nil
}
fullPath := strings.TrimSpace(skill.GetFullPath())
if fullPath == "" {
return nil
}
parseError := strings.TrimSpace(skill.GetParseError())
descriptor := &agentv1.SkillDescriptor{
Name: inferSkillName(fullPath),
Description: strings.TrimSpace(skill.GetDescription()),
FolderPath: filepath.Dir(fullPath),
Enabled: parseError == "",
ReadmeFilePath: fullPath,
PackageType: inferSkillPackageType(fullPath),
}
if parseError != "" {
descriptor.ParseError = &parseError
}
return descriptor
}
func skillDescriptorFromCursorRule(rule *agentv1.CursorRule) *agentv1.SkillDescriptor {
if rule == nil {
return nil
}
fullPath := strings.TrimSpace(rule.GetFullPath())
if fullPath == "" || !isSkillFilePath(fullPath) {
return nil
}
description := ""
if rule.GetType() != nil && rule.GetType().GetAgentFetched() != nil {
description = strings.TrimSpace(rule.GetType().GetAgentFetched().GetDescription())
}
parseError := strings.TrimSpace(rule.GetParseError())
descriptor := &agentv1.SkillDescriptor{
Name: inferSkillName(fullPath),
Description: description,
FolderPath: filepath.Dir(fullPath),
Enabled: parseError == "",
ReadmeFilePath: fullPath,
PackageType: inferSkillPackageType(fullPath),
}
if parseError != "" {
descriptor.ParseError = &parseError
}
return descriptor
}
func normalizeSkillDescriptor(descriptor *agentv1.SkillDescriptor) *agentv1.SkillDescriptor {
if descriptor == nil {
return nil
}
cloned := proto.Clone(descriptor).(*agentv1.SkillDescriptor)
cloned.Name = strings.TrimSpace(cloned.GetName())
cloned.Description = strings.TrimSpace(cloned.GetDescription())
cloned.FolderPath = strings.TrimSpace(cloned.GetFolderPath())
cloned.ReadmeFilePath = strings.TrimSpace(cloned.GetReadmeFilePath())
if cloned.ReadmeFilePath == "" && cloned.FolderPath != "" {
cloned.ReadmeFilePath = filepath.Join(cloned.FolderPath, "SKILL.md")
}
if cloned.FolderPath == "" && cloned.ReadmeFilePath != "" {
cloned.FolderPath = filepath.Dir(cloned.ReadmeFilePath)
}
if cloned.Name == "" {
cloned.Name = inferSkillName(cloned.GetReadmeFilePath())
}
if cloned.PackageType == agentv1.PackageType_PACKAGE_TYPE_UNSPECIFIED {
cloned.PackageType = inferSkillPackageType(firstNonEmpty(cloned.GetReadmeFilePath(), cloned.GetFolderPath()))
}
if cloned.GetReadmeFilePath() == "" || cloned.GetDescription() == "" {
return nil
}
return cloned
}
func mergeSkillDescriptor(target *agentv1.SkillDescriptor, candidate *agentv1.SkillDescriptor) {
if target == nil || candidate == nil {
return
}
if strings.TrimSpace(target.GetDescription()) == "" && strings.TrimSpace(candidate.GetDescription()) != "" {
target.Description = strings.TrimSpace(candidate.GetDescription())
}
if strings.TrimSpace(target.GetReadmeFilePath()) == "" && strings.TrimSpace(candidate.GetReadmeFilePath()) != "" {
target.ReadmeFilePath = strings.TrimSpace(candidate.GetReadmeFilePath())
}
if strings.TrimSpace(target.GetFolderPath()) == "" && strings.TrimSpace(candidate.GetFolderPath()) != "" {
target.FolderPath = strings.TrimSpace(candidate.GetFolderPath())
}
if !target.GetEnabled() && candidate.GetEnabled() {
target.Enabled = true
}
if target.GetPackageType() == agentv1.PackageType_PACKAGE_TYPE_UNSPECIFIED && candidate.GetPackageType() != agentv1.PackageType_PACKAGE_TYPE_UNSPECIFIED {
target.PackageType = candidate.GetPackageType()
}
if strings.TrimSpace(target.GetParseError()) == "" && strings.TrimSpace(candidate.GetParseError()) != "" {
parseError := strings.TrimSpace(candidate.GetParseError())
target.ParseError = &parseError
}
}
func skillDescriptorKey(descriptor *agentv1.SkillDescriptor) string {
if descriptor == nil {
return ""
}
if readme := strings.TrimSpace(descriptor.GetReadmeFilePath()); readme != "" {
return readme
}
if folder := strings.TrimSpace(descriptor.GetFolderPath()); folder != "" {
return folder
}
return strings.TrimSpace(descriptor.GetName())
}
func buildSkillDiscoveryMessage(requestContext *agentv1.RequestContext) string {
normalized := normalizeRequestContextForStorageMode(requestContext, true)
if normalized == nil || normalized.GetSkillOptions() == nil || len(normalized.GetSkillOptions().GetSkillDescriptors()) == 0 {
return ""
}
lines := []string{
"<agent_skills>",
"When users ask you to perform tasks, check if any of the available skills below can help complete the task more effectively. Skills provide specialized capabilities and domain knowledge. To use a skill, read the skill file at the provided absolute path using the Read tool, then follow the instructions within. When a skill is relevant, read and follow it IMMEDIATELY as your first action. NEVER just announce or mention a skill without actually reading and following it. Only use skills listed below.",
"",
`<available_skills description="Skills the agent can use. Use the Read tool with the provided absolute path to fetch full contents.">`,
}
for _, descriptor := range normalized.GetSkillOptions().GetSkillDescriptors() {
if descriptor == nil {
continue
}
fullPath := strings.TrimSpace(descriptor.GetReadmeFilePath())
description := strings.TrimSpace(descriptor.GetDescription())
if fullPath == "" || description == "" {
continue
}
lines = append(lines, fmt.Sprintf(`<agent_skill fullPath="%s">%s</agent_skill>`, escapeSkillPromptXML(fullPath), escapeSkillPromptXML(description)))
}
lines = append(lines, "</available_skills>", "</agent_skills>")
if len(lines) == 6 {
return ""
}
return strings.Join(lines, "\n\n")
}
func collectMCPToolServers(requestContext *agentv1.RequestContext) map[string]string {
if requestContext == nil {
return nil
}
servers := make(map[string]string)
addDescriptors := func(descriptors []*agentv1.McpDescriptor) {
for _, descriptor := range descriptors {
if descriptor == nil {
continue
}
serverIdentifier := firstNonEmpty(descriptor.GetServerIdentifier(), descriptor.GetServerName())
if strings.TrimSpace(serverIdentifier) == "" {
continue
}
for _, tool := range descriptor.GetTools() {
if tool == nil {
continue
}
toolName := strings.TrimSpace(tool.GetToolName())
if toolName == "" {
continue
}
if _, exists := servers[toolName]; exists {
continue
}
servers[toolName] = strings.TrimSpace(serverIdentifier)
}
}
}
addDescriptors(requestContext.GetMcpFileSystemOptions().GetMcpDescriptors())
addDescriptors(requestContext.GetMcpMetaToolOptions().GetMcpDescriptors())
if len(servers) == 0 {
return nil
}
return servers
}
func escapeSkillPromptXML(value string) string {
replacer := strings.NewReplacer(
"&", "&amp;",
`"`, "&quot;",
"<", "&lt;",
">", "&gt;",
)
return replacer.Replace(strings.TrimSpace(value))
}
func inferSkillName(path string) string {
trimmed := strings.TrimSpace(path)
if trimmed == "" {
return ""
}
if strings.EqualFold(filepath.Base(trimmed), "SKILL.md") {
return strings.TrimSpace(filepath.Base(filepath.Dir(trimmed)))
}
return strings.TrimSpace(strings.TrimSuffix(filepath.Base(trimmed), filepath.Ext(trimmed)))
}
func inferSkillPackageType(path string) agentv1.PackageType {
normalized := strings.ToLower(strings.TrimSpace(path))
switch {
case strings.Contains(normalized, "/.agents/skills/"):
return agentv1.PackageType_PACKAGE_TYPE_CURSOR_PROJECT
case strings.Contains(normalized, "/.cursor/skills/"),
strings.Contains(normalized, "/.cursor/skills-cursor/"),
strings.Contains(normalized, "/.codex/skills/"):
return agentv1.PackageType_PACKAGE_TYPE_CURSOR_PERSONAL
default:
return agentv1.PackageType_PACKAGE_TYPE_UNSPECIFIED
}
}
func isSkillFilePath(path string) bool {
trimmed := strings.TrimSpace(path)
if trimmed == "" {
return false
}
return strings.EqualFold(filepath.Base(trimmed), "SKILL.md")
}
+299
View File
@@ -0,0 +1,299 @@
package forwarder
import (
"context"
"fmt"
"strings"
"google.golang.org/protobuf/encoding/protojson"
"cursor/gen/agentv1"
)
type runRewindDecision struct {
Evaluated bool
Apply bool
Reason string
SkipReason string
IncomingMessageID string
HasClientTurnCount bool
ClientTurnCount int
ServerTailTurnSeq int64
ServerNextTurnSeq int64
TargetTurnSeq int64
TargetEntrySeq int64
TargetRequestID string
MatchCount int
DroppedEntryCount int
DroppedTurnCount int
DroppedSeqStart int64
DroppedSeqEnd int64
PrefixEntries []HistoryEntry
}
type runRewindMatch struct {
Entry HistoryEntry
}
func (service *Service) decideRunRewind(intent InboundIntent, conversation *ConversationFile) runRewindDecision {
decision := runRewindDecision{ClientTurnCount: -1}
if !shouldEvaluateRunRewind(intent) {
return decision
}
decision.Evaluated = true
decision.IncomingMessageID = strings.TrimSpace(intent.UserMessage.GetMessageId())
decision.HasClientTurnCount = intent.ConversationState != nil
if decision.HasClientTurnCount {
decision.ClientTurnCount = len(intent.ConversationState.GetTurns())
}
if conversation != nil {
decision.ServerTailTurnSeq = maxHistoryTurnSeq(conversation.Entries)
decision.ServerNextTurnSeq = conversation.NextTurnSeq
}
if decision.IncomingMessageID == "" {
decision.SkipReason = "missing_message_id"
return decision
}
if conversation == nil || len(conversation.Entries) == 0 {
decision.SkipReason = "message_id_not_found"
return decision
}
matches := findUserMessageEntriesByMessageID(conversation.Entries, decision.IncomingMessageID)
decision.MatchCount = len(matches)
if len(matches) == 0 {
decision.SkipReason = "message_id_not_found"
return decision
}
selected, selectReason := selectRunRewindMatch(matches, decision.ClientTurnCount, decision.HasClientTurnCount)
decision.TargetTurnSeq = selected.Entry.TurnSeq
decision.TargetEntrySeq = selected.Entry.Seq
decision.TargetRequestID = strings.TrimSpace(selected.Entry.RequestID)
if decision.TargetTurnSeq <= 0 {
decision.SkipReason = "target_turn_seq_missing"
return decision
}
serverTailBeyondTarget := decision.ServerTailTurnSeq > decision.TargetTurnSeq
clientBehindServerTail := decision.HasClientTurnCount && int64(decision.ClientTurnCount) < decision.ServerTailTurnSeq
if !serverTailBeyondTarget && !clientBehindServerTail {
decision.SkipReason = "message_id_at_active_tail"
return decision
}
decision.Apply = true
decision.Reason = selectReason
decision.PrefixEntries = prefixEntriesBeforeTurn(conversation.Entries, decision.TargetTurnSeq)
decision.DroppedEntryCount, decision.DroppedTurnCount, decision.DroppedSeqStart, decision.DroppedSeqEnd = droppedEntryStats(conversation.Entries, decision.TargetTurnSeq)
return decision
}
func shouldEvaluateRunRewind(intent InboundIntent) bool {
if strings.TrimSpace(intent.Kind) != "run" || intent.Prewarm || intent.UserMessage == nil {
return false
}
return conversationActionCase(intent.ClientMessage) == "user_message_action"
}
func findUserMessageEntriesByMessageID(entries []HistoryEntry, messageID string) []runRewindMatch {
messageID = strings.TrimSpace(messageID)
if messageID == "" || len(entries) == 0 {
return nil
}
matches := make([]runRewindMatch, 0, 1)
for _, entry := range entries {
if strings.TrimSpace(entry.Kind) != "user_message" || len(entry.Payload) == 0 {
continue
}
userMessage := &agentv1.UserMessage{}
if err := protojson.Unmarshal(entry.Payload, userMessage); err != nil {
continue
}
if strings.TrimSpace(userMessage.GetMessageId()) != messageID {
continue
}
matches = append(matches, runRewindMatch{Entry: entry})
}
return matches
}
func selectRunRewindMatch(matches []runRewindMatch, clientTurnCount int, hasClientTurnCount bool) (runRewindMatch, string) {
if len(matches) == 0 {
return runRewindMatch{}, "no_match"
}
if hasClientTurnCount && clientTurnCount >= 0 {
targetTurnSeq := int64(clientTurnCount) + 1
for _, match := range matches {
if match.Entry.TurnSeq == targetTurnSeq {
return match, "client_turn_count_aligned"
}
}
var candidate *runRewindMatch
for index := range matches {
match := matches[index]
if match.Entry.TurnSeq <= int64(clientTurnCount) {
continue
}
if candidate == nil || earlierHistoryEntry(match.Entry, candidate.Entry) {
candidate = &match
}
}
if candidate != nil {
return *candidate, "first_match_after_client_turn_count"
}
}
selected := matches[0]
for _, match := range matches[1:] {
if earlierHistoryEntry(match.Entry, selected.Entry) {
selected = match
}
}
return selected, "earliest_message_id_match"
}
func earlierHistoryEntry(left HistoryEntry, right HistoryEntry) bool {
if left.TurnSeq != right.TurnSeq {
return left.TurnSeq < right.TurnSeq
}
if left.Seq != right.Seq {
return left.Seq < right.Seq
}
return left.CreatedAt.Before(right.CreatedAt)
}
func maxHistoryTurnSeq(entries []HistoryEntry) int64 {
var maxTurnSeq int64
for _, entry := range entries {
if entry.TurnSeq > maxTurnSeq {
maxTurnSeq = entry.TurnSeq
}
}
return maxTurnSeq
}
func prefixEntriesBeforeTurn(entries []HistoryEntry, targetTurnSeq int64) []HistoryEntry {
if len(entries) == 0 {
return nil
}
prefix := make([]HistoryEntry, 0, len(entries))
for _, entry := range entries {
if entry.TurnSeq < targetTurnSeq {
prefix = append(prefix, entry)
}
}
return prefix
}
func droppedEntryStats(entries []HistoryEntry, targetTurnSeq int64) (int, int, int64, int64) {
if len(entries) == 0 {
return 0, 0, 0, 0
}
droppedTurns := make(map[int64]struct{})
var droppedEntries int
var seqStart int64
var seqEnd int64
for _, entry := range entries {
if entry.TurnSeq < targetTurnSeq {
continue
}
droppedEntries++
if entry.TurnSeq > 0 {
droppedTurns[entry.TurnSeq] = struct{}{}
}
if entry.Seq > 0 && (seqStart == 0 || entry.Seq < seqStart) {
seqStart = entry.Seq
}
if entry.Seq > seqEnd {
seqEnd = entry.Seq
}
}
return droppedEntries, len(droppedTurns), seqStart, seqEnd
}
func appendReplacementRunEntries(prefix []HistoryEntry, entries []HistoryEntry) []HistoryEntry {
replacement := make([]HistoryEntry, 0, len(prefix)+len(entries))
replacement = append(replacement, prefix...)
replacement = append(replacement, entries...)
return replacement
}
func (service *Service) applyRunRewindToConversation(conversation *ConversationFile, decision runRewindDecision, entries []HistoryEntry, intent InboundIntent, turnSeq int64) {
if conversation == nil || !decision.Apply {
return
}
conversation.Entries = nil
conversation.NextEntrySeq = 1
conversation.NextTurnSeq = 1
appendEntriesInPlace(conversation, appendReplacementRunEntries(decision.PrefixEntries, entries))
applyRunRewindConversationState(conversation, intent, turnSeq)
deriveConversationLoopState(conversation)
}
func applyRunRewindConversationState(conversation *ConversationFile, intent InboundIntent, turnSeq int64) {
if conversation == nil {
return
}
conversation.TokenDetailsUsedTokens = 0
if conversation.TokenDetailsMaxTokens == 0 {
conversation.TokenDetailsMaxTokens = projectedConversationMaxTokens
}
clearConversationAutoCompactionState(conversation)
conversation.LatestRequestPrefix = nil
conversation.LastProviderCall = nil
conversation.CurrentLoopID = fmt.Sprintf("%d:%s", turnSeq, strings.TrimSpace(intent.RequestID))
conversation.CurrentLoopStatus = "running"
conversation.CurrentRequestID = strings.TrimSpace(intent.RequestID)
conversation.CurrentTurnSeq = turnSeq
}
func applyRunRewindMetadata(conversation *ConversationFile, source *ConversationFile, intent InboundIntent, turnSeq int64) {
if conversation == nil {
return
}
if source != nil {
if strings.TrimSpace(source.ConversationID) != "" {
conversation.ConversationID = strings.TrimSpace(source.ConversationID)
}
if strings.TrimSpace(source.RootConversationID) != "" {
conversation.RootConversationID = strings.TrimSpace(source.RootConversationID)
}
conversation.ParentConversationID = strings.TrimSpace(source.ParentConversationID)
conversation.ParentToolCallID = strings.TrimSpace(source.ParentToolCallID)
conversation.SubagentTypeName = strings.TrimSpace(source.SubagentTypeName)
if strings.TrimSpace(source.Mode) != "" {
conversation.Mode = strings.TrimSpace(source.Mode)
}
if source.TokenDetailsMaxTokens > 0 {
conversation.TokenDetailsMaxTokens = source.TokenDetailsMaxTokens
}
}
applyRunRewindConversationState(conversation, intent, turnSeq)
}
func (service *Service) logRunRewindDecision(requestID string, conversationID string, eventName string, decision runRewindDecision) {
if service == nil || !decision.Evaluated {
return
}
fields := map[string]any{
"message_id": decision.IncomingMessageID,
"apply": decision.Apply,
"reason": decision.Reason,
"skip_reason": decision.SkipReason,
"target_turn_seq": decision.TargetTurnSeq,
"target_entry_seq": decision.TargetEntrySeq,
"target_request_id": decision.TargetRequestID,
"server_tail_turn_seq": decision.ServerTailTurnSeq,
"server_next_turn_seq": decision.ServerNextTurnSeq,
"match_count": decision.MatchCount,
"dropped_entry_count": decision.DroppedEntryCount,
"dropped_turn_count": decision.DroppedTurnCount,
"dropped_seq_start": decision.DroppedSeqStart,
"dropped_seq_end": decision.DroppedSeqEnd,
}
if decision.HasClientTurnCount {
fields["client_turn_count"] = decision.ClientTurnCount
} else {
fields["client_turn_count"] = nil
}
service.debug.LogRuntime(context.Background(), requestID, conversationID, eventName, fields)
}
@@ -0,0 +1,298 @@
package forwarder
import (
"encoding/json"
"fmt"
"strings"
"time"
"cursor/gen/agentv1"
modeladapter "cursor/internal/backend/agent/model"
)
func newRuntimeConversation(conversationID string, mode agentv1.AgentMode) (*ConversationFile, error) {
now := time.Now().UTC()
normalizedID := strings.TrimSpace(conversationID)
alias, err := modeAlias(mode)
if err != nil {
return nil, err
}
return &ConversationFile{
ConversationID: normalizedID,
RootConversationID: normalizedID,
Mode: alias,
CreatedAt: now,
UpdatedAt: now,
NextTurnSeq: 1,
NextEntrySeq: 1,
Entries: make([]HistoryEntry, 0, 16),
}, nil
}
func (service *Service) bootstrapRuntimeConversation(intent InboundIntent) (*ConversationFile, agentv1.AgentMode, int64, []HistoryEntry, error) {
if service == nil {
return nil, agentv1.AgentMode_AGENT_MODE_AGENT, 0, nil, fmt.Errorf("forwarder service is nil")
}
contextWindowTokens := service.resolveContextWindowTokens(intent.ModelID)
var conversation *ConversationFile
var err error
if service.store != nil {
conversation, err = service.store.LoadConversation(intent.ConversationID)
if err != nil {
return nil, agentv1.AgentMode_AGENT_MODE_AGENT, 0, nil, err
}
}
if conversation == nil {
conversation, err = newRuntimeConversation(intent.ConversationID, intent.Mode)
if err != nil {
return nil, agentv1.AgentMode_AGENT_MODE_AGENT, 0, nil, err
}
}
importedEntries := []HistoryEntry(nil)
if len(conversation.Entries) == 0 && intent.ConversationState != nil {
importedEntries, err = service.importConversationState(conversation, intent.ConversationState)
if err != nil {
return nil, agentv1.AgentMode_AGENT_MODE_AGENT, 0, nil, err
}
}
if strings.TrimSpace(conversation.ConversationID) == "" {
conversation.ConversationID = strings.TrimSpace(intent.ConversationID)
}
if strings.TrimSpace(conversation.RootConversationID) == "" {
conversation.RootConversationID = strings.TrimSpace(conversation.ConversationID)
}
if intent.HasExplicitMode {
alias, err := modeAlias(intent.Mode)
if err != nil {
return nil, agentv1.AgentMode_AGENT_MODE_AGENT, 0, nil, err
}
conversation.Mode = alias
} else if strings.TrimSpace(conversation.Mode) == "" {
alias, err := modeAlias(intent.Mode)
if err != nil {
return nil, agentv1.AgentMode_AGENT_MODE_AGENT, 0, nil, err
}
conversation.Mode = alias
}
if strings.TrimSpace(intent.SubagentTypeName) != "" {
conversation.SubagentTypeName = strings.TrimSpace(intent.SubagentTypeName)
}
if contextWindowTokens > 0 {
conversation.TokenDetailsMaxTokens = contextWindowTokens
} else if conversation.TokenDetailsMaxTokens == 0 {
conversation.TokenDetailsMaxTokens = projectedConversationMaxTokens
}
turnSeq := conversation.NextTurnSeq
if turnSeq <= 0 {
turnSeq = 1
}
effectiveMode, err := parseModeAlias(conversation.Mode)
if err != nil {
return nil, agentv1.AgentMode_AGENT_MODE_AGENT, 0, nil, err
}
entries, err := buildRunEntries(intent, effectiveMode, turnSeq)
if err != nil {
return nil, agentv1.AgentMode_AGENT_MODE_AGENT, 0, nil, err
}
conversation.CurrentRequestID = strings.TrimSpace(intent.RequestID)
conversation.CurrentTurnSeq = turnSeq
conversation.CurrentLoopID = fmt.Sprintf("%d:%s", turnSeq, strings.TrimSpace(intent.RequestID))
conversation.CurrentLoopStatus = "running"
if conversation.NextTurnSeq <= turnSeq {
conversation.NextTurnSeq = turnSeq + 1
}
return conversation, effectiveMode, turnSeq, append(importedEntries, entries...), nil
}
func (service *Service) syncConversationRecord(conversationID string, conversation *ConversationFile) error {
if service == nil || service.store == nil || conversation == nil {
return nil
}
mode, err := parseModeAlias(conversation.Mode)
if err != nil {
return err
}
if _, err := service.store.CreateConversation(
conversationID,
mode,
conversation.ParentConversationID,
conversation.ParentToolCallID,
conversation.RootConversationID,
); err != nil {
return err
}
_, err = service.store.UpdateConversationMeta(conversationID, func(item *ConversationFile) error {
if item == nil {
return nil
}
item.ConversationID = conversation.ConversationID
item.RootConversationID = conversation.RootConversationID
item.ParentConversationID = conversation.ParentConversationID
item.ParentToolCallID = conversation.ParentToolCallID
item.SubagentTypeName = conversation.SubagentTypeName
item.Mode = conversation.Mode
item.TokenDetailsUsedTokens = conversation.TokenDetailsUsedTokens
item.TokenDetailsMaxTokens = conversation.TokenDetailsMaxTokens
item.AutoCompactionPending = conversation.AutoCompactionPending
item.AutoCompactionPromptTokens = conversation.AutoCompactionPromptTokens
item.AutoCompactionReserveTokens = conversation.AutoCompactionReserveTokens
item.AutoCompactionTriggeredAt = conversation.AutoCompactionTriggeredAt
item.AutoCompactionSourceModelCallID = conversation.AutoCompactionSourceModelCallID
item.LatestRequestPrefix = cloneConversationRequestPrefix(conversation.LatestRequestPrefix)
item.LastProviderCall = cloneConversationProviderCall(conversation.LastProviderCall)
item.CreatedAt = conversation.CreatedAt
item.UpdatedAt = conversation.UpdatedAt
item.NextTurnSeq = conversation.NextTurnSeq
item.NextEntrySeq = conversation.NextEntrySeq
item.CurrentLoopID = conversation.CurrentLoopID
item.CurrentLoopStatus = conversation.CurrentLoopStatus
item.CurrentRequestID = conversation.CurrentRequestID
item.CurrentTurnSeq = conversation.CurrentTurnSeq
return nil
})
return err
}
func (service *Service) syncSummaryCarryForward(_ string, requestID string, modelCallID string) error {
if service == nil || strings.TrimSpace(requestID) == "" || strings.TrimSpace(modelCallID) == "" {
return nil
}
stream, ok := service.broker.Get(requestID)
if !ok || stream == nil {
return nil
}
conversation, _, _, err := service.snapshotCheckpointConversation(stream)
if err != nil {
return err
}
return service.syncSummarySnapshot(stream, conversation, requestID, modelCallID)
}
func (service *Service) syncSummarySnapshot(stream *ActiveStream, conversation *ConversationFile, requestID string, modelCallID string) error {
if service == nil || conversation == nil || strings.TrimSpace(requestID) == "" || strings.TrimSpace(modelCallID) == "" {
return nil
}
return service.syncConversationRecord(conversation.ConversationID, conversation)
}
func (service *Service) projectSummaryInputHistoryMessages(conversation *ConversationFile) ([]modeladapter.Message, error) {
if service == nil || service.projector == nil || conversation == nil {
return nil, nil
}
inputConversation := cloneConversationFile(conversation)
inputConversation.Entries = filterConversationEntriesByKind(conversation.Entries, "user_message", "request_context", "prompt_context")
return service.projector.ProjectPromptReplay(inputConversation)
}
func (service *Service) projectSummaryOutputMessages(conversation *ConversationFile) ([]modeladapter.Message, error) {
if service == nil || service.projector == nil || conversation == nil {
return nil, nil
}
outputConversation := cloneConversationFile(conversation)
outputConversation.Entries = filterConversationEntriesByKind(conversation.Entries, "assistant_text", "tool_call", "tool_result")
return service.projector.ProjectPromptReplay(outputConversation)
}
func filterConversationEntriesByKind(entries []HistoryEntry, kinds ...string) []HistoryEntry {
if len(entries) == 0 || len(kinds) == 0 {
return nil
}
allowed := make(map[string]struct{}, len(kinds))
for _, kind := range kinds {
if trimmed := strings.TrimSpace(kind); trimmed != "" {
allowed[trimmed] = struct{}{}
}
}
filtered := make([]HistoryEntry, 0, len(entries))
for _, entry := range entries {
if _, ok := allowed[strings.TrimSpace(entry.Kind)]; !ok {
continue
}
filtered = append(filtered, entry)
}
return filtered
}
func buildSummaryTurnOutcome(stream *ActiveStream, conversation *ConversationFile, requestID string, modelCallID string) map[string]any {
outcome := map[string]any{
"status": "in_progress",
"request_id": strings.TrimSpace(requestID),
"run_id": strings.TrimSpace(requestID),
"model_call_id": strings.TrimSpace(modelCallID),
}
if stream != nil {
outcome["model_id"] = strings.TrimSpace(stream.ModelID)
outcome["is_prewarm"] = false
outcome["turn_seq"] = stream.TurnSeq
if !stream.CreatedAt.IsZero() {
outcome["started_at"] = stream.CreatedAt.UTC().Format(time.RFC3339Nano)
}
}
if conversation != nil && !conversation.UpdatedAt.IsZero() {
outcome["last_event_at"] = conversation.UpdatedAt.UTC().Format(time.RFC3339Nano)
}
for index := len(conversation.Entries) - 1; index >= 0; index-- {
entry := conversation.Entries[index]
if strings.TrimSpace(entry.RequestID) != strings.TrimSpace(requestID) || strings.TrimSpace(entry.Kind) != "metadata" {
continue
}
var payload metadataPayload
if err := json.Unmarshal(entry.Payload, &payload); err != nil {
continue
}
switch strings.TrimSpace(payload.Type) {
case "turn_completed":
outcome["status"] = "completed"
case "provider_error":
outcome["status"] = "provider_error"
if errorText := strings.TrimSpace(readStringValue(payload.Value["error"])); errorText != "" {
outcome["error"] = errorText
}
case "failed":
outcome["status"] = "failed"
if errorText := strings.TrimSpace(readStringValue(payload.Value["error"])); errorText != "" {
outcome["error"] = errorText
}
default:
continue
}
if modelCall := strings.TrimSpace(readStringValue(payload.Value["model_call_id"])); modelCall != "" {
outcome["model_call_id"] = modelCall
}
if !entry.CreatedAt.IsZero() {
outcome["last_event_at"] = entry.CreatedAt.UTC().Format(time.RFC3339Nano)
}
break
}
return outcome
}
func buildSummaryAnnotations(conversation *ConversationFile, requestID string) map[string]any {
if conversation == nil {
return nil
}
for index := len(conversation.Entries) - 1; index >= 0; index-- {
entry := conversation.Entries[index]
if strings.TrimSpace(entry.RequestID) != strings.TrimSpace(requestID) || strings.TrimSpace(entry.Kind) != "metadata" {
continue
}
var payload metadataPayload
if err := json.Unmarshal(entry.Payload, &payload); err != nil {
continue
}
if strings.TrimSpace(payload.Type) != "thought_annotation" {
continue
}
if strings.TrimSpace(readStringValue(payload.Value["kind"])) != "summary_completed" {
continue
}
thought := strings.TrimSpace(readStringValue(payload.Value["thought"]))
if thought == "" {
continue
}
return map[string]any{
"summary_completed_thought": thought,
}
}
return nil
}
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,238 @@
package forwarder
import (
"fmt"
"log"
"strings"
"time"
runtimecore "cursor/internal/backend/agent/core"
)
const shellTerminalRecoveryGrace = 1500 * time.Millisecond
const (
shellRecoveryReasonForegroundDeadline = "foreground_deadline_exceeded"
shellRecoveryReasonTransportClosed = "transport_closed_without_terminal"
)
func initializePendingExecForTracking(pending runtimecore.PendingExec) runtimecore.PendingExec {
if strings.TrimSpace(pending.ExecKind) != "shell" {
return pending
}
now := time.Now().UTC()
if pending.OpenedAt.IsZero() {
pending.OpenedAt = now
}
if pending.LastShellActivityAt.IsZero() {
pending.LastShellActivityAt = pending.OpenedAt
}
if pending.ShellForegroundDeadline.IsZero() {
pending.ShellForegroundDeadline = pending.OpenedAt.Add(shellForegroundTimeoutDuration(pending.ArgsJSON) + shellTerminalRecoveryGrace)
}
return pending
}
func shellForegroundTimeoutDuration(argsJSON []byte) time.Duration {
timeoutMS := int64(30000)
args, err := runtimecore.DecodeArgsMap(argsJSON)
if err == nil {
if blockUntilMS, found, err := runtimecore.ReadFloat64Arg(args, "block_until_ms", "blockUntilMS"); err == nil && found {
if blockUntilMS <= 0 {
return 0
}
timeoutMS = int64(blockUntilMS)
}
}
return time.Duration(timeoutMS) * time.Millisecond
}
func shellForegroundTimeoutMS(argsJSON []byte) int64 {
return shellForegroundTimeoutDuration(argsJSON).Milliseconds()
}
func (service *Service) scheduleShellForegroundRecovery(requestID string, pending runtimecore.PendingExec) {
if service == nil || strings.TrimSpace(requestID) == "" || strings.TrimSpace(pending.ExecKind) != "shell" || strings.TrimSpace(pending.ExecID) == "" {
return
}
stream, ok := service.broker.Get(requestID)
if !ok || stream == nil {
return
}
deadline := pending.ShellForegroundDeadline
if deadline.IsZero() {
deadline = time.Now().UTC().Add(shellForegroundTimeoutDuration(pending.ArgsJSON) + shellTerminalRecoveryGrace)
}
service.scheduleStreamTimer(
stream,
providerTimerKey(streamTimerShellForeground, pending.ExecID),
time.Until(deadline),
streamTimerShellForeground,
pending.ExecID,
pending.MessageID,
shellRecoveryReasonForegroundDeadline,
)
}
func (service *Service) scheduleShellTransportCloseRecovery(requestID string, pending runtimecore.PendingExec) {
if service == nil || strings.TrimSpace(requestID) == "" || strings.TrimSpace(pending.ExecKind) != "shell" || strings.TrimSpace(pending.ExecID) == "" {
return
}
stream, ok := service.broker.Get(requestID)
if !ok || stream == nil {
return
}
if _, scheduled := markShellRecoveryScheduled(stream, pending.ExecID); !scheduled {
return
}
service.scheduleStreamTimer(
stream,
providerTimerKey(streamTimerShellTransportClose, pending.ExecID),
shellTerminalRecoveryGrace,
streamTimerShellTransportClose,
pending.ExecID,
pending.MessageID,
shellRecoveryReasonTransportClosed,
)
}
func snapshotPendingExecWithStatus(stream *ActiveStream, execID string) (runtimecore.PendingExec, StreamStatus, bool) {
if stream == nil || strings.TrimSpace(execID) == "" {
return runtimecore.PendingExec{}, "", false
}
stream.mu.Lock()
defer stream.mu.Unlock()
item, ok := stream.PendingExecs[strings.TrimSpace(execID)]
if !ok {
return runtimecore.PendingExec{}, stream.Status, false
}
return item, stream.Status, true
}
func markShellRecoveryScheduled(stream *ActiveStream, execID string) (runtimecore.PendingExec, bool) {
if stream == nil || strings.TrimSpace(execID) == "" {
return runtimecore.PendingExec{}, false
}
stream.mu.Lock()
defer stream.mu.Unlock()
current, ok := stream.PendingExecs[strings.TrimSpace(execID)]
if !ok || strings.TrimSpace(current.ExecKind) != "shell" || current.ShellRecoveryScheduled {
return current, false
}
current.ShellRecoveryScheduled = true
stream.PendingExecs[strings.TrimSpace(execID)] = current
stream.UpdatedAt = time.Now().UTC()
return current, true
}
func (service *Service) recoverShellWithoutTerminalIfNeeded(stream *ActiveStream, execID string, messageID uint32, reason string) error {
if stream == nil || strings.TrimSpace(execID) == "" {
return nil
}
current, status, found := snapshotPendingExecWithStatus(stream, execID)
if !found || current.MessageID != messageID || strings.TrimSpace(current.ExecKind) != "shell" || isTerminalStreamStatus(status) {
return nil
}
switch strings.TrimSpace(current.StreamState) {
case "exited", "backgrounded", "rejected", "permission_denied":
return nil
}
if reason == shellRecoveryReasonForegroundDeadline && !current.ShellForegroundDeadline.IsZero() && time.Now().UTC().Before(current.ShellForegroundDeadline) {
return nil
}
return service.recoverShellWithoutTerminal(stream, current, reason)
}
func (service *Service) recoverShellWithoutTerminal(stream *ActiveStream, pending runtimecore.PendingExec, reason string) error {
if stream == nil {
return nil
}
pending.ShellRecoveryScheduled = true
markExecCompleted(stream, pending)
resultPayload := buildSyntheticShellResultPayload(pending, reason)
log.Printf(
"forwarder synthetic shell recovery request_id=%s tool_call_id=%s message_id=%d exec_id=%s reason=%s stream_state=%s",
strings.TrimSpace(stream.RequestID),
strings.TrimSpace(pending.ToolCallID),
pending.MessageID,
strings.TrimSpace(pending.ExecID),
strings.TrimSpace(reason),
strings.TrimSpace(pending.StreamState),
)
if err := service.appendToolResult(stream, pending.ToolCallID, deriveToolNameFromPendingExec(pending), pending.ArgsJSON, resultPayload, pending.ReasoningContent, nil); err != nil {
return err
}
if _, err := service.appendConversationEntries(stream, stream.ConversationID, []HistoryEntry{
newMetadataEntry(stream.TurnSeq, stream.RequestID, "shell_stream_stalled", map[string]any{
"tool_call_id": pending.ToolCallID,
"message_id": pending.MessageID,
"exec_id": pending.ExecID,
"exec_kind": pending.ExecKind,
"reason": strings.TrimSpace(reason),
"recent_stream_state": pending.StreamState,
"chunk_count": pending.ChunkCount,
"first_chunk_at": pending.FirstChunkAt,
"last_activity_at": pending.LastShellActivityAt,
"last_heartbeat_at": pending.LastShellHeartbeatAt,
"foreground_deadline": pending.ShellForegroundDeadline,
"timeout_ms": shellForegroundTimeoutMS(pending.ArgsJSON),
"stdout_buffer_bytes": len(pending.StdoutBuffer),
"stderr_buffer_bytes": len(pending.StderrBuffer),
"shell_recovery_scheduled": pending.ShellRecoveryScheduled,
}),
}); err != nil {
return err
}
if err := service.syncSummaryCarryForward(stream.ConversationID, stream.RequestID, pending.ModelCallID); err != nil {
return err
}
if err := service.publishToolCallCompleted(stream.RequestID, pending.ToolCallID, pending.ModelCallID, nil); err != nil {
return err
}
if err := service.publishCheckpoint(stream.RequestID, stream.ConversationID); err != nil {
return err
}
return service.reconcileStream(stream)
}
func buildSyntheticShellResultPayload(pending runtimecore.PendingExec, reason string) string {
sections := make([]string, 0, 2)
if captured := summarizeCapturedShellOutput(pending.StdoutBuffer, pending.StderrBuffer); captured != "" {
sections = append(sections, captured)
}
noteLines := []string{
"<shell-incomplete>",
"Missing terminal shell stream event (expected exit or backgrounded).",
}
switch strings.TrimSpace(reason) {
case shellRecoveryReasonTransportClosed:
noteLines = append(noteLines, "The shell transport closed before a terminal event arrived.")
case shellRecoveryReasonForegroundDeadline:
noteLines = append(noteLines, fmt.Sprintf("The foreground wait window expired after %dms without a terminal event.", shellForegroundTimeoutMS(pending.ArgsJSON)))
default:
noteLines = append(noteLines, "The shell stream stopped progressing before a terminal event arrived.")
}
noteLines = append(noteLines,
"The command may still be running in the Cursor app client.",
"</shell-incomplete>",
)
sections = append(sections, strings.Join(noteLines, "\n"))
return strings.Join(sections, "\n\n")
}
func summarizeCapturedShellOutput(stdout string, stderr string) string {
trimmedStdout := strings.TrimSpace(stdout)
trimmedStderr := strings.TrimSpace(stderr)
sections := make([]string, 0, 2)
if trimmedStdout != "" {
sections = append(sections, trimmedStdout)
}
if trimmedStderr != "" {
if trimmedStdout != "" {
sections = append(sections, "<stderr>\n"+trimmedStderr+"\n</stderr>")
} else {
sections = append(sections, trimmedStderr)
}
}
return strings.Join(sections, "\n\n")
}
@@ -0,0 +1,37 @@
package forwarder
type latestSummaryState struct {
RuntimeSnapshot *ConversationFile
}
func (service *Service) loadLatestSummaryState(conversationID string) (*latestSummaryState, bool, error) {
if service == nil || service.store == nil {
return nil, false, nil
}
conversation, err := service.store.LoadConversation(conversationID)
if err != nil || conversation == nil {
return nil, false, err
}
return &latestSummaryState{RuntimeSnapshot: conversation}, true, nil
}
func (service *Service) loadLatestCarryForwardReplay(_ string) ([][]byte, bool, error) {
return nil, false, nil
}
func (service *Service) loadLatestSummaryPromptTokens(conversationID string) (int64, bool, error) {
if service == nil || service.store == nil {
return 0, false, nil
}
conversation, err := service.store.LoadConversation(conversationID)
if err != nil || conversation == nil {
return 0, false, err
}
if conversation.AutoCompactionPromptTokens > 0 {
return conversation.AutoCompactionPromptTokens, true, nil
}
if conversation.TokenDetailsUsedTokens > 0 {
return int64(conversation.TokenDetailsUsedTokens), true, nil
}
return 0, false, nil
}
+871
View File
@@ -0,0 +1,871 @@
package forwarder
import (
"encoding/json"
"fmt"
"sort"
"strconv"
"strings"
"time"
"google.golang.org/protobuf/encoding/protojson"
"google.golang.org/protobuf/proto"
"cursor/gen/agentv1"
modeladapter "cursor/internal/backend/agent/model"
)
const todoSectionReminderMessage = "<system_reminder>\nYou are currently under the todo section, be sure to track tasks and do not forget to update.\n</system_reminder>"
const (
promptContextSourceStructuredCurrentPlan = "structured_state/current_plan"
promptContextSourceStructuredTodoList = "structured_state/todo_list"
promptContextSourceStructuredTodoReminder = "structured_state/todo_reminder"
)
const todoUnsafeReplaceBaseError = "todo update rejected: merge=false may only omit completed or cancelled todos; use merge=true for incremental updates or include every active todo"
type structuredConversationState struct {
PlanText string
HasPlan bool
Plans map[string]*agentv1.PlanRegistryEntry
Todos []*agentv1.TodoItem
HasTodos bool
}
func projectConversationStructuredState(conversation *ConversationFile) (structuredConversationState, error) {
state := structuredConversationState{}
if conversation == nil {
return state, nil
}
for _, entry := range checkpointProjectionEntries(conversation.Entries) {
switch strings.TrimSpace(entry.Kind) {
case "runtime_state":
var payload runtimeStateEntryPayload
if err := json.Unmarshal(entry.Payload, &payload); err != nil {
return structuredConversationState{}, fmt.Errorf("decode runtime_state entry: %w", err)
}
if strings.TrimSpace(payload.PlanText) != "" {
state.PlanText = strings.TrimSpace(payload.PlanText)
state.HasPlan = true
}
if len(payload.Plans) > 0 {
state.Plans = clonePlanRegistryEntries(payload.Plans)
}
if len(payload.Todos) > 0 {
state.Todos = cloneTodoItems(payload.Todos)
state.HasTodos = true
}
continue
case "tool_result":
default:
continue
}
var payload toolResultEntryPayload
if err := json.Unmarshal(entry.Payload, &payload); err != nil {
return structuredConversationState{}, fmt.Errorf("decode structured state tool_result: %w", err)
}
if len(payload.ToolCall) == 0 {
continue
}
toolCall := &agentv1.ToolCall{}
if err := protojson.Unmarshal(payload.ToolCall, toolCall); err != nil {
return structuredConversationState{}, fmt.Errorf("decode structured state tool_call: %w", err)
}
entryTime := entry.CreatedAt
switch item := toolCall.GetTool().(type) {
case *agentv1.ToolCall_CreatePlanToolCall:
createPlanToolCall := item.CreatePlanToolCall
if createPlanToolCall == nil || createPlanToolCall.GetResult().GetSuccess() == nil {
continue
}
planURI := strings.TrimSpace(createPlanToolCall.GetResult().GetPlanUri())
if planURI == "" {
continue
}
state.PlanText = strings.TrimSpace(createPlanToolCall.GetArgs().GetPlan())
state.HasPlan = true
state.Plans = upsertCurrentPlanRegistryEntry(state.Plans, planURI)
todos, err := normalizeTodoItems(flattenCreatePlanTodos(createPlanToolCall.GetArgs()), entryTime, false)
if err != nil {
return structuredConversationState{}, err
}
state.Todos = todos
state.HasTodos = true
case *agentv1.ToolCall_UpdateTodosToolCall:
updateTodosToolCall := item.UpdateTodosToolCall
if updateTodosToolCall == nil || updateTodosToolCall.GetResult().GetSuccess() == nil {
continue
}
success := updateTodosToolCall.GetResult().GetSuccess()
todos, err := normalizeTodoItems(success.GetTodos(), entryTime, false)
if err != nil {
return structuredConversationState{}, err
}
if !success.GetWasMerge() && len(missingActiveTodoReplacementIDs(state.Todos, todos)) > 0 {
todos, err = mergeTodoItems(state.Todos, todos, entryTime)
if err != nil {
return structuredConversationState{}, err
}
}
state.Todos = todos
state.HasTodos = true
}
}
return state, nil
}
func refreshConversationRuntimeState(conversation *ConversationFile) error {
if conversation == nil {
return nil
}
state, err := projectConversationStructuredState(conversation)
if err != nil {
return err
}
conversation.CurrentPlanText = ""
conversation.CurrentPlans = nil
conversation.CurrentTodos = nil
if state.HasPlan && strings.TrimSpace(state.PlanText) != "" {
conversation.CurrentPlanText = strings.TrimSpace(state.PlanText)
}
if len(state.Plans) > 0 {
conversation.CurrentPlans = clonePlanRegistryEntries(state.Plans)
}
if state.HasTodos && len(state.Todos) > 0 {
conversation.CurrentTodos = cloneTodoItems(state.Todos)
}
return nil
}
func buildStructuredStatePromptContexts(conversation *ConversationFile) ([]PromptContextMessage, []modeladapter.Message, error) {
state, err := projectConversationStructuredState(conversation)
if err != nil {
return nil, nil, err
}
contexts := make([]PromptContextMessage, 0, 2)
tailMessages := make([]modeladapter.Message, 0, 1)
if state.HasPlan && strings.TrimSpace(state.PlanText) != "" {
contexts = append(contexts, newPromptContextMessage(
promptContextSourceStructuredCurrentPlan,
modeladapter.Message{
Role: "user",
Content: "<current_plan>\n" + strings.TrimSpace(state.PlanText) + "\n</current_plan>",
},
false,
))
}
if state.HasTodos {
if todoText := renderTodoList(state.Todos); todoText != "" {
contexts = append(contexts, newPromptContextMessage(
promptContextSourceStructuredTodoList,
modeladapter.Message{
Role: "user",
Content: "<todo_list>\n" + todoText + "\n</todo_list>",
},
false,
))
tailMessages = append(tailMessages, modeladapter.Message{
Role: "user",
Content: todoSectionReminderMessage,
})
}
}
return contexts, tailMessages, nil
}
func upsertCurrentPlanRegistryEntry(plans map[string]*agentv1.PlanRegistryEntry, planURI string) map[string]*agentv1.PlanRegistryEntry {
trimmedURI := strings.TrimSpace(planURI)
if trimmedURI == "" {
return clonePlanRegistryEntries(plans)
}
next := clonePlanRegistryEntries(plans)
if next == nil {
next = make(map[string]*agentv1.PlanRegistryEntry, 1)
}
key := "current"
if _, ok := next[key]; !ok && len(next) == 1 {
for existingKey := range next {
key = existingKey
}
}
next[key] = &agentv1.PlanRegistryEntry{
Id: key,
Path: trimmedURI,
}
return next
}
func clonePlanRegistryEntries(items map[string]*agentv1.PlanRegistryEntry) map[string]*agentv1.PlanRegistryEntry {
if len(items) == 0 {
return nil
}
cloned := make(map[string]*agentv1.PlanRegistryEntry, len(items))
for key, item := range items {
trimmedKey := strings.TrimSpace(key)
if trimmedKey == "" || item == nil {
continue
}
cloned[trimmedKey] = &agentv1.PlanRegistryEntry{
Id: strings.TrimSpace(item.GetId()),
Path: strings.TrimSpace(item.GetPath()),
}
if cloned[trimmedKey].GetId() == "" {
cloned[trimmedKey].Id = trimmedKey
}
}
if len(cloned) == 0 {
return nil
}
return cloned
}
func encodeConversationPlanBytes(plan string) []byte {
text := strings.TrimSpace(plan)
if text == "" {
return nil
}
payload, err := proto.Marshal(&agentv1.ConversationPlan{Plan: text})
if err != nil {
return nil
}
return payload
}
func encodeConversationTodoBytes(items []*agentv1.TodoItem) [][]byte {
if len(items) == 0 {
return nil
}
encoded := make([][]byte, 0, len(items))
for _, item := range items {
if item == nil {
continue
}
payload, err := proto.Marshal(cloneTodoItem(item))
if err != nil {
continue
}
encoded = append(encoded, payload)
}
if len(encoded) == 0 {
return nil
}
return encoded
}
func flattenCreatePlanTodos(args *agentv1.CreatePlanArgs) []*agentv1.TodoItem {
if args == nil {
return nil
}
if len(args.GetTodos()) > 0 {
return cloneTodoItems(args.GetTodos())
}
items := make([]*agentv1.TodoItem, 0, len(args.GetPhases()))
for _, phase := range args.GetPhases() {
if phase == nil {
continue
}
items = append(items, cloneTodoItems(phase.GetTodos())...)
}
return items
}
type decodedUpdateTodosArgs struct {
Args *agentv1.UpdateTodosArgs
MergeSet bool
}
func decodeUpdateTodosArgsJSON(raw []byte) (*agentv1.UpdateTodosArgs, error) {
decoded, err := decodeUpdateTodosArgsJSONWithPresence(raw)
if err != nil {
return nil, err
}
return decoded.Args, nil
}
func decodeUpdateTodosArgsJSONWithPresence(raw []byte) (decodedUpdateTodosArgs, error) {
payload, err := decodeJSONObject(raw)
if err != nil {
return decodedUpdateTodosArgs{}, err
}
todos, err := decodeTodoItemsValue(valueByAlias(payload, "todos"), false)
if err != nil {
return decodedUpdateTodosArgs{}, err
}
mergeValue, mergeSet := valueByAliasWithPresence(payload, "merge")
return decodedUpdateTodosArgs{
Args: &agentv1.UpdateTodosArgs{
Todos: todos,
Merge: boolValue(mergeValue),
},
MergeSet: mergeSet,
}, nil
}
func decodeReadTodosArgsJSON(raw []byte) (*agentv1.ReadTodosArgs, error) {
payload, err := decodeJSONObject(raw)
if err != nil {
return nil, err
}
statuses, err := decodeTodoStatusesValue(valueByAlias(payload, "status_filter", "statusFilter"))
if err != nil {
return nil, err
}
return &agentv1.ReadTodosArgs{
StatusFilter: statuses,
IdFilter: stringSliceValue(valueByAlias(payload, "id_filter", "idFilter")),
}, nil
}
func decodeJSONObject(raw []byte) (map[string]any, error) {
if len(raw) == 0 {
return map[string]any{}, nil
}
var payload map[string]any
if err := json.Unmarshal(raw, &payload); err != nil {
return nil, err
}
if payload == nil {
payload = map[string]any{}
}
return payload, nil
}
func decodeTodoItemsValue(value any, validate bool) ([]*agentv1.TodoItem, error) {
items, ok := value.([]any)
if !ok || len(items) == 0 {
return nil, nil
}
todos := make([]*agentv1.TodoItem, 0, len(items))
for _, rawItem := range items {
item, ok := rawItem.(map[string]any)
if !ok {
return nil, fmt.Errorf("todo item must be an object")
}
status, err := todoStatusFromValue(valueByAlias(item, "status"))
if err != nil {
return nil, err
}
todo := &agentv1.TodoItem{
Id: strings.TrimSpace(stringValue(valueByAlias(item, "id"))),
Content: strings.TrimSpace(stringValue(valueByAlias(item, "content"))),
Status: status,
CreatedAt: int64Value(valueByAlias(item, "created_at", "createdAt")),
UpdatedAt: int64Value(valueByAlias(item, "updated_at", "updatedAt")),
Dependencies: stringSliceValue(valueByAlias(item, "dependencies")),
}
if validate {
if todo.GetId() == "" {
return nil, fmt.Errorf("todo id is required")
}
if todo.GetContent() == "" {
return nil, fmt.Errorf("todo content is required")
}
}
todos = append(todos, todo)
}
return todos, nil
}
func decodeTodoStatusesValue(value any) ([]agentv1.TodoStatus, error) {
items, ok := value.([]any)
if !ok || len(items) == 0 {
return nil, nil
}
statuses := make([]agentv1.TodoStatus, 0, len(items))
for _, item := range items {
status, err := todoStatusFromValue(item)
if err != nil {
return nil, err
}
if status == agentv1.TodoStatus_TODO_STATUS_UNSPECIFIED {
continue
}
statuses = append(statuses, status)
}
return statuses, nil
}
func valueByAlias(values map[string]any, aliases ...string) any {
value, _ := valueByAliasWithPresence(values, aliases...)
return value
}
func valueByAliasWithPresence(values map[string]any, aliases ...string) (any, bool) {
for _, alias := range aliases {
if alias == "" {
continue
}
if value, ok := values[alias]; ok {
return value, true
}
}
return nil, false
}
func stringValue(value any) string {
text, ok := value.(string)
if !ok {
return ""
}
return strings.TrimSpace(text)
}
func boolValue(value any) bool {
switch item := value.(type) {
case bool:
return item
case string:
parsed, _ := strconv.ParseBool(strings.TrimSpace(item))
return parsed
default:
return false
}
}
func int64Value(value any) int64 {
switch item := value.(type) {
case float64:
return int64(item)
case float32:
return int64(item)
case int64:
return item
case int32:
return int64(item)
case int:
return int64(item)
case json.Number:
parsed, _ := item.Int64()
return parsed
case string:
parsed, _ := strconv.ParseInt(strings.TrimSpace(item), 10, 64)
return parsed
default:
return 0
}
}
func stringSliceValue(value any) []string {
items, ok := value.([]any)
if !ok || len(items) == 0 {
return nil
}
result := make([]string, 0, len(items))
for _, item := range items {
text := stringValue(item)
if text == "" {
continue
}
result = append(result, text)
}
return result
}
func todoStatusFromValue(value any) (agentv1.TodoStatus, error) {
switch item := value.(type) {
case nil:
return agentv1.TodoStatus_TODO_STATUS_UNSPECIFIED, nil
case float64:
return agentv1.TodoStatus(int32(item)), nil
case float32:
return agentv1.TodoStatus(int32(item)), nil
case int:
return agentv1.TodoStatus(item), nil
case int32:
return agentv1.TodoStatus(item), nil
case int64:
return agentv1.TodoStatus(item), nil
case string:
return todoStatusFromString(item)
default:
return agentv1.TodoStatus_TODO_STATUS_UNSPECIFIED, fmt.Errorf("unsupported todo status type %T", value)
}
}
func todoStatusFromString(raw string) (agentv1.TodoStatus, error) {
switch strings.ToLower(strings.TrimSpace(raw)) {
case "", "unspecified", "todo_status_unspecified":
return agentv1.TodoStatus_TODO_STATUS_UNSPECIFIED, nil
case "pending", "todo_status_pending":
return agentv1.TodoStatus_TODO_STATUS_PENDING, nil
case "in_progress", "in-progress", "inprogress", "todo_status_in_progress":
return agentv1.TodoStatus_TODO_STATUS_IN_PROGRESS, nil
case "completed", "complete", "todo_status_completed":
return agentv1.TodoStatus_TODO_STATUS_COMPLETED, nil
case "cancelled", "canceled", "todo_status_cancelled":
return agentv1.TodoStatus_TODO_STATUS_CANCELLED, nil
default:
return agentv1.TodoStatus_TODO_STATUS_UNSPECIFIED, fmt.Errorf("unsupported todo status %q", raw)
}
}
func todoStatusLabel(status agentv1.TodoStatus) string {
switch status {
case agentv1.TodoStatus_TODO_STATUS_PENDING:
return "pending"
case agentv1.TodoStatus_TODO_STATUS_IN_PROGRESS:
return "in_progress"
case agentv1.TodoStatus_TODO_STATUS_COMPLETED:
return "completed"
case agentv1.TodoStatus_TODO_STATUS_CANCELLED:
return "cancelled"
default:
return "unspecified"
}
}
func normalizeTodoItems(items []*agentv1.TodoItem, baseTime time.Time, validate bool) ([]*agentv1.TodoItem, error) {
if len(items) == 0 {
return nil, nil
}
baseMillis := timestampMillis(baseTime)
normalized := make([]*agentv1.TodoItem, 0, len(items))
for _, item := range items {
if item == nil {
continue
}
next := cloneTodoItem(item)
next.Id = strings.TrimSpace(next.GetId())
next.Content = strings.TrimSpace(next.GetContent())
if validate {
if next.GetId() == "" {
return nil, fmt.Errorf("todo id is required")
}
if next.GetContent() == "" {
return nil, fmt.Errorf("todo content is required")
}
}
if next.GetStatus() == agentv1.TodoStatus_TODO_STATUS_UNSPECIFIED {
next.Status = agentv1.TodoStatus_TODO_STATUS_PENDING
}
if next.GetCreatedAt() <= 0 {
next.CreatedAt = baseMillis
}
if next.GetUpdatedAt() <= 0 {
next.UpdatedAt = next.GetCreatedAt()
}
normalized = append(normalized, next)
}
return normalized, nil
}
func mergeTodoItems(existing []*agentv1.TodoItem, updates []*agentv1.TodoItem, baseTime time.Time) ([]*agentv1.TodoItem, error) {
result, err := normalizeTodoItems(existing, baseTime, false)
if err != nil {
return nil, err
}
indexByID := make(map[string]int, len(result))
for index, item := range result {
if item == nil || strings.TrimSpace(item.GetId()) == "" {
continue
}
indexByID[strings.TrimSpace(item.GetId())] = index
}
nowMillis := timestampMillis(baseTime)
for _, update := range updates {
if update == nil {
continue
}
id := strings.TrimSpace(update.GetId())
if id == "" {
return nil, fmt.Errorf("todo id is required")
}
content := strings.TrimSpace(update.GetContent())
if index, ok := indexByID[id]; ok {
current := result[index]
next := cloneTodoItem(current)
next.Id = id
if content != "" {
next.Content = content
}
if update.GetStatus() != agentv1.TodoStatus_TODO_STATUS_UNSPECIFIED {
next.Status = update.GetStatus()
}
if len(update.GetDependencies()) > 0 {
next.Dependencies = append([]string(nil), update.GetDependencies()...)
}
if next.GetContent() == "" {
return nil, fmt.Errorf("todo content is required")
}
if next.GetStatus() == agentv1.TodoStatus_TODO_STATUS_UNSPECIFIED {
next.Status = agentv1.TodoStatus_TODO_STATUS_PENDING
}
next.CreatedAt = current.GetCreatedAt()
if next.GetCreatedAt() <= 0 {
next.CreatedAt = nowMillis
}
next.UpdatedAt = nowMillis
result[index] = next
continue
}
next := cloneTodoItem(update)
next.Id = id
next.Content = content
if next.GetContent() == "" {
return nil, fmt.Errorf("todo content is required for new todo")
}
if next.GetStatus() == agentv1.TodoStatus_TODO_STATUS_UNSPECIFIED {
next.Status = agentv1.TodoStatus_TODO_STATUS_PENDING
}
if next.GetCreatedAt() <= 0 {
next.CreatedAt = nowMillis
}
if next.GetUpdatedAt() <= 0 {
next.UpdatedAt = next.GetCreatedAt()
}
indexByID[id] = len(result)
result = append(result, next)
}
return result, nil
}
func missingActiveTodoReplacementIDs(existing []*agentv1.TodoItem, replacement []*agentv1.TodoItem) []string {
existingIDs := make(map[string]struct{}, len(existing))
for _, item := range existing {
if item == nil || isTerminalTodoStatus(item.GetStatus()) {
continue
}
id := strings.TrimSpace(item.GetId())
if id == "" {
continue
}
existingIDs[id] = struct{}{}
}
if len(existingIDs) == 0 {
return nil
}
replacementIDs := make(map[string]struct{}, len(replacement))
for _, item := range replacement {
id := strings.TrimSpace(item.GetId())
if id == "" {
continue
}
replacementIDs[id] = struct{}{}
}
missing := make([]string, 0)
for id := range existingIDs {
if _, ok := replacementIDs[id]; !ok {
missing = append(missing, id)
}
}
sort.Strings(missing)
return missing
}
func isTerminalTodoStatus(status agentv1.TodoStatus) bool {
switch status {
case agentv1.TodoStatus_TODO_STATUS_COMPLETED,
agentv1.TodoStatus_TODO_STATUS_CANCELLED:
return true
default:
return false
}
}
func unsafeTodoReplaceError(missingIDs []string) string {
if len(missingIDs) == 0 {
return todoUnsafeReplaceBaseError
}
return fmt.Sprintf("%s; missing active todo ids: %s", todoUnsafeReplaceBaseError, strings.Join(missingIDs, ", "))
}
func filterTodoItems(items []*agentv1.TodoItem, statuses []agentv1.TodoStatus, ids []string) []*agentv1.TodoItem {
if len(items) == 0 {
return nil
}
statusFilter := make(map[agentv1.TodoStatus]struct{}, len(statuses))
for _, status := range statuses {
statusFilter[status] = struct{}{}
}
idFilter := make(map[string]struct{}, len(ids))
for _, id := range ids {
trimmed := strings.TrimSpace(id)
if trimmed == "" {
continue
}
idFilter[trimmed] = struct{}{}
}
filtered := make([]*agentv1.TodoItem, 0, len(items))
for _, item := range items {
if item == nil {
continue
}
if len(statusFilter) > 0 {
if _, ok := statusFilter[item.GetStatus()]; !ok {
continue
}
}
if len(idFilter) > 0 {
if _, ok := idFilter[strings.TrimSpace(item.GetId())]; !ok {
continue
}
}
filtered = append(filtered, cloneTodoItem(item))
}
return filtered
}
func renderTodoList(items []*agentv1.TodoItem) string {
if len(items) == 0 {
return ""
}
lines := make([]string, 0, len(items))
for _, item := range items {
if item == nil {
continue
}
lines = append(lines, fmt.Sprintf("- [%s] %s: %s", todoStatusLabel(item.GetStatus()), strings.TrimSpace(item.GetId()), strings.TrimSpace(item.GetContent())))
}
return strings.Join(lines, "\n")
}
func buildUpdateTodosToolCall(args *agentv1.UpdateTodosArgs, result *agentv1.UpdateTodosResult) *agentv1.ToolCall {
return &agentv1.ToolCall{
Tool: &agentv1.ToolCall_UpdateTodosToolCall{
UpdateTodosToolCall: &agentv1.UpdateTodosToolCall{
Args: args,
Result: result,
},
},
}
}
func buildUpdateTodosErrorResult(args *agentv1.UpdateTodosArgs, message string) *agentv1.UpdateTodosResult {
return &agentv1.UpdateTodosResult{
Result: &agentv1.UpdateTodosResult_Error{
Error: &agentv1.UpdateTodosError{Error: strings.TrimSpace(message)},
},
}
}
func buildReadTodosToolCall(args *agentv1.ReadTodosArgs, result *agentv1.ReadTodosResult) *agentv1.ToolCall {
return &agentv1.ToolCall{
Tool: &agentv1.ToolCall_ReadTodosToolCall{
ReadTodosToolCall: &agentv1.ReadTodosToolCall{
Args: args,
Result: result,
},
},
}
}
func summarizeUpdateTodosResult(result *agentv1.UpdateTodosResult) string {
if result == nil {
return "todo update missing"
}
switch item := result.GetResult().(type) {
case *agentv1.UpdateTodosResult_Success:
return fmt.Sprintf("todo update success count=%d merge=%t", item.Success.GetTotalCount(), item.Success.GetWasMerge())
case *agentv1.UpdateTodosResult_Error:
return strings.TrimSpace(item.Error.GetError())
default:
return "todo update finished"
}
}
func summarizeReadTodosResult(result *agentv1.ReadTodosResult) string {
if result == nil {
return "todo read missing"
}
switch item := result.GetResult().(type) {
case *agentv1.ReadTodosResult_Success:
return fmt.Sprintf("todo read success count=%d", item.Success.GetTotalCount())
case *agentv1.ReadTodosResult_Error:
return strings.TrimSpace(item.Error.GetError())
default:
return "todo read finished"
}
}
func buildTaskArgsFromJSON(raw []byte) *agentv1.TaskArgs {
payload, err := decodeJSONObject(raw)
if err != nil {
return &agentv1.TaskArgs{}
}
return buildTaskArgsFromMap(payload)
}
func buildTaskArgsFromMap(payload map[string]any) *agentv1.TaskArgs {
args := &agentv1.TaskArgs{
Description: strings.TrimSpace(stringValue(valueByAlias(payload, "description"))),
Prompt: strings.TrimSpace(stringValue(valueByAlias(payload, "prompt"))),
SubagentType: subagentTypeFromString(
strings.TrimSpace(stringValue(valueByAlias(payload, "subagent_type", "subagentType"))),
),
Attachments: stringSliceValue(valueByAlias(payload, "attachments")),
Mode: taskModeFromReadonly(boolValue(valueByAlias(payload, "readonly", "readOnly"))),
}
if model := strings.TrimSpace(stringValue(valueByAlias(payload, "model"))); model != "" {
args.Model = &model
}
if resume := strings.TrimSpace(stringValue(valueByAlias(payload, "resume"))); resume != "" {
args.Resume = &resume
}
if agentID := strings.TrimSpace(stringValue(valueByAlias(payload, "agent_id", "agentId"))); agentID != "" {
args.AgentId = &agentID
}
return args
}
func subagentTypeFromString(raw string) *agentv1.SubagentType {
switch strings.TrimSpace(raw) {
case "explore":
return &agentv1.SubagentType{Type: &agentv1.SubagentType_Explore{Explore: &agentv1.SubagentTypeExplore{}}}
case "browser-use", "browserUse":
return &agentv1.SubagentType{Type: &agentv1.SubagentType_BrowserUse{BrowserUse: &agentv1.SubagentTypeBrowserUse{}}}
case "shell":
return &agentv1.SubagentType{Type: &agentv1.SubagentType_Shell{Shell: &agentv1.SubagentTypeShell{}}}
case "":
return &agentv1.SubagentType{Type: &agentv1.SubagentType_Unspecified{Unspecified: &agentv1.SubagentTypeUnspecified{}}}
default:
return &agentv1.SubagentType{
Type: &agentv1.SubagentType_Custom{
Custom: &agentv1.SubagentTypeCustom{Name: strings.TrimSpace(raw)},
},
}
}
}
func taskModeFromReadonly(readonly bool) agentv1.TaskMode {
if readonly {
return agentv1.TaskMode_TASK_MODE_PLAN
}
return agentv1.TaskMode_TASK_MODE_AGENT
}
func cloneTodoItems(items []*agentv1.TodoItem) []*agentv1.TodoItem {
if len(items) == 0 {
return nil
}
cloned := make([]*agentv1.TodoItem, 0, len(items))
for _, item := range items {
if item == nil {
continue
}
cloned = append(cloned, cloneTodoItem(item))
}
return cloned
}
func cloneTodoItem(item *agentv1.TodoItem) *agentv1.TodoItem {
if item == nil {
return nil
}
return &agentv1.TodoItem{
Id: item.GetId(),
Content: item.GetContent(),
Status: item.GetStatus(),
CreatedAt: item.GetCreatedAt(),
UpdatedAt: item.GetUpdatedAt(),
Dependencies: append([]string(nil), item.GetDependencies()...),
}
}
func timestampMillis(baseTime time.Time) int64 {
now := baseTime
if now.IsZero() {
now = time.Now().UTC()
}
return now.UTC().UnixMilli()
}
@@ -0,0 +1,205 @@
package forwarder
import (
"encoding/json"
"strings"
"unicode/utf8"
"google.golang.org/protobuf/encoding/protojson"
"cursor/gen/agentv1"
"cursor/gen/aiserverv1"
modeladapter "cursor/internal/backend/agent/model"
promptengine "cursor/internal/backend/agent/prompt"
)
const (
estimatedTokensPerMessageOverhead = int64(8)
estimatedTokensPerToolCallOverhead = int64(6)
estimatedTokensPerImagePart = int64(1024)
)
func estimateCompiledPromptTokens(compiled CompiledConversation) int64 {
return estimateModelMessagesTokens(compiled.Messages) + estimateToolDescriptorsTokens(compiled.Tools)
}
func estimateModelMessagesTokens(messages []modeladapter.Message) int64 {
total := int64(0)
for _, item := range messages {
total += estimateModelMessageTokens(item)
}
return total
}
func estimateModelMessageTokens(item modeladapter.Message) int64 {
total := estimatedTokensPerMessageOverhead
total += estimateTextTokens(item.Role)
total += estimateTextTokens(item.Content)
total += estimateModelContentPartsTokens(item.Content, item.ContentParts)
total += estimateTextTokens(item.ReasoningContent)
total += estimateTextTokens(item.ReasoningSignature)
total += estimateTextTokens(item.ToolCallID)
total += estimateTextTokens(item.Name)
for _, toolCall := range item.ToolCalls {
total += estimatedTokensPerToolCallOverhead
total += estimateTextTokens(toolCall.ID)
total += estimateTextTokens(toolCall.Type)
total += estimateTextTokens(toolCall.Function.Name)
total += estimateTextTokens(toolCall.Function.Arguments)
}
return total
}
func estimateToolDescriptorsTokens(tools []json.RawMessage) int64 {
total := int64(0)
for _, item := range tools {
total += estimateTextTokens(string(item))
}
return total
}
func estimateContextItemTokens(item *aiserverv1.ContextItem) int64 {
if item == nil {
return 0
}
body, err := protojson.MarshalOptions{EmitUnpopulated: false}.Marshal(item)
if err != nil {
return estimateTextTokens(item.String())
}
return estimateTextTokens(string(body))
}
func estimateTextTokens(text string) int64 {
trimmed := strings.TrimSpace(text)
if trimmed == "" {
return 0
}
runeCount := utf8.RuneCountInString(trimmed)
if runeCount <= 0 {
return 0
}
estimated := int64((runeCount + 3) / 4)
estimated += int64(strings.Count(trimmed, "\n"))
if estimated < 1 {
return 1
}
return estimated
}
func estimateModelContentPartsTokens(content string, parts []modeladapter.ContentPart) int64 {
if len(parts) == 0 {
return 0
}
total := int64(0)
countText := strings.TrimSpace(content) == ""
for _, part := range parts {
switch strings.TrimSpace(strings.ToLower(part.Type)) {
case "", "text":
if countText {
total += estimateTextTokens(part.Text)
}
case "image":
total += estimatedTokensPerImagePart
if part.Image != nil {
total += estimateTextTokens(part.Image.MIMEType)
total += estimateTextTokens(part.Image.Path)
}
}
}
return total
}
func estimatePromptContentPartsTokens(content string, parts []promptengine.ContentPart) int64 {
if len(parts) == 0 {
return 0
}
total := int64(0)
countText := strings.TrimSpace(content) == ""
for _, part := range parts {
switch strings.TrimSpace(strings.ToLower(part.Type)) {
case "", "text":
if countText {
total += estimateTextTokens(part.Text)
}
case "image":
total += estimatedTokensPerImagePart
if part.Image != nil {
total += estimateTextTokens(part.Image.MIMEType)
total += estimateTextTokens(part.Image.Path)
}
}
}
return total
}
func estimateCheckpointPromptTokenBreakdown(compiled CompiledConversation, hasCompiled bool, usedTokens uint32, maxTokens uint32) *agentv1.PromptTokenBreakdownSnapshot {
if maxTokens == 0 {
return nil
}
categories := make([]*agentv1.PromptTokenBreakdownCategory, 0, 4)
if hasCompiled {
systemTokens, summaryTokens, conversationTokens := estimateMessageBreakdownTokens(compiled.Messages)
categories = appendPromptTokenBreakdownCategory(categories, "system_prompt", "System Prompt", systemTokens)
categories = appendPromptTokenBreakdownCategory(categories, "tools", "Tools", estimateToolDescriptorsTokens(compiled.Tools))
categories = appendPromptTokenBreakdownCategory(categories, "summarized_conversation", "Summarized Conversation", summaryTokens)
categories = appendPromptTokenBreakdownCategory(categories, "conversation", "Conversation", conversationTokens)
} else if usedTokens > 0 {
categories = appendPromptTokenBreakdownCategory(categories, "conversation", "Conversation", int64(usedTokens))
}
categoryTotal := int64(0)
for _, category := range categories {
categoryTotal += int64(category.GetEstimatedTokens())
}
totalUsedTokens := usedTokens
if categoryTotal > int64(totalUsedTokens) {
totalUsedTokens = clampInt64ToUint32(categoryTotal)
}
return &agentv1.PromptTokenBreakdownSnapshot{
TotalUsedTokens: totalUsedTokens,
MaxTokens: maxTokens,
Categories: categories,
}
}
func estimateMessageBreakdownTokens(messages []modeladapter.Message) (int64, int64, int64) {
systemTokens := int64(0)
summaryTokens := int64(0)
conversationTokens := int64(0)
for _, message := range messages {
tokens := estimateModelMessageTokens(message)
switch {
case strings.TrimSpace(message.Role) == "system":
systemTokens += tokens
case isConversationSummaryMessage(message):
summaryTokens += tokens
default:
conversationTokens += tokens
}
}
return systemTokens, summaryTokens, conversationTokens
}
func isConversationSummaryMessage(message modeladapter.Message) bool {
return strings.Contains(message.Content, "<conversation_summary>") || strings.Contains(message.Content, "</conversation_summary>")
}
func appendPromptTokenBreakdownCategory(categories []*agentv1.PromptTokenBreakdownCategory, id string, label string, estimatedTokens int64) []*agentv1.PromptTokenBreakdownCategory {
if estimatedTokens <= 0 {
return categories
}
return append(categories, &agentv1.PromptTokenBreakdownCategory{
Id: id,
Label: label,
EstimatedTokens: clampInt64ToUint32(estimatedTokens),
})
}
func clampInt64ToInt32(value int64) int32 {
if value <= 0 {
return 0
}
if value > int64(^uint32(0)>>1) {
return int32(^uint32(0) >> 1)
}
return int32(value)
}
+459
View File
@@ -0,0 +1,459 @@
package forwarder
import (
"encoding/json"
"fmt"
"strings"
"time"
"cursor/gen/agentv1"
modeladapter "cursor/internal/backend/agent/model"
promptengine "cursor/internal/backend/agent/prompt"
"google.golang.org/protobuf/proto"
)
type turnUsageSnapshot struct {
Provider string
Model string
InputTokens int64
OutputTokens int64
CacheReadTokens int64
CacheWriteTokens int64
UsagePresent bool
CacheReadPresent bool
CacheWritePresent bool
}
func (snapshot turnUsageSnapshot) hasAny() bool {
return snapshot.InputTokens > 0 ||
snapshot.OutputTokens > 0 ||
snapshot.CacheReadTokens > 0 ||
snapshot.CacheWriteTokens > 0
}
func (snapshot turnUsageSnapshot) cacheUsageComplete() bool {
return snapshot.CacheReadPresent && snapshot.CacheWritePresent
}
func (snapshot turnUsageSnapshot) promptTokensTotal() int64 {
return nonNegativeInt64(snapshot.InputTokens) +
nonNegativeInt64(snapshot.CacheReadTokens) +
nonNegativeInt64(snapshot.CacheWriteTokens)
}
func (snapshot turnUsageSnapshot) requestTokensTotal() int64 {
return snapshot.promptTokensTotal() + nonNegativeInt64(snapshot.OutputTokens)
}
func (service *Service) importConversationState(item *ConversationFile, state *agentv1.ConversationStateStructure) ([]HistoryEntry, error) {
if item == nil || state == nil {
return nil, nil
}
item.TokenDetailsUsedTokens = state.GetTokenDetails().GetUsedTokens()
entries := make([]HistoryEntry, 0, 2)
if messages, err := importedConversationStateModelMessages(state); err != nil {
return nil, err
} else {
for _, message := range messages {
entry, ok, err := newModelMessageEntry(0, "", message)
if err != nil {
return nil, err
}
if ok {
entries = append(entries, entry)
}
}
}
if len(entries) == 0 {
summary, ok, err := importedConversationStateSummary(state)
if err != nil {
return nil, err
}
if ok {
payload, err := json.Marshal(compactionSummaryEntryPayload{
Summary: strings.TrimSpace(summary),
Trigger: "imported_conversation_state",
})
if err != nil {
return nil, fmt.Errorf("encode imported summary context: %w", err)
}
entries = append(entries, HistoryEntry{
TurnSeq: 0,
Role: "system",
Kind: "compacted_summary",
Payload: payload,
})
}
}
runtimeState, ok, err := runtimeStatePayloadFromConversationState(state)
if err != nil {
return nil, err
}
if ok {
payload, err := json.Marshal(runtimeState)
if err != nil {
return nil, fmt.Errorf("encode imported runtime state context: %w", err)
}
entries = append(entries, HistoryEntry{
TurnSeq: 0,
Role: "system",
Kind: "runtime_state",
Payload: payload,
})
}
return entries, nil
}
func importedConversationStateModelMessages(state *agentv1.ConversationStateStructure) ([]modeladapter.Message, error) {
if state == nil {
return nil, nil
}
if len(state.GetRootPromptMessagesJson()) > 0 {
decoded, err := promptengine.DecodeReplayMessages(state.GetRootPromptMessagesJson())
if err != nil {
return nil, fmt.Errorf("decode imported replay messages: %w", err)
}
decoded = restoreImportedReplayUserMessages(decoded, state.GetTurns())
decoded = filterLegacyPlainWriteReplay(decoded)
decoded = filterInternalPromptContextReplay(decoded)
messages := make([]modeladapter.Message, 0, len(decoded))
for _, item := range decoded {
messages = append(messages, toModelMessage(item))
}
return normalizeReplayMessageSequence(messages), nil
}
if len(state.GetSummary()) > 0 {
return nil, nil
}
if len(state.GetTurns()) == 0 {
return nil, nil
}
messages := make([]modeladapter.Message, 0, len(state.GetTurns())*2)
for _, rawTurn := range state.GetTurns() {
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)
}
agentTurn := turn.GetAgentConversationTurn()
if agentTurn == nil {
continue
}
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))
}
}
}
return normalizeReplayMessageSequence(messages), nil
}
func importedConversationStateSummary(state *agentv1.ConversationStateStructure) (string, bool, error) {
if state == nil || len(state.GetSummary()) == 0 {
return "", false, nil
}
item := &agentv1.ConversationSummary{}
if err := proto.Unmarshal(state.GetSummary(), item); err != nil {
return "", false, fmt.Errorf("decode imported summary: %w", err)
}
text := strings.TrimSpace(item.GetSummary())
return text, text != "", nil
}
func newModelMessageEntry(turnSeq int64, requestID string, message modeladapter.Message) (HistoryEntry, bool, error) {
message.Role = strings.TrimSpace(message.Role)
if message.Role == "" {
return HistoryEntry{}, false, nil
}
if strings.TrimSpace(message.Content) == "" &&
len(message.ContentParts) == 0 &&
len(message.ToolCalls) == 0 &&
strings.TrimSpace(message.ToolCallID) == "" &&
!hasReplayableReasoningPayload(message.ReasoningContent, message.ReasoningSignature, message.ReasoningSignatureSource) &&
len(message.OpenAIResponsesReasoningSummary) == 0 {
return HistoryEntry{}, false, nil
}
payload, err := json.Marshal(modelMessageEntryPayload{Message: message})
if err != nil {
return HistoryEntry{}, false, fmt.Errorf("encode imported model message context: %w", err)
}
return HistoryEntry{
TurnSeq: turnSeq,
RequestID: strings.TrimSpace(requestID),
Role: message.Role,
Kind: "model_message",
Payload: payload,
}, true, nil
}
func runtimeStatePayloadFromConversationState(state *agentv1.ConversationStateStructure) (runtimeStateEntryPayload, bool, error) {
if state == nil {
return runtimeStateEntryPayload{}, false, nil
}
payload := runtimeStateEntryPayload{
Plans: clonePlanRegistryEntries(state.GetPlans()),
}
if len(state.GetPlan()) > 0 {
plan := &agentv1.ConversationPlan{}
if err := proto.Unmarshal(state.GetPlan(), plan); err != nil {
return runtimeStateEntryPayload{}, false, fmt.Errorf("decode imported plan: %w", err)
}
payload.PlanText = strings.TrimSpace(plan.GetPlan())
}
if len(state.GetTodos()) > 0 {
todos := make([]*agentv1.TodoItem, 0, len(state.GetTodos()))
for _, raw := range state.GetTodos() {
if len(raw) == 0 {
continue
}
item := &agentv1.TodoItem{}
if err := proto.Unmarshal(raw, item); err != nil {
return runtimeStateEntryPayload{}, false, fmt.Errorf("decode imported todo: %w", err)
}
todos = append(todos, cloneTodoItem(item))
}
payload.Todos = todos
}
ok := strings.TrimSpace(payload.PlanText) != "" || len(payload.Plans) > 0 || len(payload.Todos) > 0
return payload, ok, nil
}
func (service *Service) updateConversationTokenState(stream *ActiveStream, conversationID string, usage turnUsageSnapshot, modelCallID string, finalizeAutoCompaction bool) error {
if service == nil || !usage.hasAny() {
return nil
}
now := time.Now().UTC()
autoCompactionReserveTokens := int64(compactionAutoReserveTokens)
if finalizeAutoCompaction {
autoCompactionReserveTokens = service.resolveCompactionReserveTokens(activeStreamModelID(stream))
if autoCompactionReserveTokens <= 0 {
autoCompactionReserveTokens = compactionAutoReserveTokens
}
}
_, err := service.updateConversationMetaAndCheckpoint(stream, conversationID, func(item *ConversationFile) error {
if item == nil {
return nil
}
promptTokensTotal := usage.promptTokensTotal()
if promptTokensTotal > 0 {
item.TokenDetailsUsedTokens = clampInt64ToUint32(promptTokensTotal)
}
if item.TokenDetailsMaxTokens == 0 {
item.TokenDetailsMaxTokens = projectedConversationMaxTokens
}
if finalizeAutoCompaction {
updateConversationAutoCompactionState(item, usage.requestTokensTotal(), autoCompactionReserveTokens, modelCallID, now)
}
return nil
})
return err
}
func activeStreamModelID(stream *ActiveStream) string {
if stream == nil {
return ""
}
stream.mu.Lock()
defer stream.mu.Unlock()
return stream.ModelID
}
func updateConversationAutoCompactionState(conversation *ConversationFile, promptTokensTotal int64, reserveTokens int64, modelCallID string, triggeredAt time.Time) {
if conversation == nil {
return
}
if promptTokensTotal <= 0 {
clearConversationAutoCompactionState(conversation)
return
}
contextWindowTokens := int64(conversation.TokenDetailsMaxTokens)
if contextWindowTokens <= 0 {
contextWindowTokens = projectedConversationMaxTokens
}
if reserveTokens <= 0 {
reserveTokens = compactionAutoReserveTokens
}
remainingTokens := contextWindowTokens - promptTokensTotal
if remainingTokens > reserveTokens {
clearConversationAutoCompactionState(conversation)
return
}
conversation.AutoCompactionPending = true
conversation.AutoCompactionPromptTokens = promptTokensTotal
conversation.AutoCompactionReserveTokens = reserveTokens
conversation.AutoCompactionTriggeredAt = triggeredAt.UTC().Format(time.RFC3339Nano)
conversation.AutoCompactionSourceModelCallID = strings.TrimSpace(modelCallID)
}
func clearConversationAutoCompactionState(conversation *ConversationFile) {
if conversation == nil {
return
}
conversation.AutoCompactionPending = false
conversation.AutoCompactionPromptTokens = 0
conversation.AutoCompactionReserveTokens = 0
conversation.AutoCompactionTriggeredAt = ""
conversation.AutoCompactionSourceModelCallID = ""
}
func (service *Service) recordTurnUsageSnapshot(stream *ActiveStream, conversationID string, turnSeq int64, requestID string, modelCallID string, status string, usage turnUsageSnapshot, errorText string, turnFinalized bool) error {
_ = turnFinalized
if service == nil || strings.TrimSpace(requestID) == "" {
return nil
}
modelID := ""
modelName := ""
startedAt := time.Time{}
lastEventAt := time.Now().UTC()
if stream != nil {
stream.mu.Lock()
modelID = strings.TrimSpace(stream.ModelID)
modelName = strings.TrimSpace(stream.ModelName)
startedAt = stream.CreatedAt
if !stream.UpdatedAt.IsZero() {
lastEventAt = stream.UpdatedAt
}
stream.mu.Unlock()
}
if strings.TrimSpace(modelName) == "" {
modelName = modelID
}
provider := strings.TrimSpace(usage.Provider)
if strings.TrimSpace(usage.Model) != "" {
modelName = strings.TrimSpace(usage.Model)
}
effectiveModelCallID := firstNonEmpty(strings.TrimSpace(modelCallID), strings.TrimSpace(requestID))
if service.usageStore != nil {
if err := service.usageStore.UpsertEvent(usageFileEvent{
EventID: usageEventID(requestID, effectiveModelCallID),
Kind: usageEventKindProvider,
At: lastEventAt,
InputTokens: usage.InputTokens,
OutputTokens: usage.OutputTokens,
CacheReadTokens: usage.CacheReadTokens,
CacheWriteTokens: usage.CacheWriteTokens,
UsagePresent: usage.UsagePresent,
}); err != nil {
return err
}
}
if strings.TrimSpace(conversationID) != "" {
_, err := service.updateConversationMetaAndCheckpoint(stream, conversationID, func(item *ConversationFile) error {
if item == nil {
return nil
}
item.LastProviderCall = &ConversationProviderCall{
RequestID: strings.TrimSpace(requestID),
ModelCallID: effectiveModelCallID,
Provider: provider,
Model: modelName,
Status: strings.TrimSpace(status),
ErrorText: strings.TrimSpace(errorText),
UpdatedAt: lastEventAt,
}
return nil
})
if err != nil {
return err
}
}
_ = turnSeq
_ = startedAt
return nil
}
func (service *Service) recordTurnFinalizedSnapshot(stream *ActiveStream, conversationID string, turnSeq int64, requestID string, status string, errorText string) error {
_ = stream
_ = errorText
if service == nil || service.usageStore == nil || strings.TrimSpace(requestID) == "" {
return nil
}
eventID := turnUsageEventID(conversationID, turnSeq, requestID)
usage := usageFileEvent{}
if aggregate, ok, err := service.usageStore.LookupEvent(strings.TrimSpace(requestID)); err != nil {
return err
} else if ok {
usage = aggregate
}
return service.usageStore.UpsertEvent(usageFileEvent{
EventID: eventID,
Kind: usageEventKindTurn,
Status: normalizeUsageTurnStatus(status),
At: time.Now().UTC(),
InputTokens: usage.InputTokens,
OutputTokens: usage.OutputTokens,
CacheReadTokens: usage.CacheReadTokens,
CacheWriteTokens: usage.CacheWriteTokens,
UsagePresent: usage.UsagePresent,
})
}
func normalizeUsageTurnStatus(status string) string {
if strings.TrimSpace(status) == usageTurnStatusDone {
return usageTurnStatusDone
}
return strings.TrimSpace(status)
}
func turnUsageEventID(conversationID string, turnSeq int64, requestID string) string {
conversationID = strings.TrimSpace(conversationID)
requestID = strings.TrimSpace(requestID)
if conversationID != "" && turnSeq > 0 {
return fmt.Sprintf("turn::%s::%d", conversationID, turnSeq)
}
return "turn::" + requestID
}
func usageEventID(requestID string, modelCallID string) string {
requestID = strings.TrimSpace(requestID)
modelCallID = strings.TrimSpace(modelCallID)
if modelCallID == "" || modelCallID == requestID {
return requestID
}
return requestID + "::" + modelCallID
}
func clampInt64ToUint32(value int64) uint32 {
if value <= 0 {
return 0
}
if value > int64(^uint32(0)) {
return ^uint32(0)
}
return uint32(value)
}
func nonNegativeInt64(value int64) int64 {
if value < 0 {
return 0
}
return value
}
func maxPositiveInt64(values ...int64) int64 {
maxValue := int64(0)
for _, value := range values {
if value > maxValue {
maxValue = value
}
}
return maxValue
}
+295
View File
@@ -0,0 +1,295 @@
// tool_catalog.go 负责从静态 prompt 资产中装载并筛选 canonical tool catalog。
package forwarder
import (
"encoding/json"
"fmt"
"strings"
"cursor/gen/agentv1"
promptassets "cursor/prompt"
)
type DefaultToolCatalog struct {
}
// NewToolCatalog 创建默认工具目录实现。
func NewToolCatalog() *DefaultToolCatalog {
return &DefaultToolCatalog{}
}
// Load 按 mode 读取工具资产,并过滤出当前阶段真正允许暴露的工具。
func (catalog *DefaultToolCatalog) Load(mode agentv1.AgentMode, subagentTypeName string) ([]json.RawMessage, []string, error) {
assetMode, err := toolAssetModeForConversation(mode, subagentTypeName)
if err != nil {
return nil, nil, err
}
rawTools, err := promptassets.ReadTools(assetMode)
if err != nil {
return nil, nil, err
}
var items []json.RawMessage
if err := json.Unmarshal(rawTools, &items); err != nil {
return nil, nil, fmt.Errorf("decode tools asset failed: %w", err)
}
filtered := make([]json.RawMessage, 0, len(items))
names := make([]string, 0, len(items))
for _, item := range items {
name, err := extractToolName(item)
if err != nil {
return nil, nil, err
}
if !isToolAllowedInMode(mode, subagentTypeName, name) {
continue
}
filtered = append(filtered, item)
names = append(names, name)
}
return filtered, names, nil
}
var agentModeToolNames = map[string]struct{}{
"AskQuestion": {},
"CallMcpTool": {},
"Delete": {},
"FetchMcpResource": {},
"GenerateImage": {},
"Glob": {},
"Grep": {},
"Ls": {},
"PatchEdit": {},
"Read": {},
"ReadLints": {},
"Shell": {},
"AwaitShell": {},
"WriteShellStdin": {},
"ForceBackgroundShell": {},
"SwitchMode": {},
"Task": {},
"TodoWrite": {},
"WebFetch": {},
"WebSearch": {},
"Write": {},
}
var multitaskModeToolNames = map[string]struct{}{
"AskQuestion": {},
"CallMcpTool": {},
"Delete": {},
"FetchMcpResource": {},
"GenerateImage": {},
"Glob": {},
"Grep": {},
"Ls": {},
"PatchEdit": {},
"Read": {},
"ReadLints": {},
"Shell": {},
"AwaitShell": {},
"WriteShellStdin": {},
"ForceBackgroundShell": {},
"SwitchMode": {},
"Task": {},
"TodoWrite": {},
"WebFetch": {},
"WebSearch": {},
"Write": {},
}
var debugModeToolNames = map[string]struct{}{
"AskQuestion": {},
"CallMcpTool": {},
"Delete": {},
"FetchMcpResource": {},
"Glob": {},
"Grep": {},
"Ls": {},
"PatchEdit": {},
"Read": {},
"ReadLints": {},
"Shell": {},
"AwaitShell": {},
"WriteShellStdin": {},
"ForceBackgroundShell": {},
"Task": {},
"TodoWrite": {},
"WebFetch": {},
"WebSearch": {},
"Write": {},
}
var askModeToolNames = map[string]struct{}{
"AskQuestion": {},
"CallMcpTool": {},
"Delete": {},
"FetchMcpResource": {},
"Glob": {},
"Grep": {},
"Ls": {},
"PatchEdit": {},
"Read": {},
"ReadLints": {},
"Shell": {},
"AwaitShell": {},
"WriteShellStdin": {},
"ForceBackgroundShell": {},
"Task": {},
"TodoWrite": {},
"WebFetch": {},
"WebSearch": {},
"Write": {},
}
var planModeToolNames = map[string]struct{}{
"AskQuestion": {},
"CallMcpTool": {},
"CreatePlan": {},
"FetchMcpResource": {},
"Glob": {},
"Grep": {},
"Ls": {},
"Read": {},
"ReadLints": {},
"Shell": {},
"AwaitShell": {},
"WriteShellStdin": {},
"ForceBackgroundShell": {},
"Task": {},
"TodoWrite": {},
"WebFetch": {},
"WebSearch": {},
}
var childConversationDisallowedAgentToolNames = map[string]struct{}{
"AskQuestion": {},
}
func supportedToolNamesForMode(mode agentv1.AgentMode) map[string]struct{} {
switch normalizeMode(mode) {
case agentv1.AgentMode_AGENT_MODE_AGENT:
return agentModeToolNames
case agentv1.AgentMode_AGENT_MODE_ASK:
return askModeToolNames
case agentv1.AgentMode_AGENT_MODE_PLAN:
return planModeToolNames
case agentv1.AgentMode_AGENT_MODE_DEBUG:
return debugModeToolNames
case agentv1.AgentMode_AGENT_MODE_MULTITASK:
return multitaskModeToolNames
default:
return nil
}
}
func isToolAllowedInMode(mode agentv1.AgentMode, subagentTypeName string, toolName string) bool {
trimmedToolName := strings.TrimSpace(toolName)
if trimmedToolName == "" {
return false
}
if isChildConversationSubagentTypeName(subagentTypeName) {
if _, disallowed := childConversationDisallowedAgentToolNames[trimmedToolName]; disallowed {
return false
}
_, ok := agentModeToolNames[trimmedToolName]
return ok
}
supported := supportedToolNamesForMode(mode)
if supported == nil {
return false
}
_, ok := supported[trimmedToolName]
return ok
}
func isChildConversationSubagentTypeName(subagentTypeName string) bool {
return strings.TrimSpace(subagentTypeName) != ""
}
func selectToolsByOrderedNames(items []json.RawMessage, orderedNames []string) ([]json.RawMessage, []string, error) {
byName := make(map[string]json.RawMessage, len(items))
for _, item := range items {
name, err := extractToolName(item)
if err != nil {
return nil, nil, err
}
if _, exists := byName[name]; !exists {
byName[name] = item
}
}
filtered := make([]json.RawMessage, 0, len(orderedNames))
names := make([]string, 0, len(orderedNames))
for _, name := range orderedNames {
item, ok := byName[name]
if !ok {
return nil, nil, fmt.Errorf("tool descriptor %q not found in prompt asset", name)
}
filtered = append(filtered, item)
names = append(names, name)
}
return filtered, names, nil
}
func toolAssetModeForConversation(mode agentv1.AgentMode, subagentTypeName string) (promptassets.Mode, error) {
if isChildConversationSubagentTypeName(subagentTypeName) {
return promptassets.ModeAgent, nil
}
return mapPromptMode(mode)
}
func promptAssetModeForConversation(mode agentv1.AgentMode, subagentTypeName string) (promptassets.Mode, error) {
if isChildConversationSubagentTypeName(subagentTypeName) {
return promptassets.ModeSubagent, nil
}
return mapPromptMode(mode)
}
// mapPromptMode 把协议 mode 映射为静态 prompt 资产对应的目录名。
func mapPromptMode(mode agentv1.AgentMode) (promptassets.Mode, error) {
switch normalizeMode(mode) {
case agentv1.AgentMode_AGENT_MODE_AGENT:
return promptassets.ModeAgent, nil
case agentv1.AgentMode_AGENT_MODE_ASK:
return promptassets.ModeAsk, nil
case agentv1.AgentMode_AGENT_MODE_PLAN:
return promptassets.ModePlan, nil
case agentv1.AgentMode_AGENT_MODE_DEBUG:
return promptassets.ModeDebug, nil
case agentv1.AgentMode_AGENT_MODE_MULTITASK:
return promptassets.ModeMultitask, nil
default:
return "", fmt.Errorf("unsupported prompt asset mode: %s", mode.String())
}
}
// extractToolName 从原始 tool descriptor JSON 中提取函数名。
func extractToolName(raw json.RawMessage) (string, error) {
var wrapper struct {
Function struct {
Name string `json:"name"`
} `json:"function"`
}
if err := json.Unmarshal(raw, &wrapper); err != nil {
return "", fmt.Errorf("decode tool descriptor failed: %w", err)
}
name := strings.TrimSpace(wrapper.Function.Name)
if name == "" {
return "", fmt.Errorf("tool descriptor name is required")
}
return name, nil
}
// sanitizePromptAsset 去掉资产文件中的说明性标题,只保留真正的 prompt 文本。
func sanitizePromptAsset(text string, modelName string) string {
lines := strings.Split(text, "\n")
filtered := make([]string, 0, len(lines))
for _, line := range lines {
trimmed := strings.TrimSpace(line)
switch trimmed {
case "# 通用系统提示词", "# 模式静态补充", "---":
continue
default:
filtered = append(filtered, line)
}
}
return promptassets.RenderPromptTemplate(strings.TrimSpace(strings.Join(filtered, "\n")), modelName)
}
@@ -0,0 +1,111 @@
package forwarder
import (
"errors"
"fmt"
"strings"
"google.golang.org/protobuf/encoding/protojson"
"cursor/gen/agentv1"
runtimecore "cursor/internal/backend/agent/core"
)
type recoverableToolInvocationError struct {
cause error
}
func (err recoverableToolInvocationError) Error() string {
if err.cause == nil {
return "recoverable tool invocation error"
}
return err.cause.Error()
}
func (err recoverableToolInvocationError) Unwrap() error {
return err.cause
}
func newRecoverableToolInvocationError(cause error) error {
if cause == nil {
return nil
}
return recoverableToolInvocationError{cause: cause}
}
func recoverableToolInvocationCause(err error) (error, bool) {
if err == nil {
return nil, false
}
var recoverable recoverableToolInvocationError
if !errors.As(err, &recoverable) {
return nil, false
}
if recoverable.cause == nil {
return err, true
}
return recoverable.cause, true
}
func (service *Service) completePreDispatchToolError(
stream *ActiveStream,
invocation runtimecore.ToolInvocation,
startedToolCall *agentv1.ToolCall,
startedHistoryAppended bool,
startedEmitted bool,
cause error,
) error {
if stream == nil || cause == nil {
return cause
}
if strings.TrimSpace(invocation.CallID) == "" {
return cause
}
if startedToolCall == nil {
startedToolCall = buildStartedToolCall(invocation)
}
if !startedHistoryAppended && startedToolCall != nil {
toolCallPayload, err := protojson.Marshal(startedToolCall)
if err != nil {
return err
}
if _, err := service.appendConversationEntries(stream, stream.ConversationID, []HistoryEntry{
newToolCallEntryWithProviderMetadata(stream.TurnSeq, stream.RequestID, invocation.CallID, invocation.ToolName, invocation.ReasoningContent, invocation.ReasoningSignature, invocation.ReasoningSignatureSource, invocation.ReasoningProviderItemID, invocation.ReasoningProviderStatus, invocation.ReasoningProviderSummary, invocation.ProviderItemID, invocation.ProviderCallID, invocation.ProviderStatus, toolCallPayload),
}); err != nil {
return err
}
}
if !startedEmitted {
if err := service.broker.Publish(stream.RequestID, StreamEvent{
Message: buildToolCallStartedMessage(invocation.CallID, invocation.ModelCallID, startedToolCall),
}); err != nil {
return err
}
}
resultText := formatPreDispatchToolError(invocation, cause)
if err := service.appendToolResult(stream, invocation.CallID, strings.TrimSpace(invocation.ToolName), invocation.ArgsJSON, resultText, invocation.ReasoningContent, nil); err != nil {
return err
}
if err := service.publishToolCallCompleted(stream.RequestID, invocation.CallID, invocation.ModelCallID, nil); err != nil {
return err
}
if err := service.syncSummaryCarryForward(stream.ConversationID, stream.RequestID, invocation.ModelCallID); err != nil {
return err
}
if err := service.publishCheckpoint(stream.RequestID, stream.ConversationID); err != nil {
return err
}
return service.reconcileStream(stream)
}
func formatPreDispatchToolError(invocation runtimecore.ToolInvocation, cause error) string {
toolName := strings.TrimSpace(invocation.ToolName)
if toolName == "" {
toolName = "Tool"
}
message := "unknown error"
if cause != nil && strings.TrimSpace(cause.Error()) != "" {
message = strings.TrimSpace(cause.Error())
}
return fmt.Sprintf("%s error: %s", toolName, message)
}
@@ -0,0 +1,394 @@
package forwarder
import (
"encoding/json"
"fmt"
"strings"
"unicode/utf8"
)
const (
projectedReplayKiB = 1024
projectedReadReplayLimit = 64 * projectedReplayKiB
projectedShellReplayLimit = 128 * projectedReplayKiB
projectedShellStreamLimit = 16 * projectedReplayKiB
projectedShellInterleavedLimit = 32 * projectedReplayKiB
projectedGrepReplayLimit = 32 * projectedReplayKiB
projectedEditReplayLimit = 32 * projectedReplayKiB
projectedPatchEditReplayLimit = 4 * projectedReplayKiB
projectedWebFetchReplayLimit = 32 * projectedReplayKiB
projectedWebSearchReplayLimit = 16 * projectedReplayKiB
projectedMcpReplayLimit = 32 * projectedReplayKiB
)
func limitProjectedToolResultReplay(toolName string, content string, resultText string, fromStoredToolCall bool, historical bool) string {
if compacted, ok := compactProjectedGenerateImageResultReplay(toolName, content, resultText); ok {
return compacted
}
if historical {
if compacted, ok := compactHistoricalEditErrorReplay(toolName, content); ok {
content = compacted
fromStoredToolCall = false
}
}
if compacted, ok := compactProjectedEditToolResultReplay(toolName, content); ok {
content = compacted
fromStoredToolCall = false
}
if compacted, ok := compactProjectedShellToolResultReplay(toolName, content); ok {
content = compacted
fromStoredToolCall = false
}
limit, ok := projectedToolReplayLimit(toolName)
if !ok {
return strings.TrimSpace(content)
}
content = strings.TrimSpace(content)
if len(content) <= limit {
return content
}
if fromStoredToolCall {
fallback := strings.TrimSpace(resultText)
notice := fmt.Sprintf("[tool result replay truncated: stored ToolCall result exceeded %d bytes]", limit)
if fallback == "" {
fallback = notice
} else {
fallback += "\n\n" + notice
}
return truncateProjectedReplayText(toolName, fallback, limit)
}
return truncateProjectedReplayText(toolName, content, limit)
}
func projectedToolReplayLimit(toolName string) (int, bool) {
switch strings.TrimSpace(toolName) {
case "GenerateImage":
return projectedWebSearchReplayLimit, true
case "Read":
return projectedReadReplayLimit, true
case "Shell":
return projectedShellReplayLimit, true
case "Grep":
return projectedGrepReplayLimit, true
case "PatchEdit", "PatchEditLines", "PatchEditSpan":
return projectedPatchEditReplayLimit, true
case "Edit", "Write":
return projectedEditReplayLimit, true
case "WebFetch":
return projectedWebFetchReplayLimit, true
case "WebSearch":
return projectedWebSearchReplayLimit, true
case "CallMcpTool", "FetchMcpResource", "ListMcpResources":
return projectedMcpReplayLimit, true
default:
return 0, false
}
}
func compactProjectedGenerateImageResultReplay(toolName string, content string, resultText string) (string, bool) {
if strings.TrimSpace(toolName) != "GenerateImage" {
return "", false
}
fallback := strings.TrimSpace(resultText)
trimmed := strings.TrimSpace(content)
if trimmed == "" || trimmed == "{}" || trimmed == "null" {
if fallback == "" {
return "", false
}
return fallback, true
}
var payload any
if err := json.Unmarshal([]byte(trimmed), &payload); err != nil {
if fallback == "" {
return truncateProjectedReplayText("GenerateImage", trimmed, projectedWebSearchReplayLimit), true
}
return fallback, true
}
if compactGenerateImagePayload(payload) {
encoded, err := json.Marshal(payload)
if err == nil {
return string(encoded), true
}
}
if fallback != "" {
return fallback, true
}
return trimmed, true
}
func compactGenerateImagePayload(value any) bool {
changed := false
switch item := value.(type) {
case map[string]any:
for key, child := range item {
if text, ok := child.(string); ok {
switch key {
case "image_data", "imageData":
item[key] = fmt.Sprintf("[base64 image data omitted from replay; bytes=%d]", len(strings.TrimSpace(text)))
changed = true
continue
}
}
if compactGenerateImagePayload(child) {
changed = true
}
}
case []any:
for _, child := range item {
if compactGenerateImagePayload(child) {
changed = true
}
}
}
return changed
}
func patchEditReplayLimit(toolName string) int {
switch strings.TrimSpace(toolName) {
case "PatchEdit", "PatchEditLines", "PatchEditSpan":
return projectedPatchEditReplayLimit
default:
return projectedEditReplayLimit
}
}
func compactProjectedShellToolResultReplay(toolName string, content string) (string, bool) {
if strings.TrimSpace(toolName) != "Shell" {
return "", false
}
trimmed := strings.TrimSpace(content)
if trimmed == "" || trimmed == "{}" || trimmed == "null" {
return "", false
}
var payload any
if err := json.Unmarshal([]byte(trimmed), &payload); err != nil {
return "", false
}
if !compactProjectedShellFields(payload) {
return "", false
}
encoded, err := json.Marshal(payload)
if err != nil {
return "", false
}
return string(encoded), true
}
func compactProjectedShellFields(value any) bool {
changed := false
switch item := value.(type) {
case map[string]any:
for key, child := range item {
if text, ok := child.(string); ok {
switch key {
case "stdout", "stderr":
next := truncateProjectedReplayTextMiddle("Shell "+key, text, projectedShellStreamLimit)
if next != text {
item[key] = next
changed = true
}
continue
case "interleaved_output", "interleavedOutput":
next := truncateProjectedReplayTextMiddle("Shell interleaved output", text, projectedShellInterleavedLimit)
if next != text {
item[key] = next
changed = true
}
continue
}
}
if compactProjectedShellFields(child) {
changed = true
}
}
case []any:
for _, child := range item {
if compactProjectedShellFields(child) {
changed = true
}
}
}
return changed
}
func compactProjectedEditToolResultReplay(toolName string, content string) (string, bool) {
switch strings.TrimSpace(toolName) {
case "PatchEdit", "PatchEditLines", "PatchEditSpan", "Edit", "Write":
default:
return "", false
}
trimmed := strings.TrimSpace(content)
if trimmed == "" || trimmed == "{}" || trimmed == "null" {
return "", false
}
var payload map[string]any
if err := json.Unmarshal([]byte(trimmed), &payload); err != nil {
return "", false
}
success, ok := payload["success"].(map[string]any)
if !ok {
return "", false
}
diffString := firstJSONText(success, "diff_string", "diffString")
beforeContent := ""
afterContent := ""
if diffString == "" {
beforeContent = firstJSONText(success, "before_full_file_content", "beforeFullFileContent")
afterContent = firstJSONText(success, "after_full_file_content", "afterFullFileContent")
if beforeContent != "" {
diffString, _, _ = computeEditDiff(beforeContent, afterContent)
}
}
if diffString != "" {
diffString = truncateProjectedReplayText(firstNonEmpty(strings.TrimSpace(toolName), "PatchEdit"), diffString, patchEditReplayLimit(toolName))
encoded, err := json.Marshal(map[string]any{
"success": map[string]any{
"diff_string": diffString,
},
})
if err != nil {
return "", false
}
return string(encoded), true
}
if afterContent != "" {
afterContent = truncateProjectedReplayText(firstNonEmpty(strings.TrimSpace(toolName), "Write"), afterContent, patchEditReplayLimit(toolName))
encoded, err := json.Marshal(map[string]any{
"success": map[string]any{
"after_full_file_content": afterContent,
},
})
if err != nil {
return "", false
}
return string(encoded), true
}
return "", false
}
func compactHistoricalEditErrorReplay(toolName string, content string) (string, bool) {
switch strings.TrimSpace(toolName) {
case "PatchEdit", "PatchEditLines", "PatchEditSpan", "Edit", "Write":
default:
return "", false
}
trimmed := strings.TrimSpace(content)
if trimmed == "" {
return "", false
}
var payload map[string]any
if err := json.Unmarshal([]byte(trimmed), &payload); err != nil {
return "", false
}
errorPayload, ok := payload["error"].(map[string]any)
if !ok {
return "", false
}
if firstJSONText(errorPayload, "error", "modelVisibleError") == "" {
return "", false
}
encoded, err := json.Marshal(map[string]any{
"error": map[string]any{
"error": "historical edit error omitted from replay; re-read the file before editing",
},
})
if err != nil {
return "", false
}
return string(encoded), true
}
func firstJSONText(payload map[string]any, keys ...string) string {
for _, key := range keys {
if value, ok := payload[key].(string); ok {
return value
}
}
return ""
}
func truncateProjectedReplayText(toolName string, text string, limit int) string {
if limit <= 0 || len(text) <= limit {
return strings.TrimSpace(text)
}
original := len(text)
notice := fmt.Sprintf("\n\n[truncated: %s result exceeded %d bytes; showing %d of %d bytes]", toolName, limit, limit, original)
for {
keep := limit - len(notice)
if keep <= 0 {
return truncateProjectedUTF8(text, limit)
}
kept := truncateProjectedUTF8(text, keep)
nextNotice := fmt.Sprintf("\n\n[truncated: %s result exceeded %d bytes; showing %d of %d bytes]", toolName, limit, len(kept), original)
output := strings.TrimRight(kept, "\n") + nextNotice
if len(output) <= limit || nextNotice == notice {
return output
}
notice = nextNotice
}
}
func truncateProjectedReplayTextMiddle(toolName string, text string, limit int) string {
if limit <= 0 {
return ""
}
if len(text) <= limit {
return text
}
original := len(text)
notice := fmt.Sprintf("\n\n[truncated: %s result exceeded %d bytes; omitted middle; showing %d of %d bytes]\n\n", toolName, limit, limit, original)
for {
keep := limit - len(notice)
if keep <= 0 {
return truncateProjectedUTF8(text, limit)
}
headLimit := keep / 2
tailLimit := keep - headLimit
head := truncateProjectedUTF8(text, headLimit)
tail := truncateProjectedUTF8Suffix(text, tailLimit)
kept := len(head) + len(tail)
nextNotice := fmt.Sprintf("\n\n[truncated: %s result exceeded %d bytes; omitted middle; showing %d of %d bytes]\n\n", toolName, limit, kept, original)
output := head + nextNotice + tail
if len(output) <= limit || nextNotice == notice {
return output
}
notice = nextNotice
}
}
func truncateProjectedUTF8(text string, limit int) string {
if limit <= 0 {
return ""
}
if len(text) <= limit {
return text
}
if limit > len(text) {
limit = len(text)
}
truncated := text[:limit]
for !utf8.ValidString(truncated) && len(truncated) > 0 {
truncated = truncated[:len(truncated)-1]
}
return truncated
}
func truncateProjectedUTF8Suffix(text string, limit int) string {
if limit <= 0 {
return ""
}
if len(text) <= limit {
return text
}
start := len(text) - limit
if start < 0 {
start = 0
}
suffix := text[start:]
for !utf8.ValidString(suffix) && start < len(text) {
start++
suffix = text[start:]
}
return suffix
}
+510
View File
@@ -0,0 +1,510 @@
// types.go 定义 forwarder 的核心数据结构与最小接口边界。
package forwarder
import (
"context"
"encoding/json"
"fmt"
"strings"
"sync"
"sync/atomic"
"time"
"cursor/gen/agentv1"
runtimecore "cursor/internal/backend/agent/core"
modeladapter "cursor/internal/backend/agent/model"
)
type ConversationFile struct {
SchemaVersion int `json:"schema_version,omitempty"`
ConversationID string `json:"conversation_id"`
RootConversationID string `json:"root_conversation_id"`
ParentConversationID string `json:"parent_conversation_id"`
ParentToolCallID string `json:"parent_tool_call_id"`
SubagentTypeName string `json:"subagent_type_name,omitempty"`
Mode string `json:"mode"`
ContextVersion int64 `json:"context_version,omitempty"`
CurrentLoopID string `json:"current_loop_id,omitempty"`
CurrentLoopStatus string `json:"current_loop_status,omitempty"`
CurrentRequestID string `json:"current_request_id,omitempty"`
CurrentTurnSeq int64 `json:"current_turn_seq,omitempty"`
TokenDetailsUsedTokens uint32 `json:"token_details_used_tokens,omitempty"`
TokenDetailsMaxTokens uint32 `json:"token_details_max_tokens,omitempty"`
AutoCompactionPending bool `json:"auto_compaction_pending,omitempty"`
AutoCompactionPromptTokens int64 `json:"auto_compaction_prompt_tokens,omitempty"`
AutoCompactionReserveTokens int64 `json:"auto_compaction_reserve_tokens,omitempty"`
AutoCompactionTriggeredAt string `json:"auto_compaction_triggered_at,omitempty"`
AutoCompactionSourceModelCallID string `json:"auto_compaction_source_model_call_id,omitempty"`
CurrentPlanText string `json:"current_plan_text,omitempty"`
CurrentPlans map[string]*agentv1.PlanRegistryEntry `json:"current_plans,omitempty"`
CurrentTodos []*agentv1.TodoItem `json:"current_todos,omitempty"`
LatestRequestPrefix *ConversationRequestPrefix `json:"latest_request_prefix,omitempty"`
LastProviderCall *ConversationProviderCall `json:"last_provider_call,omitempty"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
NextTurnSeq int64 `json:"next_turn_seq"`
NextEntrySeq int64 `json:"next_entry_seq"`
Entries []HistoryEntry `json:"entries,omitempty"`
}
type ConversationRequestPrefix struct {
RequestID string `json:"request_id,omitempty"`
ModelCallID string `json:"model_call_id,omitempty"`
Provider string `json:"provider,omitempty"`
OpenAIEndpoint string `json:"openai_endpoint,omitempty"`
Model string `json:"model,omitempty"`
PromptTokensTotal int64 `json:"prompt_tokens_total,omitempty"`
ReplayMessageCount int `json:"replay_message_count,omitempty"`
CanonicalBodyHash string `json:"canonical_body_hash,omitempty"`
FrontierHash string `json:"frontier_hash,omitempty"`
FrontierPath string `json:"frontier_path,omitempty"`
BreakpointCount int `json:"breakpoint_count,omitempty"`
ExpectedCacheRead bool `json:"expected_cache_read,omitempty"`
PreviousFrontierMatched bool `json:"previous_frontier_matched,omitempty"`
UpdatedAt time.Time `json:"updated_at,omitempty"`
}
type ConversationProviderCall struct {
RequestID string `json:"request_id,omitempty"`
ModelCallID string `json:"model_call_id,omitempty"`
Provider string `json:"provider,omitempty"`
Model string `json:"model,omitempty"`
Status string `json:"status,omitempty"`
ErrorText string `json:"error_text,omitempty"`
UpdatedAt time.Time `json:"updated_at,omitempty"`
}
type HistoryEntry struct {
Seq int64 `json:"seq"`
TurnSeq int64 `json:"turn_seq"`
RequestID string `json:"request_id,omitempty"`
Role string `json:"role"`
Kind string `json:"kind"`
ToolCallID string `json:"tool_call_id,omitempty"`
ParentToolCallID string `json:"parent_tool_call_id,omitempty"`
Payload json.RawMessage `json:"payload"`
CreatedAt time.Time `json:"created_at"`
}
type ConversationSummary struct {
ConversationID string `json:"conversation_id"`
Mode string `json:"mode"`
EntriesCount int `json:"entries_count"`
NextTurnSeq int64 `json:"next_turn_seq"`
NextEntrySeq int64 `json:"next_entry_seq"`
UpdatedAt time.Time `json:"updated_at"`
}
type StreamStatus string
const (
StreamStatusCreated StreamStatus = "created"
StreamStatusStreaming StreamStatus = "streaming"
StreamStatusCompleted StreamStatus = "completed"
StreamStatusCanceled StreamStatus = "canceled"
StreamStatusFailed StreamStatus = "failed"
)
type StreamEvent struct {
Message *agentv1.AgentServerMessage
End bool
TerminalErrorCode string
TerminalErrorMessage string
}
type StreamSubscriber struct {
Signal chan struct{}
}
type ActiveStream struct {
mu sync.Mutex
RequestID string
ConversationID string
TurnSeq int64
ModelID string
ModelName string
Mode agentv1.AgentMode
LatestUserText string
Status StreamStatus
ThinkingEffort string
SubagentModelOverrides map[string]runtimecore.SubagentModelOverrideSelection
CurrentModelCallID string
ProviderActive bool
ProviderCancel func()
ProviderPassCount int
ToolInvocationCount int
ActorMailbox chan streamCommandEnvelope
ActorDone chan struct{}
Phase TurnPhase
PendingProviderAction providerAction
PendingProviderCompletion *pendingTurnCompletion
CurrentProviderToken uint64
CurrentCompactionToken uint64
TimerTokens map[string]uint64
ProviderAccumulatedText string
ProviderAccumulatedReasoning string
ProviderAccumulatedReasoningSignature string
ProviderAccumulatedReasoningSignatureSource string
ProviderAccumulatedReasoningItemID string
ProviderAccumulatedReasoningStatus string
ProviderAccumulatedReasoningSummary json.RawMessage
ProviderSyntheticThinkingStartedAt time.Time
ProviderSyntheticThinkingPublished bool
ProviderFinishReason string
ProviderUsage turnUsageSnapshot
ProviderTerminalToolInvocation bool
PendingCompaction *PendingCompaction
Backlog []StreamEvent
Subscribers map[string]*StreamSubscriber
CheckpointConversation *ConversationFile
PendingExecs map[string]runtimecore.PendingExec
PendingInteractions map[string]runtimecore.PendingInteraction
PartialToolCallIDs map[string]struct{}
PatchEditQueues map[string][]queuedPatchEditOperation
MCPToolServers map[string]string
WorkspacePaths []string
TerminalsFolder string
RequestFileContents map[string]string
RecentCompletedExecs map[uint32]time.Time
BackgroundShells map[string]*BackgroundShellState
BackgroundShellsByMessageID map[uint32]string
BackgroundShellsByExecID map[string]string
BackgroundShellActions map[string]time.Time
TerminalCleanupTimer *time.Timer
TerminalCleanupSeq atomic.Uint64
CreatedAt time.Time
UpdatedAt time.Time
}
type BackgroundShellState struct {
ShellID string
Command string
WorkingDirectory string
PID *uint32
OriginalToolCallID string
OriginalExecID string
OriginalMessageID uint32
ModelCallID string
ArgsJSON []byte
Status string
ExitCode *int32
StdoutBuffer string
StderrBuffer string
AwaitStdoutOffset int
AwaitStderrOffset int
CreatedAt time.Time
LastActivityAt time.Time
CompletedAt time.Time
StreamClosed bool
}
type pendingTurnCompletion struct {
ConversationID string
RequestID string
TurnSeq int64
ModelCallID string
ProviderPass int
Usage turnUsageSnapshot
Disposition pendingCompletionDisposition
}
type PendingCompaction struct {
Trigger string
ContextTokens int64
ContextWindowSize int64
ContextUsagePercent float64
ReserveTokens int64
MessageCount int32
MessagesToCompact int32
CompactTurnCount int32
IsFirstCompaction bool
ExistingSummary string
CompactedTurns []compactedTurnSummary
ManualInstruction string
RequestSource string
CurrentTurnSeq int64
CurrentRequestID string
CurrentUserText string
PreserveCurrentTurnInputs bool
HookMessage string
SummaryModelCallID string
StartedAt time.Time
}
type ProviderRequest struct {
RequestID string
ConversationID string
RunID string
ModelCallID string
ModelID string
Mode agentv1.AgentMode
ThinkingEffort string
Messages []modeladapter.Message
StableMessageCount int
Tools []json.RawMessage
MaxTokens int
RequestKnobs map[string]any
CompileSummary string
Observer modeladapter.LLMArtifactObserver
ArtifactPaths *modeladapter.LLMArtifactPaths
RequestBodyOverride map[string]any
}
type ProviderGateway interface {
StartStream(ctx context.Context, req ProviderRequest, sink func(modeladapter.ModelEvent) error) error
}
type ToolCatalog interface {
Load(mode agentv1.AgentMode, subagentTypeName string) ([]json.RawMessage, []string, error)
}
type PromptReminders struct {
SystemParts []string
TailMessages []modeladapter.Message
PromptContexts []PromptContextMessage
}
type ReminderInjector interface {
Inject(mode agentv1.AgentMode, conversation *ConversationFile, replayMessages []modeladapter.Message, latestUserText string, toolNames []string) PromptReminders
}
type PromptContextMessage struct {
Source string
Message modeladapter.Message
ContentHash string
Persist bool
}
type CompiledConversation struct {
Mode agentv1.AgentMode
Messages []modeladapter.Message
StableMessageCount int
Tools []json.RawMessage
CompileSummary string
}
type ToolRequestKind string
const (
ToolRequestExec ToolRequestKind = "exec"
ToolRequestInteraction ToolRequestKind = "interaction"
)
// providerTerminalError 表示底层 LLM/provider 返回的真实错误。
type providerTerminalError struct {
cause error
}
// Error 返回 provider 错误的字符串形式。
func (err providerTerminalError) Error() string {
if err.cause == nil {
return "provider error"
}
return err.cause.Error()
}
// Unwrap 允许调用方继续取到底层原始错误。
func (err providerTerminalError) Unwrap() error {
return err.cause
}
type toolResultEntryPayload struct {
ToolCallID string `json:"tool_call_id"`
ToolName string `json:"tool_name"`
Arguments string `json:"arguments,omitempty"`
ResultText string `json:"result_text,omitempty"`
ReasoningContent string `json:"reasoning_content,omitempty"`
ReasoningSignature string `json:"reasoning_signature,omitempty"`
ReasoningSignatureSource string `json:"reasoning_signature_source,omitempty"`
ReasoningItemID string `json:"reasoning_item_id,omitempty"`
ReasoningStatus string `json:"reasoning_status,omitempty"`
ReasoningSummary json.RawMessage `json:"reasoning_summary,omitempty"`
ProviderItemID string `json:"provider_item_id,omitempty"`
ProviderCallID string `json:"provider_call_id,omitempty"`
ProviderStatus string `json:"provider_status,omitempty"`
ToolCall json.RawMessage `json:"tool_call,omitempty"`
}
type toolCallEntryPayload struct {
ToolCallID string `json:"tool_call_id"`
ToolName string `json:"tool_name"`
ReasoningContent string `json:"reasoning_content,omitempty"`
ReasoningSignature string `json:"reasoning_signature,omitempty"`
ReasoningSignatureSource string `json:"reasoning_signature_source,omitempty"`
ReasoningItemID string `json:"reasoning_item_id,omitempty"`
ReasoningStatus string `json:"reasoning_status,omitempty"`
ReasoningSummary json.RawMessage `json:"reasoning_summary,omitempty"`
ProviderItemID string `json:"provider_item_id,omitempty"`
ProviderCallID string `json:"provider_call_id,omitempty"`
ProviderStatus string `json:"provider_status,omitempty"`
ToolCall json.RawMessage `json:"tool_call"`
}
type assistantTextPayload struct {
Text string `json:"text"`
ReasoningContent string `json:"reasoning_content,omitempty"`
ReasoningSignature string `json:"reasoning_signature,omitempty"`
ReasoningSignatureSource string `json:"reasoning_signature_source,omitempty"`
ReasoningItemID string `json:"reasoning_item_id,omitempty"`
ReasoningStatus string `json:"reasoning_status,omitempty"`
ReasoningSummary json.RawMessage `json:"reasoning_summary,omitempty"`
}
type metadataPayload struct {
Type string `json:"type"`
Value map[string]any `json:"value,omitempty"`
}
type promptContextEntryPayload struct {
Source string `json:"source"`
Role string `json:"role"`
Content string `json:"content"`
ContentHash string `json:"content_hash,omitempty"`
}
type modelMessageEntryPayload struct {
Message modeladapter.Message `json:"message"`
}
type runtimeStateEntryPayload struct {
PlanText string `json:"plan_text,omitempty"`
Plans map[string]*agentv1.PlanRegistryEntry `json:"plans,omitempty"`
Todos []*agentv1.TodoItem `json:"todos,omitempty"`
}
type compactionSummaryEntryPayload struct {
Summary string `json:"summary"`
Trigger string `json:"trigger,omitempty"`
CurrentTurnSeq int64 `json:"current_turn_seq,omitempty"`
CurrentRequestID string `json:"current_request_id,omitempty"`
CompactTurnCount int32 `json:"compact_turn_count,omitempty"`
MessagesToCompact int32 `json:"messages_to_compact,omitempty"`
PreserveCurrentTurnInputs bool `json:"preserve_current_turn_inputs,omitempty"`
}
type ModeSource string
const (
ModeSourceUnknown ModeSource = ""
ModeSourceUserMessage ModeSource = "user_message"
ModeSourceExecutePlanAction ModeSource = "execute_plan_action"
ModeSourceConversationState ModeSource = "conversation_state"
ModeSourceSwitchModeTool ModeSource = "switch_mode_tool"
)
type InboundIntent struct {
Kind string
RequestID string
ConversationID string
ModelID string
ModelName string
ThinkingEffort string
Mode agentv1.AgentMode
HasExplicitMode bool
ModeSource ModeSource
StartsRun bool
SubagentTypeName string
SubagentModelOverrides map[string]runtimecore.SubagentModelOverrideSelection
ConversationState *agentv1.ConversationStateStructure
UserMessage *agentv1.UserMessage
RequestContext *agentv1.RequestContext
ClientMessage *agentv1.AgentClientMessage
ExecClientMessage *agentv1.ExecClientMessage
ExecClientControlMessage *agentv1.ExecClientControlMessage
InteractionResponse *agentv1.InteractionResponse
KVClientMessage *agentv1.KvClientMessage
CancelReason string
IgnoredReason string
Prewarm bool
}
// normalizeMode 对外部传入的 mode 做最小归一化,但不再静默降级。
func normalizeMode(mode agentv1.AgentMode) agentv1.AgentMode {
return mode
}
func isSupportedActiveMode(mode agentv1.AgentMode) bool {
switch normalizeMode(mode) {
case agentv1.AgentMode_AGENT_MODE_AGENT,
agentv1.AgentMode_AGENT_MODE_ASK,
agentv1.AgentMode_AGENT_MODE_PLAN,
agentv1.AgentMode_AGENT_MODE_DEBUG,
agentv1.AgentMode_AGENT_MODE_MULTITASK:
return true
default:
return false
}
}
func validateSupportedActiveMode(mode agentv1.AgentMode) (agentv1.AgentMode, error) {
normalized := normalizeMode(mode)
if !isSupportedActiveMode(normalized) {
return agentv1.AgentMode_AGENT_MODE_UNSPECIFIED, fmt.Errorf("unsupported active mode: %s", normalized.String())
}
return normalized, nil
}
func resolveExplicitMode(mode agentv1.AgentMode, source ModeSource) (agentv1.AgentMode, ModeSource, bool, error) {
normalized, err := validateSupportedActiveMode(mode)
if err != nil {
return agentv1.AgentMode_AGENT_MODE_UNSPECIFIED, source, true, fmt.Errorf("%w (source=%s)", err, strings.TrimSpace(string(source)))
}
return normalized, source, true, nil
}
// modeAlias 把协议枚举转换为写入 JSON history 的简短模式名。
func modeAlias(mode agentv1.AgentMode) (string, error) {
switch normalizeMode(mode) {
case agentv1.AgentMode_AGENT_MODE_AGENT:
return "agent", nil
case agentv1.AgentMode_AGENT_MODE_ASK:
return "ask", nil
case agentv1.AgentMode_AGENT_MODE_PLAN:
return "plan", nil
case agentv1.AgentMode_AGENT_MODE_DEBUG:
return "debug", nil
case agentv1.AgentMode_AGENT_MODE_MULTITASK:
return "multitask", nil
default:
return "", fmt.Errorf("unsupported mode alias: %s", normalizeMode(mode).String())
}
}
// parseModeAlias 把写入 JSON history 的模式名恢复为协议枚举。
func parseModeAlias(raw string) (agentv1.AgentMode, error) {
switch strings.ToLower(strings.TrimSpace(raw)) {
case "agent":
return agentv1.AgentMode_AGENT_MODE_AGENT, nil
case "ask":
return agentv1.AgentMode_AGENT_MODE_ASK, nil
case "plan":
return agentv1.AgentMode_AGENT_MODE_PLAN, nil
case "debug":
return agentv1.AgentMode_AGENT_MODE_DEBUG, nil
case "multitask":
return agentv1.AgentMode_AGENT_MODE_MULTITASK, nil
default:
return agentv1.AgentMode_AGENT_MODE_UNSPECIFIED, fmt.Errorf("unsupported mode alias: %q", strings.TrimSpace(raw))
}
}
func parseTargetModeID(raw string) (agentv1.AgentMode, error) {
switch strings.ToLower(strings.TrimSpace(raw)) {
case "agent":
return agentv1.AgentMode_AGENT_MODE_AGENT, nil
case "ask":
return agentv1.AgentMode_AGENT_MODE_ASK, nil
case "plan":
return agentv1.AgentMode_AGENT_MODE_PLAN, nil
case "debug":
return agentv1.AgentMode_AGENT_MODE_DEBUG, nil
case "multitask":
return agentv1.AgentMode_AGENT_MODE_MULTITASK, nil
default:
return agentv1.AgentMode_AGENT_MODE_UNSPECIFIED, fmt.Errorf("unsupported target mode id: %q", strings.TrimSpace(raw))
}
}
@@ -0,0 +1,233 @@
package forwarder
import (
"context"
"fmt"
"net/http"
"strings"
"time"
"connectrpc.com/connect"
"cursor/gen/aiserverv1"
"cursor/internal/logger"
)
const (
UploadServiceUploadDocumentationProcedure = "/aiserver.v1.UploadService/UploadDocumentation"
UploadServiceGetDocProcedure = "/aiserver.v1.UploadService/GetDoc"
UploadServiceGetPagesProcedure = "/aiserver.v1.UploadService/GetPages"
UploadServiceUploadedStatusProcedure = "/aiserver.v1.UploadService/UploadedStatus"
)
func newUploadServiceHandler(service *Service) http.Handler {
mux := http.NewServeMux()
mux.Handle(
UploadServiceUploadDocumentationProcedure,
connect.NewUnaryHandler(UploadServiceUploadDocumentationProcedure, service.UploadDocumentation),
)
mux.Handle(
UploadServiceGetDocProcedure,
connect.NewUnaryHandler(UploadServiceGetDocProcedure, service.GetDoc),
)
mux.Handle(
UploadServiceGetPagesProcedure,
connect.NewUnaryHandler(UploadServiceGetPagesProcedure, service.GetPages),
)
mux.Handle(
UploadServiceUploadedStatusProcedure,
connect.NewUnaryHandler(UploadServiceUploadedStatusProcedure, service.UploadedStatus),
)
mux.Handle("/", http.NotFoundHandler())
return mux
}
func (service *Service) UploadDocumentation(_ context.Context, req *connect.Request[aiserverv1.UploadDocumentationRequest]) (*connect.Response[aiserverv1.UploadResponse], error) {
store, err := service.requireDocsIndexStore()
if err != nil {
return nil, connect.NewError(connect.CodeInternal, err)
}
identifier := strings.TrimSpace(req.Msg.GetDocIdentifier())
if identifier == "" {
identifier = stableDocsIdentifier("upload-documentation", time.Now().UTC().Format(time.RFC3339Nano))
}
record, err := store.Upsert(uploadedDocsIndexRecord(identifier))
if err != nil {
return nil, connect.NewError(connect.CodeInternal, err)
}
logger.Infof("UploadService UploadDocumentation completed doc_identifier=%s", record.Identifier)
return connect.NewResponse(&aiserverv1.UploadResponse{
Status: aiserverv1.UploadResponse_STATUS_SUCCESS,
Progress: 1,
UploadedPages: uploadedDocPages(record),
DocUuid: uploadDocUUID(record),
}), nil
}
func (service *Service) GetDoc(_ context.Context, req *connect.Request[aiserverv1.GetDocRequest]) (*connect.Response[aiserverv1.ProtoDoc], error) {
store, err := service.requireDocsIndexStore()
if err != nil {
return nil, connect.NewError(connect.CodeInternal, err)
}
identifier := strings.TrimSpace(req.Msg.GetDocIdentifier())
if identifier == "" {
return connect.NewResponse(uploadedDocNotFound(identifier)), nil
}
record, ok, err := store.Get(identifier)
if err != nil {
return nil, connect.NewError(connect.CodeInternal, err)
}
if !ok {
logger.Infof("UploadService GetDoc not_found doc_identifier=%s", identifier)
return connect.NewResponse(uploadedDocNotFound(identifier)), nil
}
logger.Infof("UploadService GetDoc completed doc_identifier=%s", record.Identifier)
return connect.NewResponse(uploadedDocsProtoDoc(record)), nil
}
func (service *Service) GetPages(_ context.Context, req *connect.Request[aiserverv1.GetPagesRequest]) (*connect.Response[aiserverv1.Pages], error) {
store, err := service.requireDocsIndexStore()
if err != nil {
return nil, connect.NewError(connect.CodeInternal, err)
}
identifier := strings.TrimSpace(req.Msg.GetDocIdentifier())
if identifier == "" {
return connect.NewResponse(&aiserverv1.Pages{}), nil
}
record, ok, err := store.Get(identifier)
if err != nil {
return nil, connect.NewError(connect.CodeInternal, err)
}
if !ok {
logger.Infof("UploadService GetPages empty doc_identifier=%s", identifier)
return connect.NewResponse(&aiserverv1.Pages{}), nil
}
pages, pageURLs := uploadedDocsPagesResponse(record)
logger.Infof("UploadService GetPages completed doc_identifier=%s pages=%d", record.Identifier, len(pages))
return connect.NewResponse(&aiserverv1.Pages{Pages: pages, PageUrls: pageURLs}), nil
}
func (service *Service) UploadedStatus(_ context.Context, req *connect.Request[aiserverv1.UploadedStatusRequest]) (*connect.Response[aiserverv1.UploadedStatus], error) {
store, err := service.requireDocsIndexStore()
if err != nil {
return nil, connect.NewError(connect.CodeInternal, err)
}
identifier := strings.TrimSpace(req.Msg.GetDocIdentifier())
if identifier == "" {
return connect.NewResponse(&aiserverv1.UploadedStatus{Status: aiserverv1.UploadedStatus_STATUS_NOT_FOUND}), nil
}
record, ok, err := store.Get(identifier)
if err != nil {
return nil, connect.NewError(connect.CodeInternal, err)
}
if !ok {
return connect.NewResponse(&aiserverv1.UploadedStatus{Status: aiserverv1.UploadedStatus_STATUS_NOT_FOUND}), nil
}
return connect.NewResponse(&aiserverv1.UploadedStatus{
Status: aiserverv1.UploadedStatus_STATUS_SUCCEEDED,
UploadedPages: uploadedDocPages(record),
}), nil
}
func uploadedDocsIndexRecord(identifier string) DocsIndexRecord {
url := docsIndexURLCandidate(identifier)
title := identifier
if url != "" {
title = docsTitleFromURL(url)
}
return DocsIndexRecord{
ID: identifier,
Identifier: identifier,
Title: title,
URL: url,
Status: docsIndexStatusIndexed,
Source: docsIndexSourceLocal,
}
}
func uploadedDocsProtoDoc(record DocsIndexRecord) *aiserverv1.ProtoDoc {
createdAt := docsIndexTimeString(record.CreatedAt)
updatedAt := docsIndexTimeString(record.UpdatedAt)
lastUploadedAt := docsIndexTimeString(record.LastIndexAt)
if lastUploadedAt == "" {
lastUploadedAt = updatedAt
}
return &aiserverv1.ProtoDoc{
Uuid: uploadDocUUID(record),
DocIdentifier: record.Identifier,
DocName: record.Title,
DocUrlRoot: record.URL,
DocUrlPrefix: record.URL,
CreatedAt: createdAt,
UpdatedAt: updatedAt,
LastUploadedAt: lastUploadedAt,
UploadStatus: &aiserverv1.UploadedStatus{
Status: aiserverv1.UploadedStatus_STATUS_SUCCEEDED,
UploadedPages: uploadedDocPages(record),
},
Pages: uploadedDocProtoPages(record),
}
}
func uploadedDocNotFound(identifier string) *aiserverv1.ProtoDoc {
return &aiserverv1.ProtoDoc{
DocIdentifier: strings.TrimSpace(identifier),
UploadStatus: &aiserverv1.UploadedStatus{Status: aiserverv1.UploadedStatus_STATUS_NOT_FOUND},
}
}
func uploadedDocProtoPages(record DocsIndexRecord) []*aiserverv1.ProtoDocPage {
pages, pageURLs := uploadedDocsPagesResponse(record)
items := make([]*aiserverv1.ProtoDocPage, 0, len(pages))
for i, page := range pages {
url := ""
if i < len(pageURLs) {
url = pageURLs[i]
}
items = append(items, &aiserverv1.ProtoDocPage{Title: page, Url: url})
}
return items
}
func uploadedDocsPagesResponse(record DocsIndexRecord) ([]string, []string) {
title := strings.TrimSpace(record.Title)
url := strings.TrimSpace(record.URL)
if title == "" || url == "" {
return nil, nil
}
return []string{title}, []string{url}
}
func uploadedDocPages(record DocsIndexRecord) []string {
pages, _ := uploadedDocsPagesResponse(record)
return pages
}
func uploadDocUUID(record DocsIndexRecord) string {
identifier := strings.TrimSpace(record.Identifier)
if identifier == "" {
identifier = strings.TrimSpace(record.ID)
}
if identifier == "" {
identifier = fmt.Sprintf("doc-%d", time.Now().UnixNano())
}
return stableDocsIdentifier("upload-doc-uuid", identifier)
}
func docsIndexTimeString(value time.Time) string {
if value.IsZero() {
return ""
}
return value.UTC().Format(time.RFC3339)
}
func docsTitleFromURL(value string) string {
trimmed := strings.TrimSpace(value)
trimmed = strings.TrimPrefix(trimmed, "https://")
trimmed = strings.TrimPrefix(trimmed, "http://")
trimmed = strings.Trim(trimmed, "/")
if trimmed == "" {
return value
}
return trimmed
}
+343
View File
@@ -0,0 +1,343 @@
package forwarder
import (
"encoding/json"
"errors"
"fmt"
"os"
"path/filepath"
"strings"
"time"
)
const (
usageFileName = "usage.json"
usageFileSchemaVersion = 2
usageRecentEventLimit = 500
usageEventKindProvider = "provider_call"
usageEventKindTurn = "turn_finalized"
usageTurnStatusDone = "completed"
)
type UsageFileStore struct {
path string
}
type usageFileDocument struct {
SchemaVersion int `json:"schema_version"`
UpdatedAt time.Time `json:"updated_at"`
Totals usageFileTotals `json:"totals"`
Daily []usageFileDaily `json:"daily"`
RecentEvents []usageFileEvent `json:"recent_events"`
EventIndex map[string]usageFileEvent `json:"event_index,omitempty"`
}
type usageFileTotals struct {
ProviderCalls int64 `json:"provider_calls"`
TurnsTotal int64 `json:"turns_total"`
ValidTurnsTotal int64 `json:"valid_turns_total"`
InvalidTurnsTotal int64 `json:"invalid_turns_total"`
InputTokens int64 `json:"input_tokens"`
OutputTokens int64 `json:"output_tokens"`
CacheReadTokens int64 `json:"cache_read_tokens"`
CacheWriteTokens int64 `json:"cache_write_tokens"`
TotalTokens int64 `json:"total_tokens"`
}
type usageFileDaily struct {
Date string `json:"date"`
ProviderCalls int64 `json:"provider_calls"`
TurnsTotal int64 `json:"turns_total"`
ValidTurnsTotal int64 `json:"valid_turns_total"`
InvalidTurnsTotal int64 `json:"invalid_turns_total"`
InputTokens int64 `json:"input_tokens"`
OutputTokens int64 `json:"output_tokens"`
CacheReadTokens int64 `json:"cache_read_tokens"`
CacheWriteTokens int64 `json:"cache_write_tokens"`
TotalTokens int64 `json:"total_tokens"`
}
type usageFileEvent struct {
EventID string `json:"event_id"`
Kind string `json:"kind,omitempty"`
Status string `json:"status,omitempty"`
At time.Time `json:"at"`
InputTokens int64 `json:"input_tokens"`
OutputTokens int64 `json:"output_tokens"`
CacheReadTokens int64 `json:"cache_read_tokens"`
CacheWriteTokens int64 `json:"cache_write_tokens"`
TotalTokens int64 `json:"total_tokens"`
UsagePresent bool `json:"usage_present"`
}
type usageFileDelta struct {
providerCalls int64
turnsTotal int64
validTurnsTotal int64
invalidTurnsTotal int64
inputTokens int64
outputTokens int64
cacheReadTokens int64
cacheWriteTokens int64
totalTokens int64
}
func NewUsageFileStore(historyRoot string) *UsageFileStore {
return &UsageFileStore{path: filepath.Join(strings.TrimSpace(historyRoot), usageFileName)}
}
func (store *UsageFileStore) UpsertEvent(event usageFileEvent) error {
if store == nil || strings.TrimSpace(store.path) == "" {
return nil
}
event.EventID = strings.TrimSpace(event.EventID)
if event.EventID == "" {
return nil
}
event.Kind = normalizeUsageEventKind(event.Kind)
event.Status = strings.TrimSpace(event.Status)
if event.At.IsZero() {
event.At = time.Now().UTC()
} else {
event.At = event.At.UTC()
}
event.InputTokens = nonNegativeInt64(event.InputTokens)
event.OutputTokens = nonNegativeInt64(event.OutputTokens)
event.CacheReadTokens = nonNegativeInt64(event.CacheReadTokens)
event.CacheWriteTokens = nonNegativeInt64(event.CacheWriteTokens)
event.TotalTokens = event.InputTokens + event.OutputTokens + event.CacheReadTokens + event.CacheWriteTokens
if err := os.MkdirAll(filepath.Dir(store.path), 0o755); err != nil {
return fmt.Errorf("create usage directory: %w", err)
}
release, err := acquireConversationLock(store.path + ".lock")
if err != nil {
return err
}
defer release()
doc, err := readUsageFileDocument(store.path)
if err != nil {
return err
}
if doc.EventIndex == nil {
doc.EventIndex = make(map[string]usageFileEvent)
}
oldEvent, found := doc.EventIndex[event.EventID]
if found {
applyUsageFileDelta(&doc, oldEvent.At, negateUsageFileDelta(usageFileEventDelta(oldEvent)))
}
applyUsageFileDelta(&doc, event.At, usageFileEventDelta(event))
doc.RecentEvents = upsertRecentUsageEvent(doc.RecentEvents, event)
doc.RecentEvents = trimRecentUsageEvents(doc.RecentEvents, usageRecentEventLimit)
doc.EventIndex = buildUsageEventIndex(doc.RecentEvents)
doc.SchemaVersion = usageFileSchemaVersion
doc.UpdatedAt = time.Now().UTC()
return writeJSONFileAtomic(store.path, doc)
}
func (store *UsageFileStore) LookupEvent(needle string) (usageFileEvent, bool, error) {
if store == nil || strings.TrimSpace(store.path) == "" {
return usageFileEvent{}, false, nil
}
doc, err := readUsageFileDocument(store.path)
if err != nil {
return usageFileEvent{}, false, err
}
trimmed := strings.TrimSpace(needle)
if trimmed == "" {
return usageFileEvent{}, false, nil
}
var aggregate usageFileEvent
found := false
events := doc.EventIndex
if len(events) == 0 {
events = make(map[string]usageFileEvent, len(doc.RecentEvents))
for _, event := range doc.RecentEvents {
if eventID := strings.TrimSpace(event.EventID); eventID != "" {
events[eventID] = event
}
}
}
for _, event := range events {
eventID := strings.TrimSpace(event.EventID)
if eventID != trimmed && !strings.HasPrefix(eventID, trimmed+"::") {
continue
}
if !found {
aggregate = usageFileEvent{EventID: trimmed, At: event.At}
found = true
}
if event.At.After(aggregate.At) {
aggregate.At = event.At
}
aggregate.InputTokens += nonNegativeInt64(event.InputTokens)
aggregate.OutputTokens += nonNegativeInt64(event.OutputTokens)
aggregate.CacheReadTokens += nonNegativeInt64(event.CacheReadTokens)
aggregate.CacheWriteTokens += nonNegativeInt64(event.CacheWriteTokens)
aggregate.TotalTokens += nonNegativeInt64(event.TotalTokens)
aggregate.UsagePresent = aggregate.UsagePresent || event.UsagePresent
}
if found {
return aggregate, true, nil
}
return usageFileEvent{}, false, nil
}
func readUsageFileDocument(path string) (usageFileDocument, error) {
body, err := os.ReadFile(path)
if err != nil {
if errors.Is(err, os.ErrNotExist) {
return usageFileDocument{
SchemaVersion: usageFileSchemaVersion,
Daily: make([]usageFileDaily, 0),
RecentEvents: make([]usageFileEvent, 0),
EventIndex: make(map[string]usageFileEvent),
}, nil
}
return usageFileDocument{}, fmt.Errorf("read usage file: %w", err)
}
var doc usageFileDocument
if err := json.Unmarshal(body, &doc); err != nil {
return usageFileDocument{}, fmt.Errorf("decode usage file: %w", err)
}
if doc.SchemaVersion == 0 {
doc.SchemaVersion = 1
}
doc.RecentEvents = trimRecentUsageEvents(doc.RecentEvents, usageRecentEventLimit)
if len(doc.EventIndex) == 0 {
doc.EventIndex = buildUsageEventIndex(doc.RecentEvents)
}
return doc, nil
}
func upsertRecentUsageEvent(items []usageFileEvent, event usageFileEvent) []usageFileEvent {
event.EventID = strings.TrimSpace(event.EventID)
if event.EventID == "" {
return items
}
next := make([]usageFileEvent, 0, len(items)+1)
next = append(next, event)
for _, item := range items {
if strings.TrimSpace(item.EventID) == event.EventID {
continue
}
next = append(next, item)
}
return next
}
func trimRecentUsageEvents(items []usageFileEvent, limit int) []usageFileEvent {
if limit <= 0 || len(items) <= limit {
return items
}
return items[:limit]
}
func buildUsageEventIndex(items []usageFileEvent) map[string]usageFileEvent {
index := make(map[string]usageFileEvent, len(items))
for _, event := range items {
event.EventID = strings.TrimSpace(event.EventID)
if event.EventID == "" {
continue
}
event.Kind = normalizeUsageEventKind(event.Kind)
index[event.EventID] = event
}
return index
}
func normalizeUsageEventKind(kind string) string {
switch strings.TrimSpace(kind) {
case usageEventKindTurn:
return usageEventKindTurn
default:
return usageEventKindProvider
}
}
func usageFileEventDelta(event usageFileEvent) usageFileDelta {
switch normalizeUsageEventKind(event.Kind) {
case usageEventKindTurn:
delta := usageFileDelta{turnsTotal: 1}
if strings.TrimSpace(event.Status) == usageTurnStatusDone {
delta.validTurnsTotal = 1
} else {
delta.invalidTurnsTotal = 1
}
return delta
default:
return usageFileDelta{
providerCalls: 1,
inputTokens: nonNegativeInt64(event.InputTokens),
outputTokens: nonNegativeInt64(event.OutputTokens),
cacheReadTokens: nonNegativeInt64(event.CacheReadTokens),
cacheWriteTokens: nonNegativeInt64(event.CacheWriteTokens),
totalTokens: nonNegativeInt64(event.TotalTokens),
}
}
}
func negateUsageFileDelta(value usageFileDelta) usageFileDelta {
return usageFileDelta{
providerCalls: -value.providerCalls,
turnsTotal: -value.turnsTotal,
validTurnsTotal: -value.validTurnsTotal,
invalidTurnsTotal: -value.invalidTurnsTotal,
inputTokens: -value.inputTokens,
outputTokens: -value.outputTokens,
cacheReadTokens: -value.cacheReadTokens,
cacheWriteTokens: -value.cacheWriteTokens,
totalTokens: -value.totalTokens,
}
}
func applyUsageFileDelta(doc *usageFileDocument, at time.Time, delta usageFileDelta) {
if doc == nil {
return
}
doc.Totals.ProviderCalls = clampNonNegativeInt64(doc.Totals.ProviderCalls + delta.providerCalls)
doc.Totals.TurnsTotal = clampNonNegativeInt64(doc.Totals.TurnsTotal + delta.turnsTotal)
doc.Totals.ValidTurnsTotal = clampNonNegativeInt64(doc.Totals.ValidTurnsTotal + delta.validTurnsTotal)
doc.Totals.InvalidTurnsTotal = clampNonNegativeInt64(doc.Totals.InvalidTurnsTotal + delta.invalidTurnsTotal)
doc.Totals.InputTokens = clampNonNegativeInt64(doc.Totals.InputTokens + delta.inputTokens)
doc.Totals.OutputTokens = clampNonNegativeInt64(doc.Totals.OutputTokens + delta.outputTokens)
doc.Totals.CacheReadTokens = clampNonNegativeInt64(doc.Totals.CacheReadTokens + delta.cacheReadTokens)
doc.Totals.CacheWriteTokens = clampNonNegativeInt64(doc.Totals.CacheWriteTokens + delta.cacheWriteTokens)
doc.Totals.TotalTokens = clampNonNegativeInt64(doc.Totals.TotalTokens + delta.totalTokens)
date := at.UTC().Format("2006-01-02")
for index := range doc.Daily {
if doc.Daily[index].Date != date {
continue
}
applyUsageDailyDelta(&doc.Daily[index], delta)
return
}
item := usageFileDaily{Date: date}
applyUsageDailyDelta(&item, delta)
doc.Daily = append(doc.Daily, item)
}
func applyUsageDailyDelta(item *usageFileDaily, delta usageFileDelta) {
if item == nil {
return
}
item.ProviderCalls = clampNonNegativeInt64(item.ProviderCalls + delta.providerCalls)
item.TurnsTotal = clampNonNegativeInt64(item.TurnsTotal + delta.turnsTotal)
item.ValidTurnsTotal = clampNonNegativeInt64(item.ValidTurnsTotal + delta.validTurnsTotal)
item.InvalidTurnsTotal = clampNonNegativeInt64(item.InvalidTurnsTotal + delta.invalidTurnsTotal)
item.InputTokens = clampNonNegativeInt64(item.InputTokens + delta.inputTokens)
item.OutputTokens = clampNonNegativeInt64(item.OutputTokens + delta.outputTokens)
item.CacheReadTokens = clampNonNegativeInt64(item.CacheReadTokens + delta.cacheReadTokens)
item.CacheWriteTokens = clampNonNegativeInt64(item.CacheWriteTokens + delta.cacheWriteTokens)
item.TotalTokens = clampNonNegativeInt64(item.TotalTokens + delta.totalTokens)
}
func clampNonNegativeInt64(value int64) int64 {
if value < 0 {
return 0
}
return value
}
+607
View File
@@ -0,0 +1,607 @@
package forwarder
import (
"context"
"crypto/sha256"
"encoding/hex"
"errors"
"fmt"
"os"
"path/filepath"
"sort"
"strings"
"sync"
"time"
"connectrpc.com/connect"
"github.com/google/uuid"
"cursor/gen/aiserverv1"
"cursor/internal/logger"
)
const sharedUserRuleExtension = ".md"
type UserRuleRecord struct {
ID string
Title string
Filename string
FullPath string
Knowledge string
CreatedAt string
IsGenerated bool
ModifiedAt time.Time
ContentHash string
}
type userRuleGroup struct {
ContentHash string
Canonical UserRuleRecord
Members []UserRuleRecord
}
type UserRuleStore struct {
root string
mu sync.Mutex
}
func NewUserRuleStore(root string) *UserRuleStore {
return &UserRuleStore{root: strings.TrimSpace(root)}
}
func (store *UserRuleStore) List() ([]UserRuleRecord, error) {
if store == nil {
return nil, nil
}
store.mu.Lock()
defer store.mu.Unlock()
groups, err := store.listGroupsLocked()
if err != nil {
return nil, err
}
records := canonicalRecordsFromGroups(groups)
sort.SliceStable(records, func(i int, j int) bool {
if records[i].ModifiedAt.Equal(records[j].ModifiedAt) {
return records[i].Filename < records[j].Filename
}
return records[i].ModifiedAt.After(records[j].ModifiedAt)
})
return records, nil
}
func (store *UserRuleStore) Add(knowledge string) (UserRuleRecord, error) {
if store == nil {
return UserRuleRecord{}, fmt.Errorf("user rule store is nil")
}
if strings.TrimSpace(knowledge) == "" {
return UserRuleRecord{}, fmt.Errorf("knowledge is required")
}
store.mu.Lock()
defer store.mu.Unlock()
groups, err := store.listGroupsLocked()
if err != nil {
return UserRuleRecord{}, err
}
contentHash := hashSharedUserRuleContent([]byte(knowledge))
if groupIndex := indexGroupByHash(groups, contentHash); groupIndex >= 0 {
if err := store.removeRuleFilesLocked(nonCanonicalMembers(groups[groupIndex].Members)); err != nil {
return UserRuleRecord{}, err
}
return groups[groupIndex].Canonical, nil
}
id := uuid.NewString()
if err := store.writeRuleLocked(id, knowledge); err != nil {
return UserRuleRecord{}, err
}
return store.loadRuleByIDLocked(id)
}
func (store *UserRuleStore) Update(id string, knowledge string) (UserRuleRecord, bool, error) {
if store == nil {
return UserRuleRecord{}, false, fmt.Errorf("user rule store is nil")
}
if strings.TrimSpace(knowledge) == "" {
return UserRuleRecord{}, false, fmt.Errorf("knowledge is required")
}
normalizedID, err := normalizeUserRuleID(id)
if err != nil {
return UserRuleRecord{}, false, err
}
store.mu.Lock()
defer store.mu.Unlock()
current, err := store.loadRuleByIDLocked(normalizedID)
if err != nil {
if errors.Is(err, os.ErrNotExist) {
return UserRuleRecord{}, false, nil
}
return UserRuleRecord{}, false, err
}
groups, err := store.listGroupsLocked()
if err != nil {
return UserRuleRecord{}, false, err
}
currentGroupIndex := indexGroupByHash(groups, current.ContentHash)
targetHash := hashSharedUserRuleContent([]byte(knowledge))
targetGroupIndex := indexGroupByHash(groups, targetHash)
if targetGroupIndex >= 0 && groups[targetGroupIndex].ContentHash == current.ContentHash {
if currentGroupIndex >= 0 {
if err := store.removeRuleFilesLocked(nonCanonicalMembers(groups[currentGroupIndex].Members)); err != nil {
return UserRuleRecord{}, false, err
}
return groups[currentGroupIndex].Canonical, true, nil
}
return current, true, nil
}
if targetGroupIndex >= 0 {
if currentGroupIndex >= 0 {
remainingCurrentMembers := groupMembersExcludingIDs(groups[currentGroupIndex], map[string]struct{}{current.ID: {}})
if err := store.removeRuleFilesLocked(nonCanonicalMembers(remainingCurrentMembers)); err != nil {
return UserRuleRecord{}, false, err
}
}
if err := store.removeRuleFilesLocked(nonCanonicalMembers(groups[targetGroupIndex].Members)); err != nil {
return UserRuleRecord{}, false, err
}
if err := store.removeRuleFilesLocked([]UserRuleRecord{current}); err != nil {
return UserRuleRecord{}, false, err
}
return groups[targetGroupIndex].Canonical, true, nil
}
if err := store.writeRuleLocked(current.ID, knowledge); err != nil {
return UserRuleRecord{}, false, err
}
if currentGroupIndex >= 0 {
remainingCurrentMembers := groupMembersExcludingIDs(groups[currentGroupIndex], map[string]struct{}{current.ID: {}})
if err := store.removeRuleFilesLocked(nonCanonicalMembers(remainingCurrentMembers)); err != nil {
return UserRuleRecord{}, false, err
}
}
record, err := store.loadRuleByIDLocked(current.ID)
if err != nil {
return UserRuleRecord{}, false, err
}
return record, true, nil
}
func (store *UserRuleStore) Remove(id string) error {
if store == nil {
return nil
}
normalizedID, err := normalizeUserRuleID(id)
if err != nil {
return err
}
store.mu.Lock()
defer store.mu.Unlock()
record, err := store.loadRuleByIDLocked(normalizedID)
if err != nil {
if errors.Is(err, os.ErrNotExist) {
return nil
}
return err
}
groups, err := store.listGroupsLocked()
if err != nil {
return err
}
groupIndex := indexGroupByHash(groups, record.ContentHash)
if groupIndex < 0 {
return store.removeRuleFilesLocked([]UserRuleRecord{record})
}
return store.removeRuleFilesLocked(groups[groupIndex].Members)
}
func (store *UserRuleStore) BuildSystemPromptSection() (string, int, int, error) {
if store == nil {
return "", 0, 0, nil
}
store.mu.Lock()
defer store.mu.Unlock()
groups, err := store.listGroupsLocked()
if err != nil {
return "", 0, 0, err
}
if len(groups) == 0 {
return "", 0, 0, nil
}
totalFiles := totalRuleFileCount(groups)
records := canonicalRecordsFromGroups(groups)
sort.SliceStable(records, func(i int, j int) bool {
return records[i].Filename < records[j].Filename
})
lines := []string{
`<shared_user_rules description="These shared local rules are loaded from the backend configuration directory and apply to every local conversation. Follow them when relevant.">`,
}
visibleCount := 0
for _, record := range records {
if strings.TrimSpace(record.Knowledge) == "" {
continue
}
lines = append(lines,
fmt.Sprintf(`<rule file="%s">`, escapeSharedRulePromptText(record.Filename)),
escapeSharedRulePromptText(record.Knowledge),
"</rule>",
)
visibleCount++
}
if visibleCount == 0 {
return "", totalFiles, 0, nil
}
lines = append(lines, "</shared_user_rules>")
return strings.Join(lines, "\n"), totalFiles, visibleCount, nil
}
func (store *UserRuleStore) listGroupsLocked() ([]userRuleGroup, error) {
records, err := store.scanRuleFilesLocked()
if err != nil {
return nil, err
}
return buildUserRuleGroups(records), nil
}
func (store *UserRuleStore) scanRuleFilesLocked() ([]UserRuleRecord, error) {
if err := store.ensureRootLocked(); err != nil {
return nil, err
}
entries, err := os.ReadDir(store.root)
if err != nil {
if errors.Is(err, os.ErrNotExist) {
return nil, nil
}
return nil, fmt.Errorf("read shared user rules directory: %w", err)
}
records := make([]UserRuleRecord, 0, len(entries))
for _, entry := range entries {
if entry.IsDir() || filepath.Ext(entry.Name()) != sharedUserRuleExtension {
continue
}
record, err := store.readRuleFileLocked(filepath.Join(store.root, entry.Name()))
if err != nil {
logger.Infof("跳过不可用的共享 user rule 文件 path=%s err=%v", filepath.Join(store.root, entry.Name()), err)
continue
}
records = append(records, record)
}
return records, nil
}
func (store *UserRuleStore) loadRuleByIDLocked(id string) (UserRuleRecord, error) {
return store.readRuleFileLocked(store.rulePathLocked(id))
}
func (store *UserRuleStore) readRuleFileLocked(path string) (UserRuleRecord, error) {
data, err := os.ReadFile(path)
if err != nil {
return UserRuleRecord{}, err
}
if strings.TrimSpace(string(data)) == "" {
return UserRuleRecord{}, fmt.Errorf("shared user rule file is empty")
}
info, err := os.Stat(path)
if err != nil {
return UserRuleRecord{}, err
}
filename := filepath.Base(path)
id, err := normalizeUserRuleID(strings.TrimSuffix(filename, filepath.Ext(filename)))
if err != nil {
return UserRuleRecord{}, err
}
modifiedAt := info.ModTime().UTC()
return UserRuleRecord{
ID: id,
Title: id,
Filename: filename,
FullPath: path,
Knowledge: string(data),
CreatedAt: modifiedAt.Format(time.RFC3339Nano),
IsGenerated: false,
ModifiedAt: modifiedAt,
ContentHash: hashSharedUserRuleContent(data),
}, nil
}
func (store *UserRuleStore) writeRuleLocked(id string, knowledge string) error {
normalizedID, err := normalizeUserRuleID(id)
if err != nil {
return err
}
if strings.TrimSpace(knowledge) == "" {
return fmt.Errorf("knowledge is required")
}
if err := store.ensureRootLocked(); err != nil {
return err
}
path := store.rulePathLocked(normalizedID)
if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil {
return fmt.Errorf("create shared user rules directory: %w", err)
}
tempPath := path + ".tmp"
if err := os.WriteFile(tempPath, []byte(knowledge), 0o644); err != nil {
return fmt.Errorf("write temp shared user rule: %w", err)
}
if err := os.Rename(tempPath, path); err != nil {
return fmt.Errorf("rename shared user rule file: %w", err)
}
return nil
}
func (store *UserRuleStore) removeRuleFilesLocked(records []UserRuleRecord) error {
for _, record := range records {
path := strings.TrimSpace(record.FullPath)
if path == "" {
path = store.rulePathLocked(record.ID)
}
if strings.TrimSpace(path) == "" {
continue
}
if err := os.Remove(path); err != nil && !errors.Is(err, os.ErrNotExist) {
return fmt.Errorf("remove shared user rule %q: %w", record.ID, err)
}
}
return nil
}
func (store *UserRuleStore) ensureRootLocked() error {
if strings.TrimSpace(store.root) == "" {
return fmt.Errorf("user rules root is empty")
}
if err := os.MkdirAll(store.root, 0o755); err != nil {
return fmt.Errorf("create user rules root: %w", err)
}
return nil
}
func (store *UserRuleStore) rulePathLocked(id string) string {
if store == nil || strings.TrimSpace(store.root) == "" {
return ""
}
return filepath.Join(store.root, id+sharedUserRuleExtension)
}
func buildUserRuleGroups(records []UserRuleRecord) []userRuleGroup {
if len(records) == 0 {
return nil
}
membersByHash := make(map[string][]UserRuleRecord)
for _, record := range records {
membersByHash[record.ContentHash] = append(membersByHash[record.ContentHash], record)
}
groups := make([]userRuleGroup, 0, len(membersByHash))
for contentHash, members := range membersByHash {
sort.SliceStable(members, func(i int, j int) bool {
return members[i].Filename < members[j].Filename
})
groups = append(groups, userRuleGroup{
ContentHash: contentHash,
Canonical: members[0],
Members: members,
})
}
sort.SliceStable(groups, func(i int, j int) bool {
return groups[i].Canonical.Filename < groups[j].Canonical.Filename
})
return groups
}
func canonicalRecordsFromGroups(groups []userRuleGroup) []UserRuleRecord {
if len(groups) == 0 {
return nil
}
records := make([]UserRuleRecord, 0, len(groups))
for _, group := range groups {
records = append(records, group.Canonical)
}
return records
}
func totalRuleFileCount(groups []userRuleGroup) int {
total := 0
for _, group := range groups {
total += len(group.Members)
}
return total
}
func indexGroupByHash(groups []userRuleGroup, contentHash string) int {
for index := range groups {
if groups[index].ContentHash == contentHash {
return index
}
}
return -1
}
func nonCanonicalMembers(records []UserRuleRecord) []UserRuleRecord {
if len(records) <= 1 {
return nil
}
output := make([]UserRuleRecord, 0, len(records)-1)
output = append(output, records[1:]...)
return output
}
func groupMembersExcludingIDs(group userRuleGroup, excludedIDs map[string]struct{}) []UserRuleRecord {
if len(group.Members) == 0 {
return nil
}
filtered := make([]UserRuleRecord, 0, len(group.Members))
for _, record := range group.Members {
if _, excluded := excludedIDs[record.ID]; excluded {
continue
}
filtered = append(filtered, record)
}
return filtered
}
func normalizeUserRuleID(raw string) (string, error) {
id := strings.TrimSpace(raw)
id = strings.TrimSuffix(id, sharedUserRuleExtension)
switch {
case id == "":
return "", fmt.Errorf("user rule id is required")
case strings.Contains(id, "/"), strings.Contains(id, "\\"), strings.Contains(id, ".."):
return "", fmt.Errorf("invalid user rule id %q", raw)
default:
return id, nil
}
}
func hashSharedUserRuleContent(data []byte) string {
sum := sha256.Sum256(data)
return hex.EncodeToString(sum[:])
}
func escapeSharedRulePromptText(value string) string {
replacer := strings.NewReplacer(
"&", "&amp;",
`"`, "&quot;",
"<", "&lt;",
">", "&gt;",
)
return replacer.Replace(value)
}
func (service *Service) KnowledgeBaseAdd(_ context.Context, req *connect.Request[aiserverv1.KnowledgeBaseAddRequest]) (*connect.Response[aiserverv1.KnowledgeBaseAddResponse], error) {
if service == nil || service.rules == nil {
return nil, connect.NewError(connect.CodeInternal, fmt.Errorf("user rule store is not initialized"))
}
if strings.TrimSpace(req.Msg.GetKnowledge()) == "" {
return nil, connect.NewError(connect.CodeInvalidArgument, fmt.Errorf("knowledge is required"))
}
record, err := service.rules.Add(req.Msg.GetKnowledge())
if err != nil {
return nil, connect.NewError(connect.CodeInternal, err)
}
if service.docsIndexStore != nil {
if _, err := service.docsIndexStore.Upsert(DocsIndexRecord{
ID: record.ID,
Identifier: record.ID,
Title: firstNonEmptyDocs(req.Msg.GetTitle(), record.Title, record.ID),
URL: firstNonEmptyDocs(docsIndexURLCandidate(req.Msg.GetKnowledge()), docsIndexURLCandidate(req.Msg.GetTitle())),
Content: req.Msg.GetKnowledge(),
GitOrigin: req.Msg.GetGitOrigin(),
Source: docsIndexSourceLocal,
}); err != nil {
logger.Errorf("docs index sync failed operation=add id=%s err=%v", record.ID, err)
}
}
return connect.NewResponse(&aiserverv1.KnowledgeBaseAddResponse{
Success: true,
Id: record.ID,
}), nil
}
func (service *Service) KnowledgeBaseList(_ context.Context, req *connect.Request[aiserverv1.KnowledgeBaseListRequest]) (*connect.Response[aiserverv1.KnowledgeBaseListResponse], error) {
if service == nil || service.rules == nil {
return nil, connect.NewError(connect.CodeInternal, fmt.Errorf("user rule store is not initialized"))
}
records, err := service.rules.List()
if err != nil {
return nil, connect.NewError(connect.CodeInternal, err)
}
if limit := req.Msg.GetLimit(); limit > 0 && len(records) > int(limit) {
records = records[:limit]
}
items := make([]*aiserverv1.KnowledgeBaseListResponse_Item, 0, len(records))
for _, record := range records {
items = append(items, &aiserverv1.KnowledgeBaseListResponse_Item{
Id: record.ID,
Knowledge: record.Knowledge,
Title: record.Title,
CreatedAt: record.CreatedAt,
IsGenerated: false,
})
}
return connect.NewResponse(&aiserverv1.KnowledgeBaseListResponse{
Success: true,
AllResults: items,
}), nil
}
func (service *Service) KnowledgeBaseUpdate(_ context.Context, req *connect.Request[aiserverv1.KnowledgeBaseUpdateRequest]) (*connect.Response[aiserverv1.KnowledgeBaseUpdateResponse], error) {
if service == nil || service.rules == nil {
return nil, connect.NewError(connect.CodeInternal, fmt.Errorf("user rule store is not initialized"))
}
if _, err := normalizeUserRuleID(req.Msg.GetId()); err != nil {
return nil, connect.NewError(connect.CodeInvalidArgument, err)
}
if strings.TrimSpace(req.Msg.GetKnowledge()) == "" {
return nil, connect.NewError(connect.CodeInvalidArgument, fmt.Errorf("knowledge is required"))
}
_, exists, err := service.rules.Update(req.Msg.GetId(), req.Msg.GetKnowledge())
if err != nil {
return nil, connect.NewError(connect.CodeInternal, err)
}
if !exists {
return nil, connect.NewError(connect.CodeNotFound, fmt.Errorf("user rule %q not found", strings.TrimSpace(req.Msg.GetId())))
}
if service.docsIndexStore != nil {
if _, err := service.docsIndexStore.Upsert(DocsIndexRecord{
ID: strings.TrimSpace(req.Msg.GetId()),
Identifier: strings.TrimSpace(req.Msg.GetId()),
Title: firstNonEmptyDocs(req.Msg.GetTitle(), strings.TrimSpace(req.Msg.GetId())),
URL: firstNonEmptyDocs(docsIndexURLCandidate(req.Msg.GetKnowledge()), docsIndexURLCandidate(req.Msg.GetTitle())),
Content: req.Msg.GetKnowledge(),
Source: docsIndexSourceLocal,
}); err != nil {
logger.Errorf("docs index sync failed operation=update id=%s err=%v", strings.TrimSpace(req.Msg.GetId()), err)
}
}
return connect.NewResponse(&aiserverv1.KnowledgeBaseUpdateResponse{Success: true}), nil
}
func (service *Service) KnowledgeBaseRemove(_ context.Context, req *connect.Request[aiserverv1.KnowledgeBaseRemoveRequest]) (*connect.Response[aiserverv1.KnowledgeBaseRemoveResponse], error) {
if service == nil || service.rules == nil {
return nil, connect.NewError(connect.CodeInternal, fmt.Errorf("user rule store is not initialized"))
}
if _, err := normalizeUserRuleID(req.Msg.GetId()); err != nil {
return nil, connect.NewError(connect.CodeInvalidArgument, err)
}
if err := service.rules.Remove(req.Msg.GetId()); err != nil {
return nil, connect.NewError(connect.CodeInternal, err)
}
if service.docsIndexStore != nil {
if err := service.docsIndexStore.Remove(req.Msg.GetId()); err != nil {
logger.Errorf("docs index sync failed operation=remove id=%s err=%v", strings.TrimSpace(req.Msg.GetId()), err)
}
}
return connect.NewResponse(&aiserverv1.KnowledgeBaseRemoveResponse{Success: true}), nil
}
+455
View File
@@ -0,0 +1,455 @@
package forwarder
import (
"encoding/json"
"fmt"
"strings"
"google.golang.org/protobuf/encoding/protojson"
"cursor/gen/agentv1"
execbridge "cursor/internal/backend/agent/bridge/exec"
runtimecore "cursor/internal/backend/agent/core"
)
const (
writeReadExecKind = "write_read"
writeWriteExecKind = "write_write"
writePostReadExecKind = "write_post_read"
)
type writeOperationArgs struct {
Path string `json:"path"`
Contents string `json:"contents"`
}
type pendingWritePayload struct {
VisibleArgs writeOperationArgs `json:"visible_args"`
ResolvedPath string `json:"resolved_path"`
BeforeContent string `json:"before_content,omitempty"`
AfterContent string `json:"after_content,omitempty"`
}
func isHiddenWriteExecKind(kind string) bool {
switch strings.TrimSpace(kind) {
case writeReadExecKind, writeWriteExecKind, writePostReadExecKind:
return true
default:
return false
}
}
func (service *Service) handleWriteToolInvocation(stream *ActiveStream, invocation runtimecore.ToolInvocation) error {
if stream == nil {
return fmt.Errorf("write stream is required")
}
writeArgs, err := decodeWriteOperationArgs(invocation.ArgsJSON)
if err != nil {
return newRecoverableToolInvocationError(err)
}
resolvedPath := strings.TrimSpace(writeArgs.Path)
writeArgs.Path = resolvedPath
visibleArgsJSON, err := writeArgs.MarshalJSON()
if err != nil {
return err
}
invocation.ArgsJSON = visibleArgsJSON
startedToolCall := buildStartedToolCall(invocation)
if startedToolCall != nil {
historyStartedToolCall := buildStartedWriteHistoryToolCall(writeArgs.Path)
toolCallPayload, err := protojson.Marshal(historyStartedToolCall)
if err != nil {
return err
}
_, err = service.appendConversationEntries(stream, stream.ConversationID, []HistoryEntry{
newToolCallEntryWithProviderMetadata(stream.TurnSeq, stream.RequestID, invocation.CallID, invocation.ToolName, invocation.ReasoningContent, invocation.ReasoningSignature, invocation.ReasoningSignatureSource, invocation.ReasoningProviderItemID, invocation.ReasoningProviderStatus, invocation.ReasoningProviderSummary, invocation.ProviderItemID, invocation.ProviderCallID, invocation.ProviderStatus, toolCallPayload),
})
if err != nil {
return err
}
}
if err := service.broker.Publish(stream.RequestID, StreamEvent{
Message: buildToolCallStartedMessage(invocation.CallID, invocation.ModelCallID, startedToolCall),
}); err != nil {
return err
}
return service.startHiddenWriteRead(stream, invocation.CallID, invocation.ModelCallID, currentProviderPass(stream), invocation.ReasoningContent, invocation.ReasoningSignature, invocation.ReasoningSignatureSource, pendingWritePayload{
VisibleArgs: writeArgs,
ResolvedPath: resolvedPath,
})
}
func (service *Service) startHiddenWriteRead(stream *ActiveStream, toolCallID string, modelCallID string, providerPass int, reasoningContent string, reasoningSignature string, reasoningSignatureSource string, payload pendingWritePayload) error {
readArgsJSON, err := json.Marshal(map[string]any{
"path": strings.TrimSpace(payload.ResolvedPath),
})
if err != nil {
return err
}
serverMessage, pendingExec, err := service.execBridge.OpenExec(execbridge.OpenExecContext{
ConversationID: stream.ConversationID,
}, runtimecore.ToolInvocation{
CallID: strings.TrimSpace(toolCallID),
ToolName: "Read",
ArgsJSON: readArgsJSON,
ModelCallID: strings.TrimSpace(modelCallID),
})
if err != nil {
return err
}
pendingArgsJSON, err := payload.MarshalJSON()
if err != nil {
return err
}
pendingExec.ModelCallID = strings.TrimSpace(modelCallID)
pendingExec.ProviderPass = providerPass
pendingExec.ToolCallID = strings.TrimSpace(toolCallID)
pendingExec.ReasoningContent = reasoningContent
pendingExec.ReasoningSignature = strings.TrimSpace(reasoningSignature)
pendingExec.ReasoningSignatureSource = strings.TrimSpace(reasoningSignatureSource)
pendingExec.ExecKind = writeReadExecKind
pendingExec.ArgsJSON = pendingArgsJSON
stream.mu.Lock()
stream.PendingExecs[pendingExec.ExecID] = pendingExec
stream.mu.Unlock()
if err := service.publishCheckpoint(stream.RequestID, stream.ConversationID); err != nil {
return err
}
return service.broker.Publish(stream.RequestID, StreamEvent{Message: serverMessage})
}
func (service *Service) startHiddenWriteExec(stream *ActiveStream, toolCallID string, modelCallID string, providerPass int, reasoningContent string, reasoningSignature string, reasoningSignatureSource string, payload pendingWritePayload, beforeContent string) error {
transportContents := prepareWriteContentsForClient(payload.ResolvedPath, payload.VisibleArgs.Contents)
writeArgsJSON, err := json.Marshal(map[string]any{
"path": strings.TrimSpace(payload.ResolvedPath),
"contents": transportContents,
})
if err != nil {
return err
}
serverMessage, pendingExec, err := service.execBridge.OpenExec(execbridge.OpenExecContext{
ConversationID: stream.ConversationID,
}, runtimecore.ToolInvocation{
CallID: strings.TrimSpace(toolCallID),
ToolName: "Write",
ArgsJSON: writeArgsJSON,
ModelCallID: strings.TrimSpace(modelCallID),
})
if err != nil {
return err
}
payload.BeforeContent = beforeContent
payload.AfterContent = payload.VisibleArgs.Contents
pendingArgsJSON, err := payload.MarshalJSON()
if err != nil {
return err
}
pendingExec.ModelCallID = strings.TrimSpace(modelCallID)
pendingExec.ProviderPass = providerPass
pendingExec.ToolCallID = strings.TrimSpace(toolCallID)
pendingExec.ReasoningContent = reasoningContent
pendingExec.ReasoningSignature = strings.TrimSpace(reasoningSignature)
pendingExec.ReasoningSignatureSource = strings.TrimSpace(reasoningSignatureSource)
pendingExec.ExecKind = writeWriteExecKind
pendingExec.ArgsJSON = pendingArgsJSON
stream.mu.Lock()
stream.PendingExecs[pendingExec.ExecID] = pendingExec
stream.mu.Unlock()
if err := service.publishCheckpoint(stream.RequestID, stream.ConversationID); err != nil {
return err
}
return service.broker.Publish(stream.RequestID, StreamEvent{Message: serverMessage})
}
func (service *Service) startHiddenWritePostRead(stream *ActiveStream, toolCallID string, modelCallID string, providerPass int, reasoningContent string, reasoningSignature string, reasoningSignatureSource string, payload pendingWritePayload) error {
readArgsJSON, err := json.Marshal(map[string]any{
"path": strings.TrimSpace(payload.ResolvedPath),
})
if err != nil {
return err
}
serverMessage, pendingExec, err := service.execBridge.OpenExec(execbridge.OpenExecContext{
ConversationID: stream.ConversationID,
}, runtimecore.ToolInvocation{
CallID: strings.TrimSpace(toolCallID),
ToolName: "Read",
ArgsJSON: readArgsJSON,
ModelCallID: strings.TrimSpace(modelCallID),
})
if err != nil {
return err
}
pendingArgsJSON, err := payload.MarshalJSON()
if err != nil {
return err
}
pendingExec.ModelCallID = strings.TrimSpace(modelCallID)
pendingExec.ProviderPass = providerPass
pendingExec.ToolCallID = strings.TrimSpace(toolCallID)
pendingExec.ReasoningContent = reasoningContent
pendingExec.ReasoningSignature = strings.TrimSpace(reasoningSignature)
pendingExec.ReasoningSignatureSource = strings.TrimSpace(reasoningSignatureSource)
pendingExec.ExecKind = writePostReadExecKind
pendingExec.ArgsJSON = pendingArgsJSON
stream.mu.Lock()
stream.PendingExecs[pendingExec.ExecID] = pendingExec
stream.mu.Unlock()
if err := service.publishCheckpoint(stream.RequestID, stream.ConversationID); err != nil {
return err
}
return service.broker.Publish(stream.RequestID, StreamEvent{Message: serverMessage})
}
func (service *Service) handleHiddenWriteExecResult(stream *ActiveStream, pending runtimecore.PendingExec, message *agentv1.ExecClientMessage) error {
if stream == nil || message == nil {
return fmt.Errorf("write exec result requires stream and message")
}
payload, err := decodePendingWritePayload(pending.ArgsJSON)
if err != nil {
return err
}
switch strings.TrimSpace(pending.ExecKind) {
case writeReadExecKind:
markExecCompleted(stream, pending)
readResult := message.GetReadResult()
if readResult == nil {
return service.finishWriteOperation(stream, pending.ToolCallID, pending.ModelCallID, pending.ProviderPass, pending.ReasoningContent, payload.VisibleArgs, buildEditErrorResult(payload.ResolvedPath, "write read result missing"))
}
switch readResult.GetResult().(type) {
case *agentv1.ReadResult_FileNotFound:
return service.startHiddenWriteExec(stream, pending.ToolCallID, pending.ModelCallID, pending.ProviderPass, pending.ReasoningContent, pending.ReasoningSignature, pending.ReasoningSignatureSource, payload, "")
default:
content, ok := extractReadContentForEdit(readResult)
if !ok {
return service.finishWriteOperation(stream, pending.ToolCallID, pending.ModelCallID, pending.ProviderPass, pending.ReasoningContent, payload.VisibleArgs, buildEditResultFromReadResult(payload.ResolvedPath, readResult))
}
return service.startHiddenWriteExec(stream, pending.ToolCallID, pending.ModelCallID, pending.ProviderPass, pending.ReasoningContent, pending.ReasoningSignature, pending.ReasoningSignatureSource, payload, content)
}
case writeWriteExecKind:
markExecCompleted(stream, pending)
writeResult := message.GetWriteResult()
if writeResult == nil {
return service.finishWriteOperation(stream, pending.ToolCallID, pending.ModelCallID, pending.ProviderPass, pending.ReasoningContent, payload.VisibleArgs, buildEditErrorResult(payload.ResolvedPath, "write result missing"))
}
switch item := writeResult.GetResult().(type) {
case *agentv1.WriteResult_Success:
payload.ResolvedPath = firstNonEmpty(strings.TrimSpace(item.Success.GetPath()), strings.TrimSpace(payload.ResolvedPath), strings.TrimSpace(payload.VisibleArgs.Path))
if item.Success.GetFileContentAfterWrite() != "" {
payload.AfterContent, _ = reconcilePostWriteObservedContent(payload.AfterContent, item.Success.GetFileContentAfterWrite())
}
payload.VisibleArgs.Path = payload.ResolvedPath
return service.startHiddenWritePostRead(stream, pending.ToolCallID, pending.ModelCallID, pending.ProviderPass, pending.ReasoningContent, pending.ReasoningSignature, pending.ReasoningSignatureSource, payload)
default:
return service.finishWriteOperation(stream, pending.ToolCallID, pending.ModelCallID, pending.ProviderPass, pending.ReasoningContent, payload.VisibleArgs, buildEditResultFromWriteResult(payload.ResolvedPath, writeResult))
}
case writePostReadExecKind:
markExecCompleted(stream, pending)
writeArgs := payload.VisibleArgs
writeArgs.Path = firstNonEmpty(strings.TrimSpace(payload.ResolvedPath), strings.TrimSpace(writeArgs.Path))
readResult := message.GetReadResult()
if readResult == nil {
return service.finishWriteOperation(stream, pending.ToolCallID, pending.ModelCallID, pending.ProviderPass, pending.ReasoningContent, writeArgs, buildSuccessfulWriteResult(writeArgs.Path, payload.BeforeContent, payload.AfterContent))
}
if content, ok := extractReadContentForEdit(readResult); ok {
if success := readResult.GetSuccess(); success != nil {
writeArgs.Path = firstNonEmpty(strings.TrimSpace(success.GetPath()), writeArgs.Path)
}
finalAfterContent, _ := reconcilePostWriteObservedContent(payload.AfterContent, content)
return service.finishWriteOperation(stream, pending.ToolCallID, pending.ModelCallID, pending.ProviderPass, pending.ReasoningContent, writeArgs, buildSuccessfulWriteResult(writeArgs.Path, payload.BeforeContent, finalAfterContent))
}
return service.finishWriteOperation(stream, pending.ToolCallID, pending.ModelCallID, pending.ProviderPass, pending.ReasoningContent, writeArgs, buildSuccessfulWriteResult(writeArgs.Path, payload.BeforeContent, payload.AfterContent))
default:
return fmt.Errorf("unsupported hidden write exec kind: %s", pending.ExecKind)
}
}
func (service *Service) handleHiddenWriteExecControl(stream *ActiveStream, pending runtimecore.PendingExec, message *agentv1.ExecClientControlMessage) error {
if stream == nil || message == nil {
return fmt.Errorf("write exec control requires stream and message")
}
if _, ok := message.GetMessage().(*agentv1.ExecClientControlMessage_Heartbeat); ok {
return nil
}
if _, ok := message.GetMessage().(*agentv1.ExecClientControlMessage_StreamClose); ok {
return nil
}
payload, err := decodePendingWritePayload(pending.ArgsJSON)
if err != nil {
return err
}
markExecCompleted(stream, pending)
if strings.TrimSpace(pending.ExecKind) == writePostReadExecKind {
return service.finishWriteOperation(stream, pending.ToolCallID, pending.ModelCallID, pending.ProviderPass, pending.ReasoningContent, payload.VisibleArgs, buildSuccessfulWriteResult(payload.ResolvedPath, payload.BeforeContent, payload.AfterContent))
}
return service.finishWriteOperation(stream, pending.ToolCallID, pending.ModelCallID, pending.ProviderPass, pending.ReasoningContent, payload.VisibleArgs, buildEditErrorResult(payload.ResolvedPath, hiddenWriteControlError(message)))
}
func (service *Service) finishWriteOperation(stream *ActiveStream, toolCallID string, modelCallID string, providerPass int, reasoningContent string, writeArgs writeOperationArgs, result *agentv1.EditResult) error {
if stream == nil {
return nil
}
if result == nil {
result = buildEditErrorResult(writeArgs.Path, "write result missing")
}
writeArgs.Path = firstNonEmpty(resultPath(result), writeArgs.Path)
historyArgsJSON, err := writeArgs.MarshalJSON()
if err != nil {
return err
}
toolCall := buildCompletedWriteToolCall(writeArgs.Path, writeArgs.Contents, result)
historyToolCall := buildCompletedWriteHistoryToolCall(writeArgs.Path, result)
if err := service.appendToolResult(stream, strings.TrimSpace(toolCallID), "Write", historyArgsJSON, summarizeWriteHistoryResult(writeArgs.Path, result), reasoningContent, historyToolCall); err != nil {
return err
}
if err := service.publishToolCallCompleted(stream.RequestID, strings.TrimSpace(toolCallID), strings.TrimSpace(modelCallID), toolCall); err != nil {
return err
}
if err := service.syncSummaryCarryForward(stream.ConversationID, stream.RequestID, strings.TrimSpace(modelCallID)); err != nil {
return err
}
if err := service.publishCheckpoint(stream.RequestID, stream.ConversationID); err != nil {
return err
}
return service.reconcileStream(stream)
}
func decodeWriteOperationArgs(raw []byte) (writeOperationArgs, error) {
var args map[string]any
if err := json.Unmarshal(raw, &args); err != nil {
return writeOperationArgs{}, fmt.Errorf("decode write args failed: %w", err)
}
result := writeOperationArgs{
Path: firstNonEmpty(readStringAny(args["path"]), readStringAny(args["file_path"])),
Contents: firstNonEmpty(readStringAny(args["contents"]), readStringAny(args["content"]), readStringAny(args["stream_content"]), readStringAny(args["streamContent"])),
}
if strings.TrimSpace(result.Path) == "" {
return writeOperationArgs{}, fmt.Errorf("write path is required")
}
if !isAbsoluteToolPath(result.Path) {
return writeOperationArgs{}, fmt.Errorf("write path must be absolute")
}
return result, nil
}
func (args writeOperationArgs) MarshalJSON() ([]byte, error) {
return json.Marshal(map[string]any{
"path": strings.TrimSpace(args.Path),
"contents": args.Contents,
})
}
func (payload pendingWritePayload) MarshalJSON() ([]byte, error) {
type alias pendingWritePayload
return json.Marshal(alias(payload))
}
func decodePendingWritePayload(raw []byte) (pendingWritePayload, error) {
var payload pendingWritePayload
if err := json.Unmarshal(raw, &payload); err != nil {
return pendingWritePayload{}, fmt.Errorf("decode write pending payload failed: %w", err)
}
return payload, nil
}
func buildSuccessfulWriteResult(path string, beforeContent string, afterContent string) *agentv1.EditResult {
diffString, linesAdded, linesRemoved := computeEditDiff(beforeContent, afterContent)
return buildSuccessfulEditResult(path, beforeContent, afterContent, diffString, linesAdded, linesRemoved, "")
}
func buildCompletedWriteToolCall(path string, streamContent string, result *agentv1.EditResult) *agentv1.ToolCall {
return &agentv1.ToolCall{
Tool: &agentv1.ToolCall_EditToolCall{
EditToolCall: &agentv1.EditToolCall{
Args: &agentv1.EditArgs{
Path: strings.TrimSpace(path),
StreamContent: literalStringPtr(streamContent),
},
Result: result,
},
},
}
}
func buildStartedWriteHistoryToolCall(path string) *agentv1.ToolCall {
return &agentv1.ToolCall{
Tool: &agentv1.ToolCall_EditToolCall{
EditToolCall: &agentv1.EditToolCall{
Args: &agentv1.EditArgs{
Path: strings.TrimSpace(path),
},
},
},
}
}
func buildCompletedWriteHistoryToolCall(path string, result *agentv1.EditResult) *agentv1.ToolCall {
return buildCompletedEditToolCall(path, compactWriteHistoryEditResult(path, result))
}
func summarizeWriteHistoryResult(path string, result *agentv1.EditResult) string {
if result == nil {
return "write result missing"
}
success := result.GetSuccess()
if success == nil {
return summarizeEditResultWithoutPath(result)
}
encoded, err := json.Marshal(map[string]any{
"success": map[string]any{
"path": firstNonEmpty(success.GetPath(), path),
},
})
if err != nil {
return fmt.Sprintf(`{"success":{"path":%q}}`, firstNonEmpty(success.GetPath(), path))
}
return string(encoded)
}
func compactWriteHistoryEditResult(path string, result *agentv1.EditResult) *agentv1.EditResult {
if result == nil {
return buildEditErrorResult("", "write result missing")
}
success := result.GetSuccess()
if success == nil {
return editResultWithoutPath(result)
}
return &agentv1.EditResult{
Result: &agentv1.EditResult_Success{
Success: &agentv1.EditSuccess{
Path: firstNonEmpty(success.GetPath(), path),
},
},
}
}
func boundedWriteDiffString(diffString string) string {
return truncateProjectedReplayText("Write", diffString, projectedEditReplayLimit)
}
func resultPath(result *agentv1.EditResult) string {
if result == nil {
return ""
}
switch item := result.GetResult().(type) {
case *agentv1.EditResult_Success:
return strings.TrimSpace(item.Success.GetPath())
case *agentv1.EditResult_FileNotFound:
return strings.TrimSpace(item.FileNotFound.GetPath())
case *agentv1.EditResult_ReadPermissionDenied:
return strings.TrimSpace(item.ReadPermissionDenied.GetPath())
case *agentv1.EditResult_WritePermissionDenied:
return strings.TrimSpace(item.WritePermissionDenied.GetPath())
case *agentv1.EditResult_Rejected:
return strings.TrimSpace(item.Rejected.GetPath())
case *agentv1.EditResult_Error:
return strings.TrimSpace(item.Error.GetPath())
default:
return ""
}
}
+775
View File
@@ -0,0 +1,775 @@
package backend
import (
"context"
"fmt"
"net"
"net/http"
"net/http/httptest"
"net/url"
"strings"
"sync"
"time"
"cursor/internal/ads"
"cursor/internal/appdata"
"cursor/internal/backend/forwarder"
"cursor/internal/backend/server"
serverconfig "cursor/internal/backend/server/config"
"cursor/internal/backend/server/upstream"
"cursor/internal/logger"
"cursor/internal/netproxy"
legacyruntime "cursor/internal/runtime"
)
const healthPath = "/healthz"
const tabServerBaseURL = "https://tab.leokun.cn"
type Host struct {
store *serverconfig.Store
listenAddr string
configs *serverconfig.Manager
healthHTTP *http.Client
runMu sync.RWMutex
httpServer *http.Server
lastRunErr error
mux http.Handler
}
func NewHost(store *serverconfig.Store) (*Host, error) {
if store == nil {
return nil, fmt.Errorf("backend config store is required")
}
configs, err := serverconfig.NewManager(context.Background(), store)
if err != nil {
return nil, err
}
cfg := configs.Current()
host := &Host{
store: store,
listenAddr: cfg.BackendListenAddr,
configs: configs,
healthHTTP: newLoopbackHTTPClient(),
}
if err := host.rebuild(cfg); err != nil {
return nil, err
}
return host, nil
}
func (host *Host) ConfigManager() *serverconfig.Manager {
if host == nil {
return nil
}
return host.configs
}
func (host *Host) LoadConfig(ctx context.Context) (serverconfig.Config, error) {
if host == nil || host.configs == nil {
return serverconfig.DefaultConfig(), nil
}
return host.configs.Load(ctx)
}
func (host *Host) SaveConfig(ctx context.Context, cfg serverconfig.Config) (serverconfig.Config, error) {
if host == nil || host.configs == nil {
return serverconfig.Config{}, fmt.Errorf("backend config manager is not initialized")
}
normalized, err := host.configs.Save(ctx, cfg)
if err != nil {
return serverconfig.Config{}, err
}
if host.httpServer == nil {
if rebuildErr := host.rebuild(normalized); rebuildErr != nil {
return serverconfig.Config{}, rebuildErr
}
}
return normalized, nil
}
func (host *Host) ListenAddr() string {
if host == nil {
return ""
}
host.runMu.RLock()
defer host.runMu.RUnlock()
return host.listenAddr
}
func (host *Host) BaseURL() string {
listenAddr := strings.TrimSpace(host.ListenAddr())
if listenAddr == "" {
return ""
}
return "http://" + listenAddr
}
func (host *Host) IsRunning() bool {
if host == nil {
return false
}
host.runMu.RLock()
defer host.runMu.RUnlock()
return host.httpServer != nil
}
func (host *Host) LastRunError() error {
if host == nil {
return nil
}
host.runMu.RLock()
defer host.runMu.RUnlock()
return host.lastRunErr
}
func (host *Host) Start() error {
if host == nil {
return fmt.Errorf("backend host is nil")
}
cfg := host.configs.Current()
host.runMu.Lock()
defer host.runMu.Unlock()
if host.httpServer != nil {
return fmt.Errorf("backend is already running")
}
if err := host.rebuildLocked(cfg); err != nil {
return err
}
httpServer := &http.Server{
Addr: host.listenAddr,
Handler: host.mux,
ReadHeaderTimeout: 10 * time.Second,
}
listener, err := net.Listen("tcp", host.listenAddr)
if err != nil {
host.lastRunErr = fmt.Errorf("监听内置后端 %s 失败: %w", host.listenAddr, err)
return host.lastRunErr
}
host.listenAddr = listener.Addr().String()
host.httpServer = httpServer
host.lastRunErr = nil
logger.Infof("内置后端监听成功 listen_addr=%s", host.listenAddr)
go func(serverInstance *http.Server, serverListener net.Listener) {
logger.Infof("内置后端开始提供服务 listen_addr=%s", serverListener.Addr().String())
if err := serverInstance.Serve(serverListener); err != nil && err != http.ErrServerClosed {
runErr := fmt.Errorf("内置后端在 %s 上异常退出: %w", serverListener.Addr().String(), err)
host.runMu.Lock()
if host.httpServer == serverInstance {
host.httpServer = nil
}
host.lastRunErr = runErr
host.runMu.Unlock()
logger.Errorf("%v", runErr)
}
}(httpServer, listener)
return nil
}
func (host *Host) Stop(ctx context.Context) error {
if host == nil {
return nil
}
host.runMu.Lock()
serverInstance := host.httpServer
host.httpServer = nil
host.runMu.Unlock()
if serverInstance == nil {
return nil
}
err := serverInstance.Shutdown(ctx)
return err
}
func (host *Host) HealthCheck(ctx context.Context) error {
if host == nil {
return fmt.Errorf("backend host is nil")
}
if runErr := host.LastRunError(); runErr != nil {
return runErr
}
request, err := http.NewRequestWithContext(ctx, http.MethodGet, host.BaseURL()+healthPath, nil)
if err != nil {
return err
}
client := host.healthHTTP
if client == nil {
client = newLoopbackHTTPClient()
}
response, err := client.Do(request)
if err != nil {
inProcessErr := host.InProcessHealthCheck()
if inProcessErr == nil {
logger.Errorf("内置后端进程内健康检查成功,但 loopback 访问失败 base_url=%s err=%v", host.BaseURL(), err)
return fmt.Errorf("内置后端进程内健康检查成功,但本机 loopback 访问失败: %w", err)
}
logger.Errorf("内置后端 loopback 与进程内健康检查均失败 loopback_err=%v in_process_err=%v", err, inProcessErr)
if runErr := host.LastRunError(); runErr != nil {
return runErr
}
return err
}
defer response.Body.Close()
if response.StatusCode != http.StatusOK {
return fmt.Errorf("内置后端健康检查返回状态码 %d", response.StatusCode)
}
return nil
}
func (host *Host) InProcessHealthCheck() error {
if host == nil {
return fmt.Errorf("backend host is nil")
}
if host.mux == nil {
return fmt.Errorf("backend handler is nil")
}
request := httptest.NewRequest(http.MethodGet, "http://inprocess"+healthPath, nil)
recorder := httptest.NewRecorder()
host.mux.ServeHTTP(recorder, request)
if recorder.Code != http.StatusOK {
return fmt.Errorf("in-process health status %d", recorder.Code)
}
body := strings.TrimSpace(recorder.Body.String())
if body != "ok" {
return fmt.Errorf("in-process health body %q", body)
}
logger.Infof("内置后端进程内健康检查成功")
return nil
}
func newLoopbackHTTPClient() *http.Client {
return &http.Client{
Transport: &http.Transport{
Proxy: nil,
DialContext: (&net.Dialer{
Timeout: 1 * time.Second,
KeepAlive: 30 * time.Second,
}).DialContext,
ForceAttemptHTTP2: false,
MaxIdleConns: 1,
MaxIdleConnsPerHost: 1,
IdleConnTimeout: 30 * time.Second,
},
}
}
func (host *Host) rebuild(cfg serverconfig.Config) error {
host.runMu.Lock()
defer host.runMu.Unlock()
return host.rebuildLocked(cfg)
}
func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
host.listenAddr = cfg.BackendListenAddr
agentModule := forwarder.NewModule(appdata.HistoryRootPath(), host.configs)
legacyBidiAppendProcedure := "/aiserver.v1.BidiService/BidiAppend"
legacyRunSSEProcedure := "/agent.v1.AgentService/RunSSE"
routeDeps := upstream.Dependencies{
SystemSettingService: &serverSystemSettings{configs: host.configs},
HTTPClient: netproxy.NewHTTPClient(30000 * time.Second),
}
host.mux = server.New(
server.Use(
server.Recover(),
server.ServerContext(),
server.PolicyMiddleware(host.configs),
server.ErrorEncoder(),
),
server.Mount(ads.RoutePrefix, ads.NewHTTPHandler(appdata.AdsRootPath())),
server.GET(healthPath,
server.Name("healthz"),
server.HTTP(),
server.Local(server.Health()),
),
server.POST(legacyBidiAppendProcedure,
server.Name("bidi_append"),
server.ConnectUnary(),
server.Local(server.HTTPHandlerAction(agentModule.LocalBidiHandler)),
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
Name: "bidi_append",
})),
),
server.POST(legacyRunSSEProcedure,
server.Name("run_sse"),
server.ConnectStream(),
server.Local(server.HTTPHandlerAction(agentModule.LocalRunSSE)),
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
Name: "run_sse",
})),
),
server.POST("/aiserver.v1.AiService/ServerTime",
server.Name("server_time"),
server.ConnectUnary(),
server.Local(upstream.MockProtoAction(routeDeps, upstream.CompatRouteConfig{
Name: "server_time",
StatusCode: http.StatusOK,
MockProtoType: "aiserver.v1.ServerTimeResponse",
MockBuilder: upstream.ServerTimeMockBuilder,
})),
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
Name: "server_time",
})),
),
server.POST("/aiserver.v1.AiService/GetServerConfig",
server.Name("server_config"),
server.ConnectUnary(),
server.Local(upstream.MockProtoAction(routeDeps, upstream.CompatRouteConfig{
Name: "server_config",
StatusCode: http.StatusOK,
MockProtoType: "aiserver.v1.GetServerConfigResponse",
MockBuilder: upstream.ServerConfigMockBuilder,
})),
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
Name: "server_config",
})),
),
server.POST("/aiserver.v1.AiService/AvailableModels",
server.Name("available_models"),
server.ConnectUnary(),
server.Local(upstream.MockProtoAction(routeDeps, upstream.CompatRouteConfig{
Name: "available_models",
StatusCode: http.StatusOK,
MockProtoType: "aiserver.v1.AvailableModelsResponse",
MockBuilder: upstream.AvailableModelsMockBuilder,
})),
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
Name: "available_models",
})),
),
server.POST("/aiserver.v1.AiService/GetDefaultModelNudgeData",
server.Name("default_model_nudge"),
server.ConnectUnary(),
server.Local(upstream.MockProtoAction(routeDeps, upstream.CompatRouteConfig{
Name: "default_model_nudge",
StatusCode: http.StatusOK,
MockProtoType: "aiserver.v1.GetDefaultModelNudgeDataResponse",
MockBuilder: upstream.DefaultModelNudgeMockBuilder,
})),
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
Name: "default_model_nudge",
})),
),
server.POST("/aiserver.v1.AnalyticsService/BootstrapStatsig",
server.Name("bootstrap_statsig"),
server.ConnectUnary(),
server.Local(upstream.MockProtoAction(routeDeps, upstream.CompatRouteConfig{
Name: "bootstrap_statsig",
StatusCode: http.StatusOK,
MockProtoType: "aiserver.v1.BootstrapStatsigResponse",
MockBuilder: upstream.BootstrapStatsigMockBuilder,
})),
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
Name: "bootstrap_statsig",
})),
),
server.POST("/aiserver.v1.AnalyticsService/GetFirstWindowStatsigDecision",
server.Name("first_window_statsig_decision"),
server.ConnectUnary(),
server.Local(upstream.MockProtoAction(routeDeps, upstream.CompatRouteConfig{
Name: "first_window_statsig_decision",
StatusCode: http.StatusOK,
MockProtoType: "aiserver.v1.GetFirstWindowStatsigDecisionResponse",
MockBuilder: upstream.FirstWindowStatsigDecisionMockBuilder,
})),
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
Name: "first_window_statsig_decision",
})),
),
server.POST("/oauth/token",
server.Name("oauth_token"),
server.HTTP(),
server.Local(upstream.MockOAuthAction(routeDeps, upstream.CompatRouteConfig{
Name: "oauth_token",
StatusCode: http.StatusOK,
})),
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
Name: "oauth_token",
})),
),
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,
})),
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
Name: "auth_service_get_email",
})),
),
tabServerUpstreamProcedure("/aiserver.v1.AiService/StreamCpp", "ai_stream_cpp", server.ConnectStream(), routeDeps),
tabServerUpstreamProcedure("/aiserver.v1.AiService/StreamNextCursorPrediction", "ai_stream_next_cursor_prediction", server.ConnectStream(), routeDeps),
tabServerUpstreamProcedure("/aiserver.v1.AiService/GetCppEditClassification", "ai_get_cpp_edit_classification", server.ConnectUnary(), routeDeps),
tabServerUpstreamProcedure("/aiserver.v1.AiService/RefreshTabContext", "ai_refresh_tab_context", server.ConnectUnary(), routeDeps),
tabServerUpstreamProcedure("/aiserver.v1.AiService/CppConfig", "ai_cpp_config", server.ConnectUnary(), routeDeps),
tabServerUpstreamProcedure("/aiserver.v1.AiService/CppEditHistoryStatus", "ai_cpp_edit_history_status", server.ConnectUnary(), routeDeps),
tabServerUpstreamProcedure("/aiserver.v1.AiService/CppAppend", "ai_cpp_append", server.ConnectUnary(), routeDeps),
tabServerUpstreamProcedure("/aiserver.v1.AiService/CppEditHistoryAppend", "ai_cpp_edit_history_append", server.ConnectUnary(), routeDeps),
tabServerUpstreamProcedure("/aiserver.v1.AiService/ReportAiCodeChangeMetrics", "ai_report_ai_code_change_metrics", server.ConnectUnary(), routeDeps),
tabServerUpstreamProcedure("/aiserver.v1.AiService/WriteGitCommitMessage", "ai_write_git_commit_message", server.ConnectUnary(), routeDeps),
tabServerUpstreamProcedure("/aiserver.v1.AiService/WriteGitBranchName", "ai_write_git_branch_name", server.ConnectUnary(), routeDeps),
repositoryServiceProcedure(forwarder.RepositoryServiceFastRepoInitHandshakeV2Procedure, "repository_fast_repo_init_handshake_v2", server.ConnectUnary(), agentModule, routeDeps),
repositoryServiceProcedure(forwarder.RepositoryServiceFastRepoInitHandshakeProcedure, "repository_fast_repo_init_handshake", server.ConnectUnary(), agentModule, routeDeps),
repositoryServiceProcedure(forwarder.RepositoryServiceFastRepoSyncCompleteProcedure, "repository_fast_repo_sync_complete", server.ConnectUnary(), agentModule, routeDeps),
repositoryServiceProcedure(forwarder.RepositoryServiceSyncMerkleSubtreeV2Procedure, "repository_sync_merkle_subtree_v2", server.ConnectUnary(), agentModule, routeDeps),
repositoryServiceProcedure(forwarder.RepositoryServiceSyncMerkleSubtreeProcedure, "repository_sync_merkle_subtree", server.ConnectUnary(), agentModule, routeDeps),
repositoryServiceProcedure(forwarder.RepositoryServiceFastUpdateFileV2Procedure, "repository_fast_update_file_v2", server.ConnectUnary(), agentModule, routeDeps),
repositoryServiceProcedure(forwarder.RepositoryServiceFastUpdateFileProcedure, "repository_fast_update_file", server.ConnectUnary(), agentModule, routeDeps),
repositoryServiceProcedure(forwarder.RepositoryServiceEnsureIndexCreatedProcedure, "repository_ensure_index_created", server.ConnectUnary(), agentModule, routeDeps),
repositoryServiceProcedure(forwarder.RepositoryServiceGetCopyStatusProcedure, "repository_get_copy_status", server.ConnectUnary(), agentModule, routeDeps),
repositoryServiceProcedure(forwarder.RepositoryServiceGetUploadLimitsProcedure, "repository_get_upload_limits", server.ConnectUnary(), agentModule, routeDeps),
repositoryServiceProcedure(forwarder.RepositoryServiceGetNumFilesToSendProcedure, "repository_get_num_files_to_send", server.ConnectUnary(), agentModule, routeDeps),
repositoryServiceProcedure(forwarder.RepositoryServiceGetAvailableChunkingStrategiesProcedure, "repository_get_available_chunking_strategies", server.ConnectUnary(), agentModule, routeDeps),
repositoryServiceProcedure(forwarder.RepositoryServiceGetHighLevelFolderDescriptionProcedure, "repository_get_high_level_folder_description", server.ConnectUnary(), agentModule, routeDeps),
repositoryServiceProcedure(forwarder.RepositoryServiceRepositoryStatusProcedure, "repository_status", server.ConnectUnary(), agentModule, routeDeps),
repositoryServiceProcedure(forwarder.RepositoryServiceBatchRepositoryStatusProcedure, "repository_batch_status", server.ConnectUnary(), agentModule, routeDeps),
uploadServiceProcedure(forwarder.UploadServiceUploadDocumentationProcedure, "upload_documentation", server.ConnectUnary(), agentModule, routeDeps),
uploadServiceProcedure(forwarder.UploadServiceGetDocProcedure, "upload_get_doc", server.ConnectUnary(), agentModule, routeDeps),
uploadServiceProcedure(forwarder.UploadServiceGetPagesProcedure, "upload_get_pages", server.ConnectUnary(), agentModule, routeDeps),
uploadServiceProcedure(forwarder.UploadServiceUploadedStatusProcedure, "upload_uploaded_status", server.ConnectUnary(), agentModule, routeDeps),
server.Any("/aiserver.v1.AiService/*",
server.Name("ai_service"),
server.HTTP(),
server.Local(server.HTTPHandlerAction(agentModule.AiHandler)),
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
Name: "ai_service",
})),
),
tabServerUpstreamProcedure("/aiserver.v1.CppService/AvailableModels", "cpp_available_models", server.ConnectUnary(), routeDeps),
tabServerUpstreamProcedure("/aiserver.v1.CppService/RecordCppFate", "cpp_record_cpp_fate", server.ConnectUnary(), routeDeps),
server.Any("/aiserver.v1.CppService/*",
server.Name("cpp_service"),
server.HTTP(),
server.Local(func(ctx *server.Context) error {
http.NotFound(ctx.Writer, ctx.Request)
return nil
}),
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
Name: "cpp_service",
})),
),
tabServerUpstreamProcedure("/aiserver.v1.FileSyncService/FSSyncFile", "file_sync_sync_file", server.ConnectUnary(), routeDeps),
tabServerUpstreamProcedure("/aiserver.v1.FileSyncService/FSIsEnabledForUser", "file_sync_is_enabled_for_user", server.ConnectUnary(), routeDeps),
tabServerUpstreamProcedure("/aiserver.v1.FileSyncService/FSConfig", "file_sync_config", server.ConnectUnary(), routeDeps),
tabServerUpstreamProcedure("/aiserver.v1.FileSyncService/FSUploadFile", "file_sync_upload_file", server.ConnectUnary(), routeDeps),
server.Any("/aiserver.v1.FileSyncService/*",
server.Name("file_sync"),
server.HTTP(),
server.Local(func(ctx *server.Context) error {
http.NotFound(ctx.Writer, ctx.Request)
return nil
}),
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
Name: "file_sync",
})),
),
server.POST("/aiserver.v1.DashboardService/GetTokenUsage",
server.Name("dashboard_token_usage"),
server.HTTP(),
server.Local(server.HTTPHandlerAction(agentModule.AiHandler)),
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
Name: "dashboard_token_usage",
})),
),
server.POST("/aiserver.v1.DashboardService/GetGlassEarlyPreviewEnrollment",
server.Name("dashboard_glass_early_preview_enrollment"),
server.ConnectUnary(),
server.Local(server.HTTPHandlerAction(agentModule.AiHandler)),
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
Name: "dashboard_glass_early_preview_enrollment",
})),
),
server.POST("/aiserver.v1.DashboardService/GetCurrentPeriodUsage",
server.Name("dashboard_current_period_usage"),
server.ConnectUnary(),
server.Local(upstream.MockProtoAction(routeDeps, upstream.CompatRouteConfig{
Name: "dashboard_current_period_usage",
StatusCode: http.StatusOK,
MockProtoType: "aiserver.v1.GetCurrentPeriodUsageResponse",
MockBuilder: upstream.DashboardCurrentPeriodUsageMockBuilder,
})),
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
Name: "dashboard_current_period_usage",
})),
),
server.POST("/aiserver.v1.DashboardService/GetTeams",
server.Name("dashboard_get_teams"),
server.ConnectUnary(),
server.Local(upstream.MockProtoAction(routeDeps, upstream.CompatRouteConfig{
Name: "dashboard_get_teams",
StatusCode: http.StatusOK,
MockProtoType: "aiserver.v1.GetTeamsResponse",
MockBuilder: upstream.DashboardTeamsMockBuilder,
})),
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
Name: "dashboard_get_teams",
})),
),
server.POST("/aiserver.v1.DashboardService/GetManagedSkills",
server.Name("dashboard_get_managed_skills"),
server.ConnectUnary(),
server.Local(upstream.MockProtoAction(routeDeps, upstream.CompatRouteConfig{
Name: "dashboard_get_managed_skills",
StatusCode: http.StatusOK,
MockProtoType: "aiserver.v1.GetManagedSkillsResponse",
MockBuilder: upstream.DashboardManagedSkillsMockBuilder,
})),
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
Name: "dashboard_get_managed_skills",
})),
),
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.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
Name: "dashboard_get_me",
})),
),
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.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
Name: "dashboard_user_privacy_mode",
})),
),
server.POST("/aiserver.v1.DashboardService/GetPlanInfo",
server.Name("dashboard_plan_info"),
server.ConnectUnary(),
server.Local(upstream.MockProtoAction(routeDeps, upstream.CompatRouteConfig{
Name: "dashboard_plan_info",
StatusCode: http.StatusOK,
MockProtoType: "aiserver.v1.GetPlanInfoResponse",
MockBuilder: upstream.DashboardPlanInfoMockBuilder,
})),
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
Name: "dashboard_plan_info",
})),
),
server.POST("/aiserver.v1.DashboardService/GetUsageLimitStatusAndActiveGrants",
server.Name("dashboard_usage_limit_status"),
server.ConnectUnary(),
server.Local(upstream.MockProtoAction(routeDeps, upstream.CompatRouteConfig{
Name: "dashboard_usage_limit_status",
StatusCode: http.StatusOK,
MockProtoType: "aiserver.v1.GetUsageLimitStatusAndActiveGrantsResponse",
MockBuilder: upstream.DashboardUsageLimitStatusAndActiveGrantsMockBuilder,
})),
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
Name: "dashboard_usage_limit_status",
})),
),
server.POST("/aiserver.v1.DashboardService/IsOnNewPricing",
server.Name("dashboard_is_on_new_pricing"),
server.ConnectUnary(),
server.Local(upstream.MockProtoAction(routeDeps, upstream.CompatRouteConfig{
Name: "dashboard_is_on_new_pricing",
StatusCode: http.StatusOK,
MockProtoType: "aiserver.v1.IsOnNewPricingResponse",
MockBuilder: upstream.DashboardIsOnNewPricingMockBuilder,
})),
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
Name: "dashboard_is_on_new_pricing",
})),
),
// tabServerUpstreamProcedure("/aiserver.v1.DashboardService/GetEffectiveUserPlugins", "dashboard_get_effective_user_plugins", server.ConnectUnary(), routeDeps),
server.Any("/aiserver.v1.DashboardService/*",
server.Name("dashboard"),
server.HTTP(),
server.Local(func(ctx *server.Context) error {
http.NotFound(ctx.Writer, ctx.Request)
return nil
}),
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
Name: "dashboard",
})),
),
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.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
Name: "network_service",
})),
),
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.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
Name: "in_app_ad",
})),
),
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.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
Name: "auth_full_stripe_profile",
})),
),
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.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
Name: "auth_stripe_profile",
})),
),
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.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
Name: "auth_has_valid_payment_method",
})),
),
server.Any("/auth/poll",
server.Name("auth_poll"),
server.HTTP(),
server.Local(upstream.MockAuthPollAction(routeDeps, upstream.CompatRouteConfig{
Name: "auth_poll",
StatusCode: http.StatusOK,
})),
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
Name: "auth_poll",
})),
),
server.POST("/auth/logout",
server.Name("auth_logout"),
server.HTTP(),
server.Local(upstream.FixedStatusAction(routeDeps, upstream.CompatRouteConfig{
Name: "auth_logout",
StatusCode: http.StatusNoContent,
})),
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
Name: "auth_logout",
})),
),
server.Any("/auth/*",
server.Name("auth_proxy"),
server.HTTP(),
server.Local(func(ctx *server.Context) error {
http.NotFound(ctx.Writer, ctx.Request)
return nil
}),
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
Name: "auth_proxy",
})),
),
)
return nil
}
func directUpstreamProcedure(pattern string, name string, protocol server.RouteOption, deps upstream.Dependencies) server.Option {
direct := upstream.DirectAction(deps, upstream.CompatRouteConfig{Name: name})
action := func(ctx *server.Context) error {
if ctx != nil && ctx.UpstreamURL == nil && ctx.Request != nil && ctx.Request.URL != nil {
targetURL := *ctx.Request.URL
targetURL.Scheme = "https"
targetURL.Host = "api2.cursor.sh:443"
ctx.UpstreamURL = &targetURL
}
return direct(ctx)
}
return server.POST(pattern,
server.Name(name),
protocol,
server.Local(action),
server.Upstream(action),
)
}
func repositoryServiceProcedure(pattern string, name string, protocol server.RouteOption, module *forwarder.Module, deps upstream.Dependencies) server.Option {
localAction := server.HTTPHandlerAction(module.RepositoryServiceHandler)
upstreamAction := upstream.DirectAction(deps, upstream.CompatRouteConfig{Name: name})
return server.POST(pattern,
server.Name(name),
protocol,
server.Local(localAction),
server.Upstream(upstreamAction),
)
}
func uploadServiceProcedure(pattern string, name string, protocol server.RouteOption, module *forwarder.Module, deps upstream.Dependencies) server.Option {
localAction := server.HTTPHandlerAction(module.UploadServiceHandler)
upstreamAction := upstream.DirectAction(deps, upstream.CompatRouteConfig{Name: name})
return server.POST(pattern,
server.Name(name),
protocol,
server.Local(localAction),
server.Upstream(upstreamAction),
)
}
func tabServerUpstreamProcedure(pattern string, name string, protocol server.RouteOption, deps upstream.Dependencies) server.Option {
direct := upstream.DirectAction(deps, upstream.CompatRouteConfig{Name: name})
action := func(ctx *server.Context) error {
if ctx != nil && ctx.Request != nil && ctx.Request.URL != nil {
baseURL, err := url.Parse(tabServerBaseURL)
if err != nil {
return fmt.Errorf("解析 tab server 地址失败: %w", err)
}
targetURL := *ctx.Request.URL
targetURL.Scheme = baseURL.Scheme
targetURL.Host = baseURL.Host
ctx.UpstreamURL = &targetURL
}
return direct(ctx)
}
return server.POST(pattern,
server.Name(name),
protocol,
server.Local(action),
server.Upstream(action),
)
}
type serverSystemSettings struct {
configs *serverconfig.Manager
}
func (settings *serverSystemSettings) ResolveModelAdapters(ctx context.Context) ([]legacyruntime.ModelAdapterConfig, error) {
snapshot, err := settings.configs.LegacyRuntimeSnapshot(ctx)
if err != nil {
return nil, err
}
return snapshot.ModelAdapters, nil
}
@@ -0,0 +1,46 @@
package config
import (
"context"
legacyruntime "cursor/internal/runtime"
)
func (store *Store) LegacyRuntimeSnapshot(ctx context.Context) (legacyruntime.RuntimeConfigSnapshot, error) {
cfg, err := store.Load(ctx)
if err != nil {
return legacyruntime.RuntimeConfigSnapshot{}, err
}
adapters := make([]legacyruntime.ModelAdapterConfig, 0, len(cfg.ModelAdapters))
for _, item := range cfg.ModelAdapters {
adapters = append(adapters, legacyruntime.ModelAdapterConfig{
ID: item.ID,
DisplayName: item.DisplayName,
Type: item.Type,
BaseURL: item.BaseURL,
APIKey: item.APIKey,
TooltipData: item.TooltipData,
ModelID: item.ModelID,
ReasoningEffort: item.ReasoningEffort,
OpenAIEndpoint: item.OpenAIEndpoint,
OpenAIExtraParamsEnabled: item.OpenAIExtraParamsEnabled,
OpenAIExtraParamsJSON: item.OpenAIExtraParamsJSON,
CustomHeadersEnabled: item.CustomHeadersEnabled,
CustomHeadersJSON: item.CustomHeadersJSON,
AnthropicExtraParamsEnabled: item.AnthropicExtraParamsEnabled,
AnthropicExtraParamsJSON: item.AnthropicExtraParamsJSON,
ContextWindowTokens: item.ContextWindowTokens,
MaxCompletionTokens: item.MaxCompletionTokens,
AnthropicMaxTokens: item.AnthropicMaxTokens,
AnthropicThinkingEffort: item.AnthropicThinkingEffort,
ThinkingBudgetTokens: item.ThinkingBudgetTokens,
})
}
return legacyruntime.RuntimeConfigSnapshot{
ObservabilityLogEnabled: cfg.Log,
ProviderStreamIdleTimeout: cfg.ProviderStreamIdleTimeout,
ModelAdapters: adapters,
}, nil
}

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