mirror of
https://wget.la/https://github.com/leookun/cursor-byok
synced 2026-08-18 03:57:06 +08:00
Merge branch 'main' into fix/cli-local-mode-endpoints
# Conflicts: # internal/backend/server/middleware.go
This commit is contained in:
@@ -1,6 +1,6 @@
|
||||
# Backend 架构说明
|
||||
|
||||
`internal/backend` 当前支持本地助手模式与直连上游模式。
|
||||
`internal/backend` 只支持本地助手模式。
|
||||
|
||||
关于 backend agent「最小事实集合」的第一阶段研究文档,见 [`../../docs/backend-agent-minimum-facts-phase1.md`](../../docs/backend-agent-minimum-facts-phase1.md)。
|
||||
|
||||
@@ -34,7 +34,6 @@ internal/backend/
|
||||
errors.go
|
||||
local.go
|
||||
middleware.go
|
||||
policy.go
|
||||
route.go
|
||||
url.go
|
||||
|
||||
@@ -134,7 +133,7 @@ history/
|
||||
## 请求流
|
||||
|
||||
1. 请求进入 backend 根路由。
|
||||
2. `PolicyMiddleware` 根据 `routing.mode` 与 `X-Server-Upstream-URL` 选择本地或上游分支。
|
||||
2. `ServerContext` 解析 MITM 带入的原始目标地址,路由始终执行本地 action。
|
||||
3. `BidiAppend` / `RunSSE` 进入 `forwarder`。
|
||||
4. `forwarder` 先把当前 loop 状态写入 `state.json`,再把已发生语义事件追加到 `context.json`。
|
||||
5. 发给 LLM 的 prompt 只由 `context.json` 投射生成;`state.json` 不保存可投射历史。
|
||||
|
||||
@@ -62,6 +62,46 @@ func buildUserReplayMessage(text string, selectedContext *agentv1.SelectedContex
|
||||
}, true
|
||||
}
|
||||
|
||||
// BuildSelectedCursorCommandsReplayMessage renders command content for new history entries.
|
||||
// Keeping this separate from BuildUserMessageReplayMessage prevents old user_message entries
|
||||
// from changing their model-visible meaning after a backend upgrade.
|
||||
func BuildSelectedCursorCommandsReplayMessage(userMessage *agentv1.UserMessage) (Message, bool) {
|
||||
if userMessage == nil {
|
||||
return Message{}, false
|
||||
}
|
||||
content := buildSelectedCursorCommandsPromptSection(userMessage.GetSelectedContext())
|
||||
if content == "" {
|
||||
return Message{}, false
|
||||
}
|
||||
return Message{Role: "user", Content: content}, true
|
||||
}
|
||||
|
||||
func buildSelectedCursorCommandsPromptSection(selectedContext *agentv1.SelectedContext) string {
|
||||
if selectedContext == nil || len(selectedContext.GetCursorCommands()) == 0 {
|
||||
return ""
|
||||
}
|
||||
entries := make([]string, 0, len(selectedContext.GetCursorCommands()))
|
||||
for _, command := range selectedContext.GetCursorCommands() {
|
||||
if command == nil {
|
||||
continue
|
||||
}
|
||||
content := strings.TrimSpace(command.GetContent())
|
||||
if content == "" {
|
||||
continue
|
||||
}
|
||||
name := strings.TrimSpace(command.GetName())
|
||||
if name == "" {
|
||||
entries = append(entries, "<cursor_command>\n"+content+"\n</cursor_command>")
|
||||
continue
|
||||
}
|
||||
entries = append(entries, fmt.Sprintf("<cursor_command name=\"%s\">\n%s\n</cursor_command>", escapePromptXML(name), content))
|
||||
}
|
||||
if len(entries) == 0 {
|
||||
return ""
|
||||
}
|
||||
return "<cursor_commands>\n" + strings.Join(entries, "\n\n") + "\n</cursor_commands>"
|
||||
}
|
||||
|
||||
func buildSelectedIDEStatePromptSection(selectedContext *agentv1.SelectedContext) string {
|
||||
if selectedContext == nil || selectedContext.GetInvocationContext() == nil {
|
||||
return ""
|
||||
|
||||
@@ -524,7 +524,7 @@ func (service *Service) applyProviderModelEvent(stream *ActiveStream, event mode
|
||||
stream.mu.Unlock()
|
||||
}
|
||||
if shouldEmitSyntheticThinking {
|
||||
if err := service.broker.Publish(requestID, StreamEvent{Message: buildThinkingDeltaMessage("Thinking is encrypted. Please wait a moment.", event.ThinkingStyle)}); err != nil {
|
||||
if err := service.broker.Publish(requestID, StreamEvent{Message: buildThinkingDeltaMessage("The reasoning process is encrypted. Please wait a moment. (This message does not affect any functionality; it only indicates the current reasoning status.)", event.ThinkingStyle)}); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
@@ -77,7 +77,7 @@ func (service *Service) maybeCompactBeforeProvider(stream *ActiveStream, convers
|
||||
if service == nil || stream == nil || conversation == nil {
|
||||
return false, nil
|
||||
}
|
||||
manualInstruction, manual := parseManualCompactionDirective(stream.LatestUserText)
|
||||
manualInstruction, manual := streamManualCompactionDirective(stream)
|
||||
plan, err := service.buildCompactionPlan(stream, conversation, compiled, manual, manualInstruction)
|
||||
if err != nil {
|
||||
return false, err
|
||||
@@ -827,7 +827,7 @@ func buildFallbackCompactionSummary(plan *PendingCompaction) string {
|
||||
sections = append(sections, "Compaction note:\n"+truncateCompactionText(plan.HookMessage, 800))
|
||||
}
|
||||
if strings.TrimSpace(plan.ManualInstruction) != "" {
|
||||
sections = append(sections, "Manual compact instruction:\n"+truncateCompactionText(plan.ManualInstruction, 800))
|
||||
sections = append(sections, "Manual summarize instruction:\n"+truncateCompactionText(plan.ManualInstruction, 800))
|
||||
}
|
||||
return strings.TrimSpace(truncateCompactionText(strings.Join(sections, "\n\n"), compactionSummaryMaxChars))
|
||||
}
|
||||
@@ -875,13 +875,62 @@ func (service *Service) resolveCompactionReserveTokens(modelID string) int64 {
|
||||
return compactionAutoReserveTokens
|
||||
}
|
||||
|
||||
func parseManualCompactionRequest(userMessage *agentv1.UserMessage) (string, bool) {
|
||||
if userMessage == nil {
|
||||
return "", false
|
||||
}
|
||||
userText := strings.TrimSpace(userMessage.GetText())
|
||||
if instruction, ok := parseManualCompactionDirective(userText); ok {
|
||||
return instruction, true
|
||||
}
|
||||
if userText != "" {
|
||||
return "", false
|
||||
}
|
||||
selectedContext := userMessage.GetSelectedContext()
|
||||
if selectedContext == nil {
|
||||
return "", false
|
||||
}
|
||||
for _, command := range selectedContext.GetCursorCommands() {
|
||||
if !isCursorSummarizeCommand(command) {
|
||||
continue
|
||||
}
|
||||
instruction, _ := parseManualCompactionDirective(command.GetContent())
|
||||
return instruction, true
|
||||
}
|
||||
return "", false
|
||||
}
|
||||
|
||||
func streamManualCompactionDirective(stream *ActiveStream) (string, bool) {
|
||||
if stream == nil {
|
||||
return "", false
|
||||
}
|
||||
stream.mu.Lock()
|
||||
defer stream.mu.Unlock()
|
||||
if stream.ManualCompaction.Requested {
|
||||
return strings.TrimSpace(stream.ManualCompaction.Instruction), true
|
||||
}
|
||||
return parseManualCompactionDirective(stream.LatestUserText)
|
||||
}
|
||||
|
||||
func isCursorSummarizeCommand(command *agentv1.SelectedCursorCommand) bool {
|
||||
if command == nil {
|
||||
return false
|
||||
}
|
||||
if strings.EqualFold(strings.TrimSpace(command.GetName()), "glass-action-summarize") {
|
||||
return true
|
||||
}
|
||||
_, ok := parseManualCompactionDirective(command.GetContent())
|
||||
return ok
|
||||
}
|
||||
|
||||
func parseManualCompactionDirective(latestUserText string) (string, bool) {
|
||||
trimmed := strings.TrimSpace(latestUserText)
|
||||
const directive = "/summarize"
|
||||
switch {
|
||||
case trimmed == "/compact":
|
||||
case trimmed == directive:
|
||||
return "", true
|
||||
case strings.HasPrefix(trimmed, "/compact "):
|
||||
return strings.TrimSpace(strings.TrimPrefix(trimmed, "/compact")), true
|
||||
case strings.HasPrefix(trimmed, directive+" "):
|
||||
return strings.TrimSpace(strings.TrimPrefix(trimmed, directive)), true
|
||||
default:
|
||||
return "", false
|
||||
}
|
||||
|
||||
@@ -6,6 +6,7 @@ import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
@@ -438,7 +439,11 @@ func (store *ConversationFileStore) writeConversationLocked(conversationID strin
|
||||
if err := store.writeContextLocked(conversationID, conversation); err != nil {
|
||||
return err
|
||||
}
|
||||
return store.writeConversationMetaLocked(conversationID, conversation)
|
||||
if err := store.writeConversationMetaLocked(conversationID, conversation); err != nil {
|
||||
return err
|
||||
}
|
||||
store.syncCursorTranscriptBestEffort(conversationID, conversation)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (store *ConversationFileStore) writeConversationMetaLocked(conversationID string, conversation *ConversationFile) error {
|
||||
@@ -477,6 +482,70 @@ func (store *ConversationFileStore) writeContextLocked(conversationID string, co
|
||||
return writeJSONFileAtomic(store.contextPath(conversationID), context)
|
||||
}
|
||||
|
||||
func (store *ConversationFileStore) syncCursorTranscriptBestEffort(conversationID string, conversation *ConversationFile) {
|
||||
if store == nil || conversation == nil {
|
||||
return
|
||||
}
|
||||
folder := normalizeAgentTranscriptsFolder(conversation.AgentTranscriptsFolder)
|
||||
if folder == "" {
|
||||
return
|
||||
}
|
||||
if err := store.syncCursorTranscript(conversationID, conversation, folder); err != nil {
|
||||
log.Printf("forwarder transcript sync failed conversation_id=%s err=%v", strings.TrimSpace(conversationID), err)
|
||||
}
|
||||
}
|
||||
|
||||
func (store *ConversationFileStore) syncCursorTranscript(conversationID string, conversation *ConversationFile, transcriptsFolder string) error {
|
||||
return store.syncCursorTranscriptWithLatestStatus(conversationID, conversation, transcriptsFolder, false)
|
||||
}
|
||||
|
||||
func (store *ConversationFileStore) syncCursorTranscriptWithLatestStatus(conversationID string, conversation *ConversationFile, transcriptsFolder string, includeLatestStatus bool) error {
|
||||
if store == nil || conversation == nil {
|
||||
return nil
|
||||
}
|
||||
path, err := cursorTranscriptPath(transcriptsFolder, conversationID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
data, err := projectCursorTranscriptJSONLWithLatestStatus(conversation, includeLatestStatus)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if len(data) == 0 {
|
||||
return nil
|
||||
}
|
||||
data = preserveCursorAppendedTurnEnded(path, data)
|
||||
return writeCursorTranscriptAtomic(path, data)
|
||||
}
|
||||
|
||||
func (store *ConversationFileStore) SyncAllCursorTranscriptsBestEffort() {
|
||||
if store == nil {
|
||||
return
|
||||
}
|
||||
conversationIDs, err := store.ListConversationIDs()
|
||||
if err != nil {
|
||||
log.Printf("forwarder transcript backfill scan failed err=%v", err)
|
||||
return
|
||||
}
|
||||
for _, conversationID := range conversationIDs {
|
||||
conversation, err := store.LoadConversation(conversationID)
|
||||
if err != nil {
|
||||
log.Printf("forwarder transcript backfill load failed conversation_id=%s err=%v", conversationID, err)
|
||||
continue
|
||||
}
|
||||
if conversation == nil || conversation.AgentTranscriptsFolder == "" {
|
||||
continue
|
||||
}
|
||||
info, err := os.Stat(conversation.AgentTranscriptsFolder)
|
||||
if err != nil || !info.IsDir() {
|
||||
continue
|
||||
}
|
||||
if err := store.syncCursorTranscriptWithLatestStatus(conversationID, conversation, conversation.AgentTranscriptsFolder, true); err != nil {
|
||||
log.Printf("forwarder transcript backfill failed conversation_id=%s err=%v", conversationID, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func contextVersionForEntries(entries []HistoryEntry) int64 {
|
||||
var version int64
|
||||
for _, entry := range entries {
|
||||
@@ -663,6 +732,9 @@ func mergeConversationMetadata(target *ConversationFile, source *ConversationFil
|
||||
target.ParentConversationID = strings.TrimSpace(source.ParentConversationID)
|
||||
target.ParentToolCallID = strings.TrimSpace(source.ParentToolCallID)
|
||||
target.SubagentTypeName = strings.TrimSpace(source.SubagentTypeName)
|
||||
if folder := normalizeAgentTranscriptsFolder(source.AgentTranscriptsFolder); folder != "" {
|
||||
target.AgentTranscriptsFolder = folder
|
||||
}
|
||||
if strings.TrimSpace(source.Mode) != "" {
|
||||
target.Mode = strings.TrimSpace(source.Mode)
|
||||
}
|
||||
@@ -718,6 +790,10 @@ func normalizeLoadedConversation(conversationID string, conversation *Conversati
|
||||
if conversation.Entries == nil {
|
||||
conversation.Entries = make([]HistoryEntry, 0, 16)
|
||||
}
|
||||
conversation.AgentTranscriptsFolder = normalizeAgentTranscriptsFolder(conversation.AgentTranscriptsFolder)
|
||||
if conversation.AgentTranscriptsFolder == "" {
|
||||
conversation.AgentTranscriptsFolder = agentTranscriptsFolderFromEntries(conversation.Entries)
|
||||
}
|
||||
for _, entry := range conversation.Entries {
|
||||
if entry.Seq >= conversation.NextEntrySeq {
|
||||
conversation.NextEntrySeq = entry.Seq + 1
|
||||
|
||||
@@ -9,6 +9,8 @@ import (
|
||||
modeladapter "cursor/internal/backend/agent/model"
|
||||
)
|
||||
|
||||
const promptContextSourceSelectedCursorCommands = "selected_cursor_commands"
|
||||
|
||||
func newPromptContextMessage(source string, message modeladapter.Message, persist bool) PromptContextMessage {
|
||||
context := PromptContextMessage{
|
||||
Source: strings.TrimSpace(source),
|
||||
|
||||
@@ -18,6 +18,10 @@ const (
|
||||
promptGuardSelectedFileChars = 16000
|
||||
promptGuardSelectedFilesTotalChars = 64000
|
||||
promptGuardSelectedFilesMaxCount = 12
|
||||
promptGuardCursorCommandNameChars = 256
|
||||
promptGuardCursorCommandChars = 12000
|
||||
promptGuardCursorCommandsTotalChars = 32000
|
||||
promptGuardCursorCommandsMaxCount = 8
|
||||
promptGuardRequestFileChars = 16000
|
||||
promptGuardRequestFilesTotalChars = 64000
|
||||
promptGuardRequestFilesMaxCount = 12
|
||||
@@ -106,11 +110,42 @@ func guardSelectedContext(selectedContext *agentv1.SelectedContext) *agentv1.Sel
|
||||
return selectedContext
|
||||
}
|
||||
cloned.Files = guardSelectedFiles(cloned.GetFiles())
|
||||
cloned.CursorCommands = guardSelectedCursorCommands(cloned.GetCursorCommands())
|
||||
cloned.SelectedSkills = guardAgentSkills(cloned.GetSelectedSkills())
|
||||
cloned.ExtraContext = guardStringSlice(cloned.GetExtraContext(), "selected_context.extra_context", promptGuardRealtimeTextChars, promptGuardRealtimeTextChars, promptGuardAgentSkillsMaxCount)
|
||||
return cloned
|
||||
}
|
||||
|
||||
func guardSelectedCursorCommands(commands []*agentv1.SelectedCursorCommand) []*agentv1.SelectedCursorCommand {
|
||||
if len(commands) == 0 {
|
||||
return nil
|
||||
}
|
||||
result := make([]*agentv1.SelectedCursorCommand, 0, minInt(len(commands), promptGuardCursorCommandsMaxCount))
|
||||
remaining := promptGuardCursorCommandsTotalChars
|
||||
for _, command := range commands {
|
||||
if command == nil || len(result) >= promptGuardCursorCommandsMaxCount {
|
||||
continue
|
||||
}
|
||||
content := strings.TrimSpace(command.GetContent())
|
||||
if content == "" {
|
||||
continue
|
||||
}
|
||||
limit := minInt(promptGuardCursorCommandChars, remaining)
|
||||
if limit <= 0 {
|
||||
break
|
||||
}
|
||||
cloned, ok := proto.Clone(command).(*agentv1.SelectedCursorCommand)
|
||||
if !ok || cloned == nil {
|
||||
continue
|
||||
}
|
||||
cloned.Name = truncatePromptGuardText("selected_context.cursor_commands.name", strings.TrimSpace(cloned.GetName()), promptGuardCursorCommandNameChars)
|
||||
cloned.Content = truncatePromptGuardText("selected_context.cursor_commands.content", content, limit)
|
||||
remaining -= promptGuardRuneCount(cloned.GetContent())
|
||||
result = append(result, cloned)
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
func guardSelectedFiles(files []*agentv1.SelectedFile) []*agentv1.SelectedFile {
|
||||
if len(files) == 0 {
|
||||
return nil
|
||||
|
||||
@@ -260,6 +260,9 @@ func applyRunRewindMetadata(conversation *ConversationFile, source *Conversation
|
||||
conversation.ParentConversationID = strings.TrimSpace(source.ParentConversationID)
|
||||
conversation.ParentToolCallID = strings.TrimSpace(source.ParentToolCallID)
|
||||
conversation.SubagentTypeName = strings.TrimSpace(source.SubagentTypeName)
|
||||
if folder := normalizeAgentTranscriptsFolder(source.AgentTranscriptsFolder); folder != "" {
|
||||
conversation.AgentTranscriptsFolder = folder
|
||||
}
|
||||
if strings.TrimSpace(source.Mode) != "" {
|
||||
conversation.Mode = strings.TrimSpace(source.Mode)
|
||||
}
|
||||
|
||||
@@ -23,6 +23,7 @@ import (
|
||||
interactionbridge "cursor/internal/backend/agent/bridge/interaction"
|
||||
runtimecore "cursor/internal/backend/agent/core"
|
||||
modeladapter "cursor/internal/backend/agent/model"
|
||||
promptengine "cursor/internal/backend/agent/prompt"
|
||||
protocol "cursor/internal/backend/agent/protocol"
|
||||
)
|
||||
|
||||
@@ -301,6 +302,7 @@ func NewService(historyRoot string, resolver modeladapter.ChannelResolver) *Serv
|
||||
appendSeq: newAppendSequenceTracker(),
|
||||
}
|
||||
service.startHistoryMaintenance()
|
||||
store.SyncAllCursorTranscriptsBestEffort()
|
||||
return service
|
||||
}
|
||||
|
||||
@@ -672,6 +674,7 @@ func (service *Service) decodeInboundIntent(requestID string, message *agentv1.A
|
||||
default:
|
||||
return InboundIntent{}, fmt.Errorf("unsupported client message kind: %s", clientKind)
|
||||
}
|
||||
intent.ManualCompaction = resolveInboundManualCompaction(message, intent.UserMessage)
|
||||
return intent, nil
|
||||
}
|
||||
|
||||
@@ -689,6 +692,11 @@ func (service *Service) handleRunIntent(intent InboundIntent) error {
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if intent.RequestContext != nil {
|
||||
if folder := normalizeAgentTranscriptsFolder(intent.RequestContext.GetEnv().GetAgentTranscriptsFolder()); folder != "" {
|
||||
conversation.AgentTranscriptsFolder = folder
|
||||
}
|
||||
}
|
||||
rewindDecision := service.decideRunRewind(intent, conversation)
|
||||
if rewindDecision.Evaluated && !rewindDecision.Apply {
|
||||
service.logRunRewindDecision(intent.RequestID, intent.ConversationID, "rewind_skipped", rewindDecision)
|
||||
@@ -751,6 +759,7 @@ func (service *Service) handleRunIntent(intent InboundIntent) error {
|
||||
stream.mu.Lock()
|
||||
stream.ThinkingEffort = strings.TrimSpace(intent.ThinkingEffort)
|
||||
stream.SubagentModelOverrides = cloneSubagentModelOverrides(intent.SubagentModelOverrides)
|
||||
stream.ManualCompaction = intent.ManualCompaction
|
||||
stream.PendingProviderAction = providerActionNone
|
||||
stream.PendingCompaction = nil
|
||||
stream.PendingExecs = make(map[string]runtimecore.PendingExec)
|
||||
@@ -788,6 +797,7 @@ func (service *Service) handleRunIntent(intent InboundIntent) error {
|
||||
"subagent_model_override_count": len(intent.SubagentModelOverrides),
|
||||
"subagent_model_overrides": subagentModelOverrideSummaries(intent.SubagentModelOverrides),
|
||||
"latest_user_text": userMessageText(intent.UserMessage),
|
||||
"manual_compaction_requested": intent.ManualCompaction.Requested,
|
||||
})
|
||||
if err := service.publishCheckpoint(intent.RequestID, intent.ConversationID); err != nil {
|
||||
return err
|
||||
@@ -2332,7 +2342,8 @@ func buildRunEntries(intent InboundIntent, effectiveMode agentv1.AgentMode, turn
|
||||
}
|
||||
}
|
||||
if intent.UserMessage != nil {
|
||||
payload, err := protojson.Marshal(normalizeUserMessageForStorage(intent.UserMessage))
|
||||
normalized := normalizeUserMessageForStorage(intent.UserMessage)
|
||||
payload, err := protojson.Marshal(normalized)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -2343,6 +2354,13 @@ func buildRunEntries(intent InboundIntent, effectiveMode agentv1.AgentMode, turn
|
||||
Kind: "user_message",
|
||||
Payload: payload,
|
||||
})
|
||||
if commandMessage, ok := promptengine.BuildSelectedCursorCommandsReplayMessage(normalized); ok {
|
||||
entries = append(entries, newPromptContextEntry(turnSeq, intent.RequestID, newPromptContextMessage(
|
||||
promptContextSourceSelectedCursorCommands,
|
||||
modeladapter.Message{Role: commandMessage.Role, Content: commandMessage.Content},
|
||||
true,
|
||||
)))
|
||||
}
|
||||
}
|
||||
modeEntry, err := newModeMetadataEntry(turnSeq, intent.RequestID, effectiveMode, intent.HasExplicitMode, intent.ModeSource)
|
||||
if err != nil {
|
||||
@@ -2643,6 +2661,38 @@ func conversationActionIsResume(action *agentv1.ConversationAction) bool {
|
||||
return ok
|
||||
}
|
||||
|
||||
func inboundConversationAction(message *agentv1.AgentClientMessage) *agentv1.ConversationAction {
|
||||
if message == nil {
|
||||
return nil
|
||||
}
|
||||
if action := message.GetConversationAction(); action != nil {
|
||||
return action
|
||||
}
|
||||
if runRequest := message.GetRunRequest(); runRequest != nil {
|
||||
return runRequest.GetAction()
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func conversationActionIsSummarize(action *agentv1.ConversationAction) bool {
|
||||
if action == nil {
|
||||
return false
|
||||
}
|
||||
_, ok := action.GetAction().(*agentv1.ConversationAction_SummarizeAction)
|
||||
return ok
|
||||
}
|
||||
|
||||
func resolveInboundManualCompaction(message *agentv1.AgentClientMessage, userMessage *agentv1.UserMessage) manualCompactionDirective {
|
||||
instruction, requested := parseManualCompactionRequest(userMessage)
|
||||
if conversationActionIsSummarize(inboundConversationAction(message)) {
|
||||
requested = true
|
||||
}
|
||||
return manualCompactionDirective{
|
||||
Requested: requested,
|
||||
Instruction: instruction,
|
||||
}
|
||||
}
|
||||
|
||||
func conversationActionStartsRun(action *agentv1.ConversationAction) bool {
|
||||
if action == nil {
|
||||
return false
|
||||
@@ -2650,6 +2700,7 @@ func conversationActionStartsRun(action *agentv1.ConversationAction) bool {
|
||||
switch action.GetAction().(type) {
|
||||
case *agentv1.ConversationAction_UserMessageAction,
|
||||
*agentv1.ConversationAction_ResumeAction,
|
||||
*agentv1.ConversationAction_SummarizeAction,
|
||||
*agentv1.ConversationAction_StartPlanAction,
|
||||
*agentv1.ConversationAction_ExecutePlanAction:
|
||||
return true
|
||||
|
||||
@@ -0,0 +1,449 @@
|
||||
package forwarder
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"regexp"
|
||||
"strings"
|
||||
|
||||
"google.golang.org/protobuf/encoding/protojson"
|
||||
|
||||
"cursor/gen/agentv1"
|
||||
modeladapter "cursor/internal/backend/agent/model"
|
||||
promptengine "cursor/internal/backend/agent/prompt"
|
||||
)
|
||||
|
||||
var (
|
||||
transcriptContextTagPatterns = compileTranscriptContextTagPatterns([]string{
|
||||
"user_info",
|
||||
"project_layout",
|
||||
"rules",
|
||||
"always_applied_workspace_rules",
|
||||
"agent_requestable_workspace_rules",
|
||||
"user_rules",
|
||||
"agent_skills",
|
||||
"available_skills",
|
||||
"cloud_instructions",
|
||||
"cloud_task_instructions",
|
||||
"open_and_recently_viewed_files",
|
||||
"system_reminder",
|
||||
"system-reminder",
|
||||
"mcp_instructions",
|
||||
"mcp_file_system",
|
||||
"mcp_file_system_servers",
|
||||
"git_status",
|
||||
"agent_transcripts",
|
||||
"cursor_rules_context",
|
||||
"attached_files",
|
||||
"system_notification",
|
||||
"task_notification",
|
||||
"agent_notification",
|
||||
})
|
||||
transcriptThinkingPattern = regexp.MustCompile(`(?is)<(?:think|thinking)>.*?</(?:think|thinking)>`)
|
||||
transcriptBlankLinesPattern = regexp.MustCompile(`\n{3,}`)
|
||||
)
|
||||
|
||||
type cursorTranscriptLine struct {
|
||||
Role string `json:"role,omitempty"`
|
||||
Message *cursorTranscriptMessage `json:"message,omitempty"`
|
||||
Type string `json:"type,omitempty"`
|
||||
Status string `json:"status,omitempty"`
|
||||
Error string `json:"error,omitempty"`
|
||||
}
|
||||
|
||||
type cursorTranscriptMessage struct {
|
||||
Content []cursorTranscriptContent `json:"content"`
|
||||
}
|
||||
|
||||
type cursorTranscriptContent struct {
|
||||
Type string `json:"type"`
|
||||
Text string `json:"text,omitempty"`
|
||||
Name string `json:"name,omitempty"`
|
||||
Input any `json:"input,omitempty"`
|
||||
}
|
||||
|
||||
// projectCursorTranscriptJSONL projects the local semantic history into Cursor's
|
||||
// current agent transcript JSONL contract. context.json remains the source of truth.
|
||||
func projectCursorTranscriptJSONL(conversation *ConversationFile) ([]byte, error) {
|
||||
return projectCursorTranscriptJSONLWithLatestStatus(conversation, false)
|
||||
}
|
||||
|
||||
func projectCursorTranscriptJSONLWithLatestStatus(conversation *ConversationFile, includeLatestStatus bool) ([]byte, error) {
|
||||
if conversation == nil {
|
||||
return nil, nil
|
||||
}
|
||||
lines := make([]cursorTranscriptLine, 0, len(conversation.Entries))
|
||||
maxTurnSeq := int64(0)
|
||||
for _, entry := range conversation.Entries {
|
||||
if entry.TurnSeq > maxTurnSeq {
|
||||
maxTurnSeq = entry.TurnSeq
|
||||
}
|
||||
}
|
||||
currentTurnSeq := int64(0)
|
||||
pendingTurnStatus := cursorTranscriptLine{}
|
||||
flushTurnStatus := func() {
|
||||
if currentTurnSeq > 0 && (includeLatestStatus || currentTurnSeq < maxTurnSeq) && pendingTurnStatus.Type != "" {
|
||||
lines = append(lines, pendingTurnStatus)
|
||||
}
|
||||
pendingTurnStatus = cursorTranscriptLine{}
|
||||
}
|
||||
for _, entry := range conversation.Entries {
|
||||
if entry.TurnSeq > 0 && entry.TurnSeq != currentTurnSeq {
|
||||
flushTurnStatus()
|
||||
currentTurnSeq = entry.TurnSeq
|
||||
}
|
||||
projected, ok, err := projectCursorTranscriptEntry(entry)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if ok {
|
||||
lines = append(lines, projected)
|
||||
}
|
||||
if status, ok := cursorTranscriptTurnStatus(entry); ok {
|
||||
pendingTurnStatus = status
|
||||
}
|
||||
}
|
||||
flushTurnStatus()
|
||||
|
||||
if len(lines) == 0 {
|
||||
return nil, nil
|
||||
}
|
||||
var output bytes.Buffer
|
||||
encoder := json.NewEncoder(&output)
|
||||
encoder.SetEscapeHTML(false)
|
||||
for _, line := range lines {
|
||||
if err := encoder.Encode(line); err != nil {
|
||||
return nil, fmt.Errorf("encode cursor transcript line: %w", err)
|
||||
}
|
||||
}
|
||||
return output.Bytes(), nil
|
||||
}
|
||||
|
||||
func projectCursorTranscriptEntry(entry HistoryEntry) (cursorTranscriptLine, bool, error) {
|
||||
switch strings.TrimSpace(entry.Kind) {
|
||||
case "user_message":
|
||||
message := &agentv1.UserMessage{}
|
||||
if err := protojson.Unmarshal(entry.Payload, message); err != nil {
|
||||
return cursorTranscriptLine{}, false, fmt.Errorf("decode transcript user_message: %w", err)
|
||||
}
|
||||
text := cleanCursorTranscriptUserText(message.GetText())
|
||||
if text == "" {
|
||||
return cursorTranscriptLine{}, false, nil
|
||||
}
|
||||
return cursorTranscriptTextLine("user", text), true, nil
|
||||
case "assistant_text":
|
||||
var payload assistantTextPayload
|
||||
if err := json.Unmarshal(entry.Payload, &payload); err != nil {
|
||||
return cursorTranscriptLine{}, false, fmt.Errorf("decode transcript assistant_text: %w", err)
|
||||
}
|
||||
text := cleanCursorTranscriptAssistantText(payload.Text)
|
||||
thinking := strings.TrimSpace(payload.ReasoningContent)
|
||||
content := joinTranscriptText(text, thinking)
|
||||
if content == "" {
|
||||
return cursorTranscriptLine{}, false, nil
|
||||
}
|
||||
return cursorTranscriptTextLine("assistant", content), true, nil
|
||||
case "tool_call":
|
||||
var payload toolCallEntryPayload
|
||||
if err := json.Unmarshal(entry.Payload, &payload); err != nil {
|
||||
return cursorTranscriptLine{}, false, fmt.Errorf("decode transcript tool_call: %w", err)
|
||||
}
|
||||
toolCall := &agentv1.ToolCall{}
|
||||
if err := protojson.Unmarshal(payload.ToolCall, toolCall); err != nil {
|
||||
return cursorTranscriptLine{}, false, fmt.Errorf("decode transcript tool_call payload: %w", err)
|
||||
}
|
||||
descriptor, ok := promptengine.BuildToolCallReplayDescriptor(firstNonEmpty(payload.ToolCallID, entry.ToolCallID), toolCall)
|
||||
if !ok {
|
||||
return cursorTranscriptLine{}, false, nil
|
||||
}
|
||||
return cursorTranscriptToolCallLine(descriptor.Function.Name, descriptor.Function.Arguments, payload.ReasoningContent), true, nil
|
||||
case "model_message":
|
||||
var payload modelMessageEntryPayload
|
||||
if err := json.Unmarshal(entry.Payload, &payload); err != nil {
|
||||
return cursorTranscriptLine{}, false, fmt.Errorf("decode transcript model_message: %w", err)
|
||||
}
|
||||
return projectCursorTranscriptModelMessage(payload.Message)
|
||||
default:
|
||||
return cursorTranscriptLine{}, false, nil
|
||||
}
|
||||
}
|
||||
|
||||
func projectCursorTranscriptModelMessage(message modeladapter.Message) (cursorTranscriptLine, bool, error) {
|
||||
role := strings.TrimSpace(message.Role)
|
||||
if role == "" || role == "system" || role == "tool" {
|
||||
return cursorTranscriptLine{}, false, nil
|
||||
}
|
||||
content := make([]cursorTranscriptContent, 0, len(message.ToolCalls)+1)
|
||||
texts := make([]string, 0, len(message.ContentParts)+2)
|
||||
if text := strings.TrimSpace(message.Content); text != "" {
|
||||
texts = append(texts, text)
|
||||
}
|
||||
for _, part := range message.ContentParts {
|
||||
switch strings.TrimSpace(strings.ToLower(part.Type)) {
|
||||
case "text", "":
|
||||
if text := strings.TrimSpace(part.Text); text != "" {
|
||||
texts = append(texts, text)
|
||||
}
|
||||
case "image":
|
||||
texts = append(texts, "[Image]")
|
||||
}
|
||||
}
|
||||
if thinking := strings.TrimSpace(message.ReasoningContent); thinking != "" {
|
||||
texts = append(texts, thinking)
|
||||
}
|
||||
if len(texts) > 0 {
|
||||
text := strings.Join(texts, "\n\n")
|
||||
if role == "user" {
|
||||
text = cleanCursorTranscriptUserText(text)
|
||||
} else if role == "assistant" {
|
||||
text = cleanCursorTranscriptAssistantText(text)
|
||||
}
|
||||
if text != "" {
|
||||
content = append(content, cursorTranscriptContent{Type: "text", Text: text})
|
||||
}
|
||||
}
|
||||
for _, call := range message.ToolCalls {
|
||||
name := strings.TrimSpace(call.Function.Name)
|
||||
if name == "" {
|
||||
continue
|
||||
}
|
||||
content = append(content, cursorTranscriptContent{
|
||||
Type: "tool_use",
|
||||
Name: name,
|
||||
Input: decodeTranscriptToolInput(call.Function.Arguments),
|
||||
})
|
||||
}
|
||||
if len(content) == 0 {
|
||||
return cursorTranscriptLine{}, false, nil
|
||||
}
|
||||
return cursorTranscriptLine{Role: role, Message: &cursorTranscriptMessage{Content: content}}, true, nil
|
||||
}
|
||||
|
||||
func cursorTranscriptTextLine(role string, text string) cursorTranscriptLine {
|
||||
if strings.TrimSpace(text) == "" {
|
||||
return cursorTranscriptLine{}
|
||||
}
|
||||
return cursorTranscriptLine{
|
||||
Role: strings.TrimSpace(role),
|
||||
Message: &cursorTranscriptMessage{Content: []cursorTranscriptContent{{
|
||||
Type: "text",
|
||||
Text: text,
|
||||
}}},
|
||||
}
|
||||
}
|
||||
|
||||
func cursorTranscriptToolCallLine(name string, arguments string, reasoning string) cursorTranscriptLine {
|
||||
content := make([]cursorTranscriptContent, 0, 2)
|
||||
if thinking := strings.TrimSpace(reasoning); thinking != "" {
|
||||
content = append(content, cursorTranscriptContent{Type: "text", Text: thinking})
|
||||
}
|
||||
content = append(content, cursorTranscriptContent{
|
||||
Type: "tool_use",
|
||||
Name: strings.TrimSpace(name),
|
||||
Input: decodeTranscriptToolInput(arguments),
|
||||
})
|
||||
return cursorTranscriptLine{Role: "assistant", Message: &cursorTranscriptMessage{Content: content}}
|
||||
}
|
||||
|
||||
func decodeTranscriptToolInput(arguments string) any {
|
||||
trimmed := strings.TrimSpace(arguments)
|
||||
if trimmed == "" {
|
||||
return map[string]any{}
|
||||
}
|
||||
var decoded any
|
||||
if err := json.Unmarshal([]byte(trimmed), &decoded); err == nil {
|
||||
return decoded
|
||||
}
|
||||
return trimmed
|
||||
}
|
||||
|
||||
func cursorTranscriptTurnStatus(entry HistoryEntry) (cursorTranscriptLine, bool) {
|
||||
if strings.TrimSpace(entry.Kind) != "metadata" || entry.TurnSeq <= 0 {
|
||||
return cursorTranscriptLine{}, false
|
||||
}
|
||||
var payload metadataPayload
|
||||
if err := json.Unmarshal(entry.Payload, &payload); err != nil {
|
||||
return cursorTranscriptLine{}, false
|
||||
}
|
||||
switch strings.TrimSpace(payload.Type) {
|
||||
case "turn_completed":
|
||||
return cursorTranscriptLine{Type: "turn_ended", Status: "success"}, true
|
||||
case "provider_error", "failed":
|
||||
return cursorTranscriptLine{
|
||||
Type: "turn_ended",
|
||||
Status: "error",
|
||||
Error: firstNonEmpty(readStringValue(payload.Value["error"]), readStringValue(payload.Value["message"]), "Request failed"),
|
||||
}, true
|
||||
case "control":
|
||||
if strings.TrimSpace(readStringValue(payload.Value["status"])) != "canceled" {
|
||||
return cursorTranscriptLine{}, false
|
||||
}
|
||||
return cursorTranscriptLine{
|
||||
Type: "turn_ended",
|
||||
Status: "aborted",
|
||||
Error: firstNonEmpty(readStringValue(payload.Value["reason"]), readStringValue(payload.Value["message"]), "User aborted request"),
|
||||
}, true
|
||||
default:
|
||||
return cursorTranscriptLine{}, false
|
||||
}
|
||||
}
|
||||
|
||||
func cleanCursorTranscriptUserText(text string) string {
|
||||
return cleanTranscriptContextTags(text)
|
||||
}
|
||||
|
||||
func cleanCursorTranscriptAssistantText(text string) string {
|
||||
cleaned := transcriptThinkingPattern.ReplaceAllString(text, "")
|
||||
return collapseTranscriptBlankLines(cleaned)
|
||||
}
|
||||
|
||||
func cleanTranscriptContextTags(text string) string {
|
||||
cleaned := text
|
||||
for _, pattern := range transcriptContextTagPatterns {
|
||||
cleaned = pattern.ReplaceAllString(cleaned, "")
|
||||
}
|
||||
return collapseTranscriptBlankLines(cleaned)
|
||||
}
|
||||
|
||||
func compileTranscriptContextTagPatterns(tags []string) []*regexp.Regexp {
|
||||
patterns := make([]*regexp.Regexp, 0, len(tags))
|
||||
for _, tag := range tags {
|
||||
patterns = append(patterns, regexp.MustCompile(`(?is)<`+regexp.QuoteMeta(tag)+`(?:\s[^>]*)?>.*?</`+regexp.QuoteMeta(tag)+`>`))
|
||||
}
|
||||
return patterns
|
||||
}
|
||||
|
||||
func collapseTranscriptBlankLines(text string) string {
|
||||
return strings.TrimSpace(transcriptBlankLinesPattern.ReplaceAllString(text, "\n\n"))
|
||||
}
|
||||
|
||||
func joinTranscriptText(text string, thinking string) string {
|
||||
parts := make([]string, 0, 2)
|
||||
if strings.TrimSpace(text) != "" {
|
||||
parts = append(parts, strings.TrimSpace(text))
|
||||
}
|
||||
if strings.TrimSpace(thinking) != "" {
|
||||
parts = append(parts, strings.TrimSpace(thinking))
|
||||
}
|
||||
return strings.Join(parts, "\n\n")
|
||||
}
|
||||
|
||||
func normalizeAgentTranscriptsFolder(path string) string {
|
||||
trimmed := strings.TrimSpace(path)
|
||||
if trimmed == "" || !filepath.IsAbs(trimmed) {
|
||||
return ""
|
||||
}
|
||||
cleaned := filepath.Clean(trimmed)
|
||||
if filepath.Base(cleaned) != "agent-transcripts" {
|
||||
return ""
|
||||
}
|
||||
return cleaned
|
||||
}
|
||||
|
||||
func agentTranscriptsFolderFromEntries(entries []HistoryEntry) string {
|
||||
for _, entry := range entries {
|
||||
if strings.TrimSpace(entry.Kind) != "request_context" {
|
||||
continue
|
||||
}
|
||||
requestContext := &agentv1.RequestContext{}
|
||||
if err := protojson.Unmarshal(entry.Payload, requestContext); err != nil {
|
||||
continue
|
||||
}
|
||||
if folder := normalizeAgentTranscriptsFolder(requestContext.GetEnv().GetAgentTranscriptsFolder()); folder != "" {
|
||||
return folder
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func cursorTranscriptPath(transcriptsFolder string, conversationID string) (string, error) {
|
||||
folder := normalizeAgentTranscriptsFolder(transcriptsFolder)
|
||||
if folder == "" {
|
||||
return "", fmt.Errorf("invalid agent transcripts folder")
|
||||
}
|
||||
id, err := validateConversationID(conversationID)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return filepath.Join(folder, id, id+".jsonl"), nil
|
||||
}
|
||||
|
||||
func preserveCursorAppendedTurnEnded(path string, projected []byte) []byte {
|
||||
existing, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
return projected
|
||||
}
|
||||
lastLine := lastNonEmptyJSONLLine(existing)
|
||||
if len(lastLine) == 0 {
|
||||
return projected
|
||||
}
|
||||
var terminal cursorTranscriptLine
|
||||
if json.Unmarshal(lastLine, &terminal) != nil || terminal.Type != "turn_ended" {
|
||||
return projected
|
||||
}
|
||||
if countTranscriptTurnEnded(existing) <= countTranscriptTurnEnded(projected) {
|
||||
return projected
|
||||
}
|
||||
result := append([]byte(nil), projected...)
|
||||
if len(result) > 0 && result[len(result)-1] != '\n' {
|
||||
result = append(result, '\n')
|
||||
}
|
||||
result = append(result, lastLine...)
|
||||
return append(result, '\n')
|
||||
}
|
||||
|
||||
func lastNonEmptyJSONLLine(data []byte) []byte {
|
||||
lines := bytes.Split(data, []byte{'\n'})
|
||||
for index := len(lines) - 1; index >= 0; index-- {
|
||||
if line := bytes.TrimSpace(lines[index]); len(line) > 0 {
|
||||
return append([]byte(nil), line...)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func countTranscriptTurnEnded(data []byte) int {
|
||||
count := 0
|
||||
for _, line := range bytes.Split(data, []byte{'\n'}) {
|
||||
trimmed := bytes.TrimSpace(line)
|
||||
if len(trimmed) == 0 {
|
||||
continue
|
||||
}
|
||||
var item cursorTranscriptLine
|
||||
if json.Unmarshal(trimmed, &item) == nil && item.Type == "turn_ended" {
|
||||
count++
|
||||
}
|
||||
}
|
||||
return count
|
||||
}
|
||||
|
||||
func writeCursorTranscriptAtomic(path string, data []byte) error {
|
||||
if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil {
|
||||
return fmt.Errorf("create transcript directory: %w", err)
|
||||
}
|
||||
file, tempPath, err := openUniqueArtifactTempFile(path)
|
||||
if err != nil {
|
||||
return fmt.Errorf("open transcript temp file: %w", err)
|
||||
}
|
||||
renamed := false
|
||||
defer func() {
|
||||
if !renamed {
|
||||
_ = os.Remove(tempPath)
|
||||
}
|
||||
}()
|
||||
if _, err := file.Write(data); err != nil {
|
||||
_ = file.Close()
|
||||
return fmt.Errorf("write transcript temp file: %w", err)
|
||||
}
|
||||
if err := file.Close(); err != nil {
|
||||
return fmt.Errorf("close transcript temp file: %w", err)
|
||||
}
|
||||
if err := renameArtifactTempFile(tempPath, path); err != nil {
|
||||
return fmt.Errorf("rename transcript temp file: %w", err)
|
||||
}
|
||||
renamed = true
|
||||
return syncDirectory(filepath.Dir(path))
|
||||
}
|
||||
@@ -0,0 +1,257 @@
|
||||
package forwarder
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"google.golang.org/protobuf/encoding/protojson"
|
||||
|
||||
"cursor/gen/agentv1"
|
||||
)
|
||||
|
||||
func TestProjectCursorTranscriptJSONLMatchesCursorContract(t *testing.T) {
|
||||
toolCall := transcriptTestEditToolCall(t, "file.txt")
|
||||
conversation := transcriptTestConversation([]HistoryEntry{
|
||||
transcriptTestUserMessageEntry(t, 1, "request-1", "<user_info>hidden</user_info>\n\nchange the file"),
|
||||
newAssistantTextEntry(1, "request-1", "<thinking>hidden</thinking>\nDone", "checked carefully", ""),
|
||||
newToolCallEntry(1, "request-1", "call-1", "Edit", "", "", toolCall),
|
||||
newToolResultEntry(1, "request-1", "call-1", "Edit", `{"path":"file.txt"}`, "edited", "", toolCall),
|
||||
newMetadataEntry(1, "request-1", "turn_completed", nil),
|
||||
transcriptTestUserMessageEntry(t, 2, "request-2", "next question"),
|
||||
newMetadataEntry(2, "request-2", "turn_completed", nil),
|
||||
})
|
||||
|
||||
data, err := projectCursorTranscriptJSONL(conversation)
|
||||
if err != nil {
|
||||
t.Fatalf("projectCursorTranscriptJSONL() error = %v", err)
|
||||
}
|
||||
lines := decodeCursorTranscriptLines(t, data)
|
||||
if len(lines) != 5 {
|
||||
t.Fatalf("transcript lines = %d, want 5\n%s", len(lines), data)
|
||||
}
|
||||
|
||||
if lines[0].Role != "user" || transcriptLineText(lines[0]) != "change the file" {
|
||||
t.Fatalf("user line = %#v", lines[0])
|
||||
}
|
||||
if lines[1].Role != "assistant" || transcriptLineText(lines[1]) != "Done\n\nchecked carefully" {
|
||||
t.Fatalf("assistant line = %#v", lines[1])
|
||||
}
|
||||
if lines[2].Role != "assistant" || lines[2].Message == nil || len(lines[2].Message.Content) != 1 {
|
||||
t.Fatalf("tool line = %#v", lines[2])
|
||||
}
|
||||
toolUse := lines[2].Message.Content[0]
|
||||
if toolUse.Type != "tool_use" || toolUse.Name != "Edit" {
|
||||
t.Fatalf("tool use = %#v", toolUse)
|
||||
}
|
||||
input, ok := toolUse.Input.(map[string]any)
|
||||
if !ok || input["path"] != "file.txt" {
|
||||
t.Fatalf("tool input = %#v", toolUse.Input)
|
||||
}
|
||||
if lines[3].Type != "turn_ended" || lines[3].Status != "success" {
|
||||
t.Fatalf("turn status = %#v", lines[3])
|
||||
}
|
||||
if lines[4].Role != "user" || transcriptLineText(lines[4]) != "next question" {
|
||||
t.Fatalf("current user line = %#v", lines[4])
|
||||
}
|
||||
}
|
||||
|
||||
func TestConversationFileStoreSyncsCursorTranscript(t *testing.T) {
|
||||
historyRoot := filepath.Join(t.TempDir(), "history")
|
||||
transcriptsFolder := filepath.Join(t.TempDir(), "agent-transcripts")
|
||||
store := NewConversationFileStore(historyRoot)
|
||||
conversation := transcriptTestConversation(nil)
|
||||
conversation.AgentTranscriptsFolder = transcriptsFolder
|
||||
|
||||
persisted, err := store.SaveConversationWithEntries(conversation.ConversationID, conversation, []HistoryEntry{
|
||||
transcriptTestUserMessageEntry(t, 1, "request-1", "hello"),
|
||||
newAssistantTextEntry(1, "request-1", "hi", "", ""),
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("SaveConversationWithEntries() error = %v", err)
|
||||
}
|
||||
if persisted.AgentTranscriptsFolder != transcriptsFolder {
|
||||
t.Fatalf("persisted transcript folder = %q", persisted.AgentTranscriptsFolder)
|
||||
}
|
||||
|
||||
path := filepath.Join(transcriptsFolder, conversation.ConversationID, conversation.ConversationID+".jsonl")
|
||||
data, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
t.Fatalf("read synced transcript: %v", err)
|
||||
}
|
||||
lines := decodeCursorTranscriptLines(t, data)
|
||||
if len(lines) != 2 || lines[0].Role != "user" || lines[1].Role != "assistant" {
|
||||
t.Fatalf("synced transcript = %s", data)
|
||||
}
|
||||
|
||||
reloaded, err := store.LoadConversation(conversation.ConversationID)
|
||||
if err != nil {
|
||||
t.Fatalf("LoadConversation() error = %v", err)
|
||||
}
|
||||
if reloaded.AgentTranscriptsFolder != transcriptsFolder {
|
||||
t.Fatalf("reloaded transcript folder = %q", reloaded.AgentTranscriptsFolder)
|
||||
}
|
||||
}
|
||||
|
||||
func TestConversationFileStoreBackfillsTranscriptOnStartup(t *testing.T) {
|
||||
historyRoot := filepath.Join(t.TempDir(), "history")
|
||||
transcriptsFolder := filepath.Join(t.TempDir(), "agent-transcripts")
|
||||
if err := os.MkdirAll(transcriptsFolder, 0o755); err != nil {
|
||||
t.Fatalf("create transcript root: %v", err)
|
||||
}
|
||||
store := NewConversationFileStore(historyRoot)
|
||||
conversation := transcriptTestConversation(nil)
|
||||
conversation.AgentTranscriptsFolder = transcriptsFolder
|
||||
_, err := store.SaveConversationWithEntries(conversation.ConversationID, conversation, []HistoryEntry{
|
||||
transcriptTestUserMessageEntry(t, 1, "request-1", "hello"),
|
||||
newAssistantTextEntry(1, "request-1", "hi", "", ""),
|
||||
newMetadataEntry(1, "request-1", "turn_completed", nil),
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("SaveConversationWithEntries() error = %v", err)
|
||||
}
|
||||
path := filepath.Join(transcriptsFolder, conversation.ConversationID, conversation.ConversationID+".jsonl")
|
||||
if err := os.RemoveAll(filepath.Dir(path)); err != nil {
|
||||
t.Fatalf("remove generated transcript: %v", err)
|
||||
}
|
||||
if err := os.MkdirAll(transcriptsFolder, 0o755); err != nil {
|
||||
t.Fatalf("restore transcript root: %v", err)
|
||||
}
|
||||
|
||||
store.SyncAllCursorTranscriptsBestEffort()
|
||||
data, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
t.Fatalf("read backfilled transcript: %v", err)
|
||||
}
|
||||
lines := decodeCursorTranscriptLines(t, data)
|
||||
if len(lines) != 3 || lines[2].Type != "turn_ended" || lines[2].Status != "success" {
|
||||
t.Fatalf("backfilled transcript = %s", data)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNormalizeAgentTranscriptsFolderRejectsUnexpectedPaths(t *testing.T) {
|
||||
root := t.TempDir()
|
||||
if got := normalizeAgentTranscriptsFolder(filepath.Join(root, "agent-transcripts")); got == "" {
|
||||
t.Fatal("valid transcript folder was rejected")
|
||||
}
|
||||
if got := normalizeAgentTranscriptsFolder(filepath.Join(root, "other")); got != "" {
|
||||
t.Fatalf("unexpected folder accepted: %q", got)
|
||||
}
|
||||
if got := normalizeAgentTranscriptsFolder("agent-transcripts"); got != "" {
|
||||
t.Fatalf("relative folder accepted: %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPreserveCursorAppendedTurnEnded(t *testing.T) {
|
||||
path := filepath.Join(t.TempDir(), "conversation.jsonl")
|
||||
existing := []byte("{\"role\":\"user\",\"message\":{\"content\":[{\"type\":\"text\",\"text\":\"hello\"}]}}\n{\"type\":\"turn_ended\",\"status\":\"success\"}\n")
|
||||
if err := os.WriteFile(path, existing, 0o644); err != nil {
|
||||
t.Fatalf("write existing transcript: %v", err)
|
||||
}
|
||||
projected := []byte("{\"role\":\"user\",\"message\":{\"content\":[{\"type\":\"text\",\"text\":\"hello\"}]}}\n{\"role\":\"assistant\",\"message\":{\"content\":[{\"type\":\"text\",\"text\":\"hi\"}]}}\n")
|
||||
preserved := preserveCursorAppendedTurnEnded(path, projected)
|
||||
if countTranscriptTurnEnded(preserved) != 1 {
|
||||
t.Fatalf("preserved transcript = %s", preserved)
|
||||
}
|
||||
if !strings.HasSuffix(string(preserved), "{\"type\":\"turn_ended\",\"status\":\"success\"}\n") {
|
||||
t.Fatalf("terminal line not preserved: %s", preserved)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAgentTranscriptsFolderRecoveredFromLegacyRequestContext(t *testing.T) {
|
||||
folder := filepath.Join(t.TempDir(), "agent-transcripts")
|
||||
payload, err := protojson.Marshal(&agentv1.RequestContext{
|
||||
Env: &agentv1.RequestContextEnv{AgentTranscriptsFolder: folder},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("marshal request context: %v", err)
|
||||
}
|
||||
conversation := transcriptTestConversation([]HistoryEntry{{
|
||||
TurnSeq: 1,
|
||||
Role: "user",
|
||||
Kind: "request_context",
|
||||
Payload: payload,
|
||||
}})
|
||||
conversation.AgentTranscriptsFolder = ""
|
||||
normalizeLoadedConversation(conversation.ConversationID, conversation)
|
||||
if conversation.AgentTranscriptsFolder != folder {
|
||||
t.Fatalf("recovered transcript folder = %q", conversation.AgentTranscriptsFolder)
|
||||
}
|
||||
}
|
||||
|
||||
func transcriptTestConversation(entries []HistoryEntry) *ConversationFile {
|
||||
conversation := &ConversationFile{
|
||||
ConversationID: "conversation-1",
|
||||
RootConversationID: "conversation-1",
|
||||
Mode: "agent",
|
||||
NextTurnSeq: 1,
|
||||
NextEntrySeq: 1,
|
||||
Entries: make([]HistoryEntry, 0, len(entries)),
|
||||
}
|
||||
appendEntriesInPlace(conversation, entries)
|
||||
return conversation
|
||||
}
|
||||
|
||||
func transcriptTestUserMessageEntry(t *testing.T, turnSeq int64, requestID string, text string) HistoryEntry {
|
||||
t.Helper()
|
||||
payload, err := protojson.Marshal(&agentv1.UserMessage{Text: text, MessageId: fmt.Sprintf("message-%d", turnSeq)})
|
||||
if err != nil {
|
||||
t.Fatalf("marshal user message: %v", err)
|
||||
}
|
||||
return HistoryEntry{
|
||||
TurnSeq: turnSeq,
|
||||
RequestID: requestID,
|
||||
Role: "user",
|
||||
Kind: "user_message",
|
||||
Payload: payload,
|
||||
}
|
||||
}
|
||||
|
||||
func transcriptTestEditToolCall(t *testing.T, path string) []byte {
|
||||
t.Helper()
|
||||
payload, err := protojson.Marshal(&agentv1.ToolCall{
|
||||
Tool: &agentv1.ToolCall_EditToolCall{
|
||||
EditToolCall: &agentv1.EditToolCall{
|
||||
Args: &agentv1.EditArgs{Path: path},
|
||||
},
|
||||
},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("marshal edit tool call: %v", err)
|
||||
}
|
||||
return payload
|
||||
}
|
||||
|
||||
func decodeCursorTranscriptLines(t *testing.T, data []byte) []cursorTranscriptLine {
|
||||
t.Helper()
|
||||
lines := make([]cursorTranscriptLine, 0)
|
||||
scanner := bufio.NewScanner(strings.NewReader(string(data)))
|
||||
for scanner.Scan() {
|
||||
var line cursorTranscriptLine
|
||||
if err := json.Unmarshal(scanner.Bytes(), &line); err != nil {
|
||||
t.Fatalf("decode transcript line %q: %v", scanner.Text(), err)
|
||||
}
|
||||
lines = append(lines, line)
|
||||
}
|
||||
if err := scanner.Err(); err != nil {
|
||||
t.Fatalf("scan transcript: %v", err)
|
||||
}
|
||||
return lines
|
||||
}
|
||||
|
||||
func transcriptLineText(line cursorTranscriptLine) string {
|
||||
if line.Message == nil {
|
||||
return ""
|
||||
}
|
||||
texts := make([]string, 0, len(line.Message.Content))
|
||||
for _, content := range line.Message.Content {
|
||||
if content.Type == "text" {
|
||||
texts = append(texts, content.Text)
|
||||
}
|
||||
}
|
||||
return strings.Join(texts, "\n\n")
|
||||
}
|
||||
@@ -22,6 +22,7 @@ type ConversationFile struct {
|
||||
ParentConversationID string `json:"parent_conversation_id"`
|
||||
ParentToolCallID string `json:"parent_tool_call_id"`
|
||||
SubagentTypeName string `json:"subagent_type_name,omitempty"`
|
||||
AgentTranscriptsFolder string `json:"agent_transcripts_folder,omitempty"`
|
||||
Mode string `json:"mode"`
|
||||
ContextVersion int64 `json:"context_version,omitempty"`
|
||||
CurrentLoopID string `json:"current_loop_id,omitempty"`
|
||||
@@ -116,6 +117,11 @@ type StreamSubscriber struct {
|
||||
Signal chan struct{}
|
||||
}
|
||||
|
||||
type manualCompactionDirective struct {
|
||||
Requested bool
|
||||
Instruction string
|
||||
}
|
||||
|
||||
type ActiveStream struct {
|
||||
mu sync.Mutex
|
||||
|
||||
@@ -126,6 +132,7 @@ type ActiveStream struct {
|
||||
ModelName string
|
||||
Mode agentv1.AgentMode
|
||||
LatestUserText string
|
||||
ManualCompaction manualCompactionDirective
|
||||
Status StreamStatus
|
||||
ThinkingEffort string
|
||||
SubagentModelOverrides map[string]runtimecore.SubagentModelOverrideSelection
|
||||
@@ -407,6 +414,7 @@ type InboundIntent struct {
|
||||
HasExplicitMode bool
|
||||
ModeSource ModeSource
|
||||
StartsRun bool
|
||||
ManualCompaction manualCompactionDirective
|
||||
SubagentTypeName string
|
||||
SubagentModelOverrides map[string]runtimecore.SubagentModelOverrideSelection
|
||||
ConversationState *agentv1.ConversationStateStructure
|
||||
|
||||
+121
-220
@@ -27,10 +27,11 @@ const healthPath = "/healthz"
|
||||
const tabServerBaseURL = "https://tab.leokun.cn"
|
||||
|
||||
type Host struct {
|
||||
store *serverconfig.Store
|
||||
listenAddr string
|
||||
configs *serverconfig.Manager
|
||||
healthHTTP *http.Client
|
||||
store *serverconfig.Store
|
||||
listenAddr string
|
||||
configs *serverconfig.Manager
|
||||
healthHTTP *http.Client
|
||||
controlPlaneAuth upstream.AuthorizationProvider
|
||||
|
||||
runMu sync.RWMutex
|
||||
httpServer *http.Server
|
||||
@@ -40,7 +41,7 @@ type Host struct {
|
||||
mux http.Handler
|
||||
}
|
||||
|
||||
func NewHost(store *serverconfig.Store) (*Host, error) {
|
||||
func NewHost(store *serverconfig.Store, controlPlaneAuth upstream.AuthorizationProvider) (*Host, error) {
|
||||
if store == nil {
|
||||
return nil, fmt.Errorf("backend config store is required")
|
||||
}
|
||||
@@ -50,10 +51,11 @@ func NewHost(store *serverconfig.Store) (*Host, error) {
|
||||
}
|
||||
cfg := configs.Current()
|
||||
host := &Host{
|
||||
store: store,
|
||||
listenAddr: cfg.BackendListenAddr,
|
||||
configs: configs,
|
||||
healthHTTP: newLoopbackHTTPClient(),
|
||||
store: store,
|
||||
listenAddr: cfg.BackendListenAddr,
|
||||
configs: configs,
|
||||
healthHTTP: newLoopbackHTTPClient(),
|
||||
controlPlaneAuth: controlPlaneAuth,
|
||||
}
|
||||
if err := host.rebuild(cfg); err != nil {
|
||||
return nil, err
|
||||
@@ -279,7 +281,6 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
server.Use(
|
||||
server.Recover(),
|
||||
server.ServerContext(),
|
||||
server.PolicyMiddleware(host.configs),
|
||||
server.ErrorEncoder(),
|
||||
),
|
||||
server.Mount(ads.RoutePrefix, ads.NewHTTPHandler(appdata.AdsRootPath())),
|
||||
@@ -292,17 +293,11 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
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"),
|
||||
@@ -313,9 +308,6 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
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"),
|
||||
@@ -326,9 +318,6 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
MockProtoType: "aiserver.v1.GetServerConfigResponse",
|
||||
MockBuilder: upstream.ServerConfigMockBuilder,
|
||||
})),
|
||||
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
|
||||
Name: "server_config",
|
||||
})),
|
||||
),
|
||||
server.POST("/aiserver.v1.ServerConfigService/GetServerConfig",
|
||||
server.Name("server_config_service_get_server_config"),
|
||||
@@ -339,9 +328,6 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
MockProtoType: "aiserver.v1.GetServerConfigResponse",
|
||||
MockBuilder: upstream.ServerConfigMockBuilder,
|
||||
})),
|
||||
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
|
||||
Name: "server_config_service_get_server_config",
|
||||
})),
|
||||
),
|
||||
server.POST("/aiserver.v1.AiService/AvailableModels",
|
||||
server.Name("available_models"),
|
||||
@@ -352,9 +338,6 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
MockProtoType: "aiserver.v1.AvailableModelsResponse",
|
||||
MockBuilder: upstream.AvailableModelsMockBuilder,
|
||||
})),
|
||||
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
|
||||
Name: "available_models",
|
||||
})),
|
||||
),
|
||||
server.POST("/aiserver.v1.AiService/GetUsableModels",
|
||||
server.Name("usable_models"),
|
||||
@@ -365,9 +348,6 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
MockProtoType: "aiserver.v1.GetUsableModelsResponse",
|
||||
MockBuilder: upstream.UsableModelsMockBuilder,
|
||||
})),
|
||||
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
|
||||
Name: "usable_models",
|
||||
})),
|
||||
),
|
||||
server.POST("/aiserver.v1.AiService/GetDefaultModelForCli",
|
||||
server.Name("default_model_for_cli"),
|
||||
@@ -378,9 +358,6 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
MockProtoType: "aiserver.v1.GetDefaultModelForCliResponse",
|
||||
MockBuilder: upstream.DefaultModelForCliMockBuilder,
|
||||
})),
|
||||
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
|
||||
Name: "default_model_for_cli",
|
||||
})),
|
||||
),
|
||||
server.POST("/aiserver.v1.AiService/GetDefaultModel",
|
||||
server.Name("default_model"),
|
||||
@@ -391,9 +368,6 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
MockProtoType: "aiserver.v1.GetDefaultModelResponse",
|
||||
MockBuilder: upstream.DefaultModelMockBuilder,
|
||||
})),
|
||||
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
|
||||
Name: "default_model",
|
||||
})),
|
||||
),
|
||||
server.POST("/aiserver.v1.AiService/GetDefaultModelNudgeData",
|
||||
server.Name("default_model_nudge"),
|
||||
@@ -404,9 +378,6 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
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"),
|
||||
@@ -417,9 +388,6 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
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"),
|
||||
@@ -430,9 +398,6 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
MockProtoType: "aiserver.v1.GetFirstWindowStatsigDecisionResponse",
|
||||
MockBuilder: upstream.FirstWindowStatsigDecisionMockBuilder,
|
||||
})),
|
||||
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
|
||||
Name: "first_window_statsig_decision",
|
||||
})),
|
||||
),
|
||||
server.POST("/aiserver.v1.AnalyticsService/SubmitLogs",
|
||||
server.Name("analytics_submit_logs"),
|
||||
@@ -443,9 +408,6 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
MockProtoType: "aiserver.v1.SubmitLogsResponse",
|
||||
MockBuilder: upstream.SubmitLogsMockBuilder,
|
||||
})),
|
||||
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
|
||||
Name: "analytics_submit_logs",
|
||||
})),
|
||||
),
|
||||
server.POST("/aiserver.v1.AnalyticsService/TrackEvents",
|
||||
server.Name("analytics_track_events"),
|
||||
@@ -456,9 +418,6 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
MockProtoType: "aiserver.v1.TrackEventsResponse",
|
||||
MockBuilder: upstream.EmptyMockBuilder,
|
||||
})),
|
||||
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
|
||||
Name: "analytics_track_events",
|
||||
})),
|
||||
),
|
||||
server.POST("/v1/traces",
|
||||
server.Name("otlp_traces"),
|
||||
@@ -467,9 +426,6 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
Name: "otlp_traces",
|
||||
StatusCode: http.StatusOK,
|
||||
})),
|
||||
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
|
||||
Name: "otlp_traces",
|
||||
})),
|
||||
),
|
||||
server.POST("/oauth/token",
|
||||
server.Name("oauth_token"),
|
||||
@@ -478,9 +434,6 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
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"),
|
||||
@@ -489,50 +442,44 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
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),
|
||||
tabServerProcedure("/aiserver.v1.AiService/StreamCpp", "ai_stream_cpp", server.ConnectStream(), routeDeps),
|
||||
tabServerProcedure("/aiserver.v1.AiService/StreamNextCursorPrediction", "ai_stream_next_cursor_prediction", server.ConnectStream(), routeDeps),
|
||||
tabServerProcedure("/aiserver.v1.AiService/GetCppEditClassification", "ai_get_cpp_edit_classification", server.ConnectUnary(), routeDeps),
|
||||
tabServerProcedure("/aiserver.v1.AiService/RefreshTabContext", "ai_refresh_tab_context", server.ConnectUnary(), routeDeps),
|
||||
tabServerProcedure("/aiserver.v1.AiService/CppConfig", "ai_cpp_config", server.ConnectUnary(), routeDeps),
|
||||
tabServerProcedure("/aiserver.v1.AiService/CppEditHistoryStatus", "ai_cpp_edit_history_status", server.ConnectUnary(), routeDeps),
|
||||
tabServerProcedure("/aiserver.v1.AiService/CppAppend", "ai_cpp_append", server.ConnectUnary(), routeDeps),
|
||||
tabServerProcedure("/aiserver.v1.AiService/CppEditHistoryAppend", "ai_cpp_edit_history_append", server.ConnectUnary(), routeDeps),
|
||||
tabServerProcedure("/aiserver.v1.AiService/ReportAiCodeChangeMetrics", "ai_report_ai_code_change_metrics", server.ConnectUnary(), routeDeps),
|
||||
tabServerProcedure("/aiserver.v1.AiService/WriteGitCommitMessage", "ai_write_git_commit_message", server.ConnectUnary(), routeDeps),
|
||||
tabServerProcedure("/aiserver.v1.AiService/WriteGitBranchName", "ai_write_git_branch_name", server.ConnectUnary(), routeDeps),
|
||||
repositoryServiceProcedure(forwarder.RepositoryServiceFastRepoInitHandshakeV2Procedure, "repository_fast_repo_init_handshake_v2", server.ConnectUnary(), agentModule),
|
||||
repositoryServiceProcedure(forwarder.RepositoryServiceFastRepoInitHandshakeProcedure, "repository_fast_repo_init_handshake", server.ConnectUnary(), agentModule),
|
||||
repositoryServiceProcedure(forwarder.RepositoryServiceFastRepoSyncCompleteProcedure, "repository_fast_repo_sync_complete", server.ConnectUnary(), agentModule),
|
||||
repositoryServiceProcedure(forwarder.RepositoryServiceSyncMerkleSubtreeV2Procedure, "repository_sync_merkle_subtree_v2", server.ConnectUnary(), agentModule),
|
||||
repositoryServiceProcedure(forwarder.RepositoryServiceSyncMerkleSubtreeProcedure, "repository_sync_merkle_subtree", server.ConnectUnary(), agentModule),
|
||||
repositoryServiceProcedure(forwarder.RepositoryServiceFastUpdateFileV2Procedure, "repository_fast_update_file_v2", server.ConnectUnary(), agentModule),
|
||||
repositoryServiceProcedure(forwarder.RepositoryServiceFastUpdateFileProcedure, "repository_fast_update_file", server.ConnectUnary(), agentModule),
|
||||
repositoryServiceProcedure(forwarder.RepositoryServiceEnsureIndexCreatedProcedure, "repository_ensure_index_created", server.ConnectUnary(), agentModule),
|
||||
repositoryServiceProcedure(forwarder.RepositoryServiceGetCopyStatusProcedure, "repository_get_copy_status", server.ConnectUnary(), agentModule),
|
||||
repositoryServiceProcedure(forwarder.RepositoryServiceGetUploadLimitsProcedure, "repository_get_upload_limits", server.ConnectUnary(), agentModule),
|
||||
repositoryServiceProcedure(forwarder.RepositoryServiceGetNumFilesToSendProcedure, "repository_get_num_files_to_send", server.ConnectUnary(), agentModule),
|
||||
repositoryServiceProcedure(forwarder.RepositoryServiceGetAvailableChunkingStrategiesProcedure, "repository_get_available_chunking_strategies", server.ConnectUnary(), agentModule),
|
||||
repositoryServiceProcedure(forwarder.RepositoryServiceGetHighLevelFolderDescriptionProcedure, "repository_get_high_level_folder_description", server.ConnectUnary(), agentModule),
|
||||
repositoryServiceProcedure(forwarder.RepositoryServiceRepositoryStatusProcedure, "repository_status", server.ConnectUnary(), agentModule),
|
||||
repositoryServiceProcedure(forwarder.RepositoryServiceBatchRepositoryStatusProcedure, "repository_batch_status", server.ConnectUnary(), agentModule),
|
||||
uploadServiceProcedure(forwarder.UploadServiceUploadDocumentationProcedure, "upload_documentation", server.ConnectUnary(), agentModule),
|
||||
uploadServiceProcedure(forwarder.UploadServiceGetDocProcedure, "upload_get_doc", server.ConnectUnary(), agentModule),
|
||||
uploadServiceProcedure(forwarder.UploadServiceGetPagesProcedure, "upload_get_pages", server.ConnectUnary(), agentModule),
|
||||
uploadServiceProcedure(forwarder.UploadServiceUploadedStatusProcedure, "upload_uploaded_status", server.ConnectUnary(), agentModule),
|
||||
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),
|
||||
tabServerProcedure("/aiserver.v1.CppService/AvailableModels", "cpp_available_models", server.ConnectUnary(), routeDeps),
|
||||
tabServerProcedure("/aiserver.v1.CppService/RecordCppFate", "cpp_record_cpp_fate", server.ConnectUnary(), routeDeps),
|
||||
server.Any("/aiserver.v1.CppService/*",
|
||||
server.Name("cpp_service"),
|
||||
server.HTTP(),
|
||||
@@ -540,14 +487,11 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) 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),
|
||||
tabServerProcedure("/aiserver.v1.FileSyncService/FSSyncFile", "file_sync_sync_file", server.ConnectUnary(), routeDeps),
|
||||
tabServerProcedure("/aiserver.v1.FileSyncService/FSIsEnabledForUser", "file_sync_is_enabled_for_user", server.ConnectUnary(), routeDeps),
|
||||
tabServerProcedure("/aiserver.v1.FileSyncService/FSConfig", "file_sync_config", server.ConnectUnary(), routeDeps),
|
||||
tabServerProcedure("/aiserver.v1.FileSyncService/FSUploadFile", "file_sync_upload_file", server.ConnectUnary(), routeDeps),
|
||||
server.Any("/aiserver.v1.FileSyncService/*",
|
||||
server.Name("file_sync"),
|
||||
server.HTTP(),
|
||||
@@ -555,25 +499,16 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) 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"),
|
||||
@@ -584,9 +519,6 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
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"),
|
||||
@@ -597,22 +529,21 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
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.Local(cursorControlPlaneAction(
|
||||
host.controlPlaneAuth,
|
||||
routeDeps,
|
||||
"dashboard_get_managed_skills",
|
||||
upstream.MockProtoAction(routeDeps, upstream.CompatRouteConfig{
|
||||
Name: "dashboard_get_managed_skills",
|
||||
StatusCode: http.StatusOK,
|
||||
MockProtoType: "aiserver.v1.GetManagedSkillsResponse",
|
||||
MockBuilder: upstream.DashboardManagedSkillsMockBuilder,
|
||||
}),
|
||||
)),
|
||||
),
|
||||
server.POST("/aiserver.v1.DashboardService/GetTeamAdminSettingsOrEmptyIfNotInTeam",
|
||||
server.Name("dashboard_get_team_admin_settings_or_empty"),
|
||||
@@ -623,9 +554,6 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
MockProtoType: "aiserver.v1.GetTeamAdminSettingsResponse",
|
||||
MockBuilder: upstream.EmptyMockBuilder,
|
||||
})),
|
||||
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
|
||||
Name: "dashboard_get_team_admin_settings_or_empty",
|
||||
})),
|
||||
),
|
||||
server.POST("/aiserver.v1.DashboardService/GetTeamReposOrEmptyIfNotInTeam",
|
||||
server.Name("dashboard_get_team_repos_or_empty"),
|
||||
@@ -636,9 +564,6 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
MockProtoType: "aiserver.v1.GetTeamReposResponse",
|
||||
MockBuilder: upstream.EmptyMockBuilder,
|
||||
})),
|
||||
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
|
||||
Name: "dashboard_get_team_repos_or_empty",
|
||||
})),
|
||||
),
|
||||
server.POST("/aiserver.v1.DashboardService/ListMarketplaces",
|
||||
server.Name("dashboard_list_marketplaces"),
|
||||
@@ -649,9 +574,6 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
MockProtoType: "aiserver.v1.ListMarketplacesResponse",
|
||||
MockBuilder: upstream.EmptyMockBuilder,
|
||||
})),
|
||||
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
|
||||
Name: "dashboard_list_marketplaces",
|
||||
})),
|
||||
),
|
||||
server.POST("/aiserver.v1.DashboardService/GetGlobalCommands",
|
||||
server.Name("dashboard_get_global_commands"),
|
||||
@@ -662,9 +584,6 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
MockProtoType: "aiserver.v1.GetGlobalCommandsResponse",
|
||||
MockBuilder: upstream.EmptyMockBuilder,
|
||||
})),
|
||||
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
|
||||
Name: "dashboard_get_global_commands",
|
||||
})),
|
||||
),
|
||||
server.POST("/aiserver.v1.DashboardService/GetEffectiveUserPlugins",
|
||||
server.Name("dashboard_get_effective_user_plugins"),
|
||||
@@ -675,9 +594,6 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
MockProtoType: "aiserver.v1.GetEffectiveUserPluginsResponse",
|
||||
MockBuilder: upstream.EmptyMockBuilder,
|
||||
})),
|
||||
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
|
||||
Name: "dashboard_get_effective_user_plugins",
|
||||
})),
|
||||
),
|
||||
server.POST("/aiserver.v1.DashboardService/RegisterMarketplaceAndPlugins",
|
||||
server.Name("dashboard_register_marketplace_and_plugins"),
|
||||
@@ -688,9 +604,6 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
MockProtoType: "aiserver.v1.RegisterMarketplaceAndPluginsResponse",
|
||||
MockBuilder: upstream.EmptyMockBuilder,
|
||||
})),
|
||||
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
|
||||
Name: "dashboard_register_marketplace_and_plugins",
|
||||
})),
|
||||
),
|
||||
server.POST("/aiserver.v1.DashboardService/GetCliDownloadUrl",
|
||||
server.Name("dashboard_get_cli_download_url"),
|
||||
@@ -701,9 +614,6 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
MockProtoType: "aiserver.v1.GetCliDownloadUrlResponse",
|
||||
MockBuilder: upstream.EmptyMockBuilder,
|
||||
})),
|
||||
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
|
||||
Name: "dashboard_get_cli_download_url",
|
||||
})),
|
||||
),
|
||||
server.POST("/aiserver.v1.DashboardService/GetMe",
|
||||
server.Name("dashboard_get_me"),
|
||||
@@ -714,9 +624,6 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
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"),
|
||||
@@ -727,9 +634,6 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
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"),
|
||||
@@ -740,9 +644,6 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
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"),
|
||||
@@ -753,9 +654,6 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
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"),
|
||||
@@ -766,11 +664,25 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
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),
|
||||
cursorControlPlaneProcedure("/aiserver.v1.DashboardService/AddMarketplace", "dashboard_add_marketplace", server.ConnectUnary(), host.controlPlaneAuth, routeDeps),
|
||||
cursorControlPlaneProcedure("/aiserver.v1.DashboardService/AddMcpServersFromPlugin", "dashboard_add_mcp_servers_from_plugin", server.ConnectUnary(), host.controlPlaneAuth, routeDeps),
|
||||
cursorControlPlaneProcedure("/aiserver.v1.DashboardService/BatchGetPluginMcpConfig", "dashboard_batch_get_plugin_mcp_config", server.ConnectUnary(), host.controlPlaneAuth, routeDeps),
|
||||
cursorControlPlaneProcedure("/aiserver.v1.DashboardService/GetAvailableMcpServers", "dashboard_get_available_mcp_servers", server.ConnectUnary(), host.controlPlaneAuth, routeDeps),
|
||||
cursorControlPlaneProcedure("/aiserver.v1.DashboardService/GetEffectiveUserPlugins", "dashboard_get_effective_user_plugins", server.ConnectUnary(), host.controlPlaneAuth, routeDeps),
|
||||
cursorControlPlaneProcedure("/aiserver.v1.DashboardService/GetPlugin", "dashboard_get_plugin", server.ConnectUnary(), host.controlPlaneAuth, routeDeps),
|
||||
cursorControlPlaneProcedure("/aiserver.v1.DashboardService/GetPluginMcpConfig", "dashboard_get_plugin_mcp_config", server.ConnectUnary(), host.controlPlaneAuth, routeDeps),
|
||||
cursorControlPlaneProcedure("/aiserver.v1.DashboardService/InstallUserPlugin", "dashboard_install_user_plugin", server.ConnectUnary(), host.controlPlaneAuth, routeDeps),
|
||||
cursorControlPlaneProcedure("/aiserver.v1.DashboardService/ListMarketplacePlugins", "dashboard_list_marketplace_plugins", server.ConnectUnary(), host.controlPlaneAuth, routeDeps),
|
||||
cursorControlPlaneProcedure("/aiserver.v1.DashboardService/ListMarketplaces", "dashboard_list_marketplaces", server.ConnectUnary(), host.controlPlaneAuth, routeDeps),
|
||||
cursorControlPlaneProcedure("/aiserver.v1.DashboardService/ListUserPluginInstalls", "dashboard_list_user_plugin_installs", server.ConnectUnary(), host.controlPlaneAuth, routeDeps),
|
||||
cursorControlPlaneProcedure("/aiserver.v1.DashboardService/RefreshMarketplace", "dashboard_refresh_marketplace", server.ConnectUnary(), host.controlPlaneAuth, routeDeps),
|
||||
cursorControlPlaneProcedure("/aiserver.v1.DashboardService/RegisterMarketplaceAndPlugins", "dashboard_register_marketplace_and_plugins", server.ConnectUnary(), host.controlPlaneAuth, routeDeps),
|
||||
cursorControlPlaneProcedure("/aiserver.v1.DashboardService/RemoveMarketplace", "dashboard_remove_marketplace", server.ConnectUnary(), host.controlPlaneAuth, routeDeps),
|
||||
cursorControlPlaneProcedure("/aiserver.v1.DashboardService/ResolvePluginsByRef", "dashboard_resolve_plugins_by_ref", server.ConnectUnary(), host.controlPlaneAuth, routeDeps),
|
||||
cursorControlPlaneProcedure("/aiserver.v1.DashboardService/UninstallUserPlugin", "dashboard_uninstall_user_plugin", server.ConnectUnary(), host.controlPlaneAuth, routeDeps),
|
||||
cursorControlPlaneProcedure("/aiserver.v1.DashboardService/UpdateUserPluginInstall", "dashboard_update_user_plugin_install", server.ConnectUnary(), host.controlPlaneAuth, routeDeps),
|
||||
cursorControlPlaneProcedure("/aiserver.v1.MCPRegistryService/GetKnownServers", "mcp_registry_get_known_servers", server.ConnectUnary(), host.controlPlaneAuth, routeDeps),
|
||||
server.Any("/aiserver.v1.DashboardService/*",
|
||||
server.Name("dashboard"),
|
||||
server.HTTP(),
|
||||
@@ -778,9 +690,6 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) 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"),
|
||||
@@ -789,9 +698,6 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) 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"),
|
||||
@@ -800,9 +706,6 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) 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"),
|
||||
@@ -811,9 +714,6 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
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"),
|
||||
@@ -822,9 +722,6 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
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"),
|
||||
@@ -836,9 +733,6 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
"hasValidPaymentMethod": true,
|
||||
},
|
||||
})),
|
||||
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
|
||||
Name: "auth_has_valid_payment_method",
|
||||
})),
|
||||
),
|
||||
server.Any("/auth/poll",
|
||||
server.Name("auth_poll"),
|
||||
@@ -847,9 +741,6 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
Name: "auth_poll",
|
||||
StatusCode: http.StatusOK,
|
||||
})),
|
||||
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
|
||||
Name: "auth_poll",
|
||||
})),
|
||||
),
|
||||
server.POST("/auth/logout",
|
||||
server.Name("auth_logout"),
|
||||
@@ -858,9 +749,6 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
Name: "auth_logout",
|
||||
StatusCode: http.StatusNoContent,
|
||||
})),
|
||||
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
|
||||
Name: "auth_logout",
|
||||
})),
|
||||
),
|
||||
server.Any("/auth/*",
|
||||
server.Name("auth_proxy"),
|
||||
@@ -869,58 +757,32 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) 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 {
|
||||
func repositoryServiceProcedure(pattern string, name string, protocol server.RouteOption, module *forwarder.Module) 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 {
|
||||
func uploadServiceProcedure(pattern string, name string, protocol server.RouteOption, module *forwarder.Module) 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})
|
||||
func tabServerProcedure(pattern string, name string, protocol server.RouteOption, deps upstream.Dependencies) server.Option {
|
||||
forward := upstream.ForwardAction(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)
|
||||
@@ -932,16 +794,55 @@ func tabServerUpstreamProcedure(pattern string, name string, protocol server.Rou
|
||||
targetURL.Host = baseURL.Host
|
||||
ctx.UpstreamURL = &targetURL
|
||||
}
|
||||
return direct(ctx)
|
||||
return forward(ctx)
|
||||
}
|
||||
return server.POST(pattern,
|
||||
server.Name(name),
|
||||
protocol,
|
||||
server.Local(action),
|
||||
server.Upstream(action),
|
||||
)
|
||||
}
|
||||
|
||||
func cursorControlPlaneProcedure(
|
||||
pattern string,
|
||||
name string,
|
||||
protocol server.RouteOption,
|
||||
authorizationProvider upstream.AuthorizationProvider,
|
||||
deps upstream.Dependencies,
|
||||
) server.Option {
|
||||
notFound := func(ctx *server.Context) error {
|
||||
http.NotFound(ctx.Writer, ctx.Request)
|
||||
return nil
|
||||
}
|
||||
return server.POST(pattern,
|
||||
server.Name(name),
|
||||
protocol,
|
||||
server.Local(cursorControlPlaneAction(authorizationProvider, deps, name, notFound)),
|
||||
)
|
||||
}
|
||||
|
||||
func cursorControlPlaneAction(
|
||||
authorizationProvider upstream.AuthorizationProvider,
|
||||
deps upstream.Dependencies,
|
||||
name string,
|
||||
fallback server.HandlerFunc,
|
||||
) server.HandlerFunc {
|
||||
forward := upstream.AuthenticatedForwardAction(deps, upstream.CompatRouteConfig{Name: name}, authorizationProvider)
|
||||
return func(ctx *server.Context) error {
|
||||
if authorizationProvider == nil || !authorizationProvider.SignedIn() {
|
||||
return fallback(ctx)
|
||||
}
|
||||
if ctx == nil || ctx.Request == nil || ctx.Request.URL == nil {
|
||||
return fmt.Errorf("Cursor 控制面请求上下文无效")
|
||||
}
|
||||
targetURL := *ctx.Request.URL
|
||||
targetURL.Scheme = "https"
|
||||
targetURL.Host = "api2.cursor.sh:443"
|
||||
ctx.UpstreamURL = &targetURL
|
||||
return forward(ctx)
|
||||
}
|
||||
}
|
||||
|
||||
type serverSystemSettings struct {
|
||||
configs *serverconfig.Manager
|
||||
}
|
||||
|
||||
@@ -171,20 +171,6 @@ func (manager *Manager) LegacyRuntimeSnapshot(_ context.Context) (legacyruntime.
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (manager *Manager) RouteMode(hasUpstreamURL bool) string {
|
||||
if !hasUpstreamURL {
|
||||
return DefaultRoutingMode
|
||||
}
|
||||
if manager == nil {
|
||||
return DefaultRoutingMode
|
||||
}
|
||||
mode := normalizeRoutingMode(manager.Current().Routing.Mode)
|
||||
if mode == "" {
|
||||
return DefaultRoutingMode
|
||||
}
|
||||
return mode
|
||||
}
|
||||
|
||||
func (manager *Manager) setCurrent(cfg Config) {
|
||||
next := cfg
|
||||
manager.current.Store(&next)
|
||||
|
||||
@@ -136,6 +136,9 @@ func (store *Store) saveLocked(normalized Config) error {
|
||||
}
|
||||
|
||||
func shouldPersistNormalizedConfig(raw []byte, current Config, normalized Config) bool {
|
||||
if yamlHasKey(raw, "routing") {
|
||||
return true
|
||||
}
|
||||
if !yamlHasKey(raw, "backendListenAddr") || !yamlHasKey(raw, "proxyListenAddr") {
|
||||
return true
|
||||
}
|
||||
|
||||
@@ -15,7 +15,6 @@ const (
|
||||
DefaultBackendListenAddr = "127.0.0.1:18090"
|
||||
DefaultProxyListenAddr = "127.0.0.1:18080"
|
||||
DefaultFrontendBaseURL = "http://127.0.0.1"
|
||||
DefaultRoutingMode = "local"
|
||||
DefaultProviderStreamIdleTimeoutSeconds = 240
|
||||
MinProviderStreamIdleTimeoutSeconds = 30
|
||||
)
|
||||
@@ -43,10 +42,6 @@ type ModelAdapterConfig struct {
|
||||
ThinkingBudgetTokens int `json:"thinkingBudgetTokens" yaml:"thinkingBudgetTokens"`
|
||||
}
|
||||
|
||||
type RoutingConfig struct {
|
||||
Mode string `json:"mode" yaml:"mode"`
|
||||
}
|
||||
|
||||
type HomeMetricsConfig struct {
|
||||
IncludeCacheWriteInHitRate bool `json:"includeCacheWriteInHitRate" yaml:"includeCacheWriteInHitRate"`
|
||||
}
|
||||
@@ -57,7 +52,6 @@ type Config struct {
|
||||
BackendListenAddr string `json:"backendListenAddr" yaml:"backendListenAddr"`
|
||||
ProxyListenAddr string `json:"proxyListenAddr" yaml:"proxyListenAddr"`
|
||||
ModelAdapters []ModelAdapterConfig `json:"modelAdapters" yaml:"modelAdapters"`
|
||||
Routing RoutingConfig `json:"routing" yaml:"routing"`
|
||||
HomeMetrics HomeMetricsConfig `json:"homeMetrics" yaml:"homeMetrics"`
|
||||
LastAgentModelHash string `json:"lastAgentModelHash" yaml:"lastAgentModelHash"`
|
||||
}
|
||||
@@ -69,9 +63,6 @@ func DefaultConfig() Config {
|
||||
BackendListenAddr: DefaultBackendListenAddr,
|
||||
ProxyListenAddr: DefaultProxyListenAddr,
|
||||
ModelAdapters: []ModelAdapterConfig{},
|
||||
Routing: RoutingConfig{
|
||||
Mode: DefaultRoutingMode,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
@@ -91,10 +82,6 @@ func NormalizeConfig(input Config) (Config, error) {
|
||||
output.ProxyListenAddr = proxyListenAddr
|
||||
output.HomeMetrics.IncludeCacheWriteInHitRate = input.HomeMetrics.IncludeCacheWriteInHitRate
|
||||
output.LastAgentModelHash = strings.TrimSpace(input.LastAgentModelHash)
|
||||
output.Routing.Mode = normalizeRoutingMode(input.Routing.Mode)
|
||||
if output.Routing.Mode == "" {
|
||||
output.Routing.Mode = DefaultRoutingMode
|
||||
}
|
||||
adapters, err := NormalizeModelAdapterConfigs(input.ModelAdapters)
|
||||
if err != nil {
|
||||
return Config{}, err
|
||||
@@ -280,14 +267,3 @@ func normalizeModelAdapterType(value string) string {
|
||||
return ""
|
||||
}
|
||||
}
|
||||
|
||||
func normalizeRoutingMode(value string) string {
|
||||
switch strings.ToLower(strings.TrimSpace(value)) {
|
||||
case "", "local":
|
||||
return "local"
|
||||
case "upstream":
|
||||
return "upstream"
|
||||
default:
|
||||
return ""
|
||||
}
|
||||
}
|
||||
|
||||
@@ -34,7 +34,6 @@ type Context struct {
|
||||
StartedAt time.Time
|
||||
|
||||
UpstreamURL *url.URL
|
||||
Mode ExecutionMode
|
||||
LastError error
|
||||
|
||||
Logger *slog.Logger
|
||||
@@ -48,7 +47,6 @@ func newContext(writer http.ResponseWriter, request *http.Request, route Route)
|
||||
Protocol: route.Protocol,
|
||||
StartedAt: time.Now(),
|
||||
Logger: slog.Default(),
|
||||
Mode: ModeLocal,
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -1,14 +1,12 @@
|
||||
package server
|
||||
|
||||
import (
|
||||
"cursor/internal/logger"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"runtime/debug"
|
||||
"strings"
|
||||
|
||||
serverconfig "cursor/internal/backend/server/config"
|
||||
legacyruntime "cursor/internal/runtime"
|
||||
)
|
||||
|
||||
@@ -39,20 +37,6 @@ func ServerContext() Middleware {
|
||||
}
|
||||
}
|
||||
|
||||
func PolicyMiddleware(configs *serverconfig.Manager) Middleware {
|
||||
return func(next HandlerFunc) HandlerFunc {
|
||||
return func(ctx *Context) error {
|
||||
ctx.Mode = parseExecutionMode(configs.RouteMode(ctx.UpstreamURL != nil))
|
||||
path := ""
|
||||
if ctx.Request != nil && ctx.Request.URL != nil {
|
||||
path = ctx.Request.URL.Path
|
||||
}
|
||||
logger.Infof("ctx.Mode=%s upstream=%t path=%s", ctx.Mode, ctx.UpstreamURL != nil, path)
|
||||
return next(ctx)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func ErrorEncoder() Middleware {
|
||||
return func(next HandlerFunc) HandlerFunc {
|
||||
return func(ctx *Context) error {
|
||||
|
||||
@@ -1,19 +0,0 @@
|
||||
package server
|
||||
|
||||
type ExecutionMode string
|
||||
|
||||
const (
|
||||
// ModeLocal 表示本地模式,适用于直接处理请求的情况。
|
||||
ModeLocal ExecutionMode = "local"
|
||||
// ModeUpstream 表示直连上游模式,适用于将请求转发到原始地址。
|
||||
ModeUpstream ExecutionMode = "upstream"
|
||||
)
|
||||
|
||||
func parseExecutionMode(value string) ExecutionMode {
|
||||
switch value {
|
||||
case string(ModeUpstream):
|
||||
return ModeUpstream
|
||||
default:
|
||||
return ModeLocal
|
||||
}
|
||||
}
|
||||
@@ -17,7 +17,6 @@ type Route struct {
|
||||
Protocol ProtocolClass
|
||||
Middleware []Middleware
|
||||
Local HandlerFunc
|
||||
Upstream HandlerFunc
|
||||
}
|
||||
|
||||
type App struct {
|
||||
@@ -129,12 +128,6 @@ func Local(action HandlerFunc) RouteOption {
|
||||
}
|
||||
}
|
||||
|
||||
func Upstream(action HandlerFunc) RouteOption {
|
||||
return func(route *Route) {
|
||||
route.Upstream = action
|
||||
}
|
||||
}
|
||||
|
||||
func (app *App) registerRoute(route Route) {
|
||||
handler := app.buildRouteHandler(route)
|
||||
if route.Method == "" {
|
||||
@@ -148,12 +141,6 @@ func (app *App) buildRouteHandler(route Route) http.HandlerFunc {
|
||||
chain := append([]Middleware{}, app.globalMiddlewares...)
|
||||
chain = append(chain, route.Middleware...)
|
||||
final := Chain(chain...)(func(ctx *Context) error {
|
||||
if shouldUseUpstreamAction(ctx, route) && route.Upstream != nil {
|
||||
return route.Upstream(ctx)
|
||||
}
|
||||
if shouldUseUpstreamAction(ctx, route) && ctx.UpstreamURL != nil {
|
||||
return fmt.Errorf("route %s is missing upstream action while request targets upstream %s", route.Name, ctx.UpstreamURL.String())
|
||||
}
|
||||
if route.Local != nil {
|
||||
return route.Local(ctx)
|
||||
}
|
||||
@@ -168,14 +155,6 @@ func (app *App) buildRouteHandler(route Route) http.HandlerFunc {
|
||||
}
|
||||
}
|
||||
|
||||
func shouldUseUpstreamAction(ctx *Context, route Route) bool {
|
||||
_ = route
|
||||
if ctx == nil {
|
||||
return false
|
||||
}
|
||||
return ctx.Mode == ModeUpstream
|
||||
}
|
||||
|
||||
func Chain(middlewares ...Middleware) Middleware {
|
||||
return func(final HandlerFunc) HandlerFunc {
|
||||
wrapped := final
|
||||
|
||||
@@ -2,6 +2,7 @@ package upstream
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"strings"
|
||||
@@ -19,7 +20,7 @@ type CompatRouteConfig struct {
|
||||
ConsoleLog bool
|
||||
}
|
||||
|
||||
func DirectAction(deps Dependencies, cfg CompatRouteConfig) server.HandlerFunc {
|
||||
func ForwardAction(deps Dependencies, cfg CompatRouteConfig) server.HandlerFunc {
|
||||
return func(ctx *server.Context) error {
|
||||
reqCtx, route, err := newCompatRouteObjects(ctx, deps, cfg)
|
||||
if err != nil {
|
||||
@@ -29,6 +30,34 @@ func DirectAction(deps Dependencies, cfg CompatRouteConfig) server.HandlerFunc {
|
||||
}
|
||||
}
|
||||
|
||||
// AuthenticatedForwardAction forwards a Cursor control-plane request with the
|
||||
// independent desktop account after the local-mode identity rewrite has run.
|
||||
func AuthenticatedForwardAction(deps Dependencies, cfg CompatRouteConfig, authorizationProvider AuthorizationProvider) server.HandlerFunc {
|
||||
return func(ctx *server.Context) error {
|
||||
reqCtx, _, err := newCompatRouteObjects(ctx, deps, cfg)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if reqCtx == nil || reqCtx.Request == nil {
|
||||
return fmt.Errorf("Cursor 控制面请求上下文无效")
|
||||
}
|
||||
if authorizationProvider == nil {
|
||||
return fmt.Errorf("Cursor 账号服务未初始化")
|
||||
}
|
||||
authorization, err := authorizationProvider.Authorization(reqCtx.Request.Context())
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
_, err = ForwardToUpstream(reqCtx, ForwardOptions{
|
||||
PatchHeaders: func(headers http.Header) {
|
||||
headers.Set("Authorization", authorization)
|
||||
headers.Set("x-cursor-checksum", BuildCursorChecksum(authorization))
|
||||
},
|
||||
})
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
func FixedStatusAction(deps Dependencies, cfg CompatRouteConfig) server.HandlerFunc {
|
||||
return func(ctx *server.Context) error {
|
||||
reqCtx, route, err := newCompatRouteObjects(ctx, deps, cfg)
|
||||
@@ -133,7 +162,6 @@ func newCompatRouteObjects(ctx *server.Context, deps Dependencies, cfg CompatRou
|
||||
Headers: ctx.Request.Header.Clone(),
|
||||
ContentType: strings.TrimSpace(ctx.Request.Header.Get("content-type")),
|
||||
RequestBody: body,
|
||||
Mode: ctx.Mode,
|
||||
Deps: &deps,
|
||||
HTTPRequestID: resolveHTTPRequestID(ctx.Request),
|
||||
}
|
||||
|
||||
@@ -15,7 +15,6 @@ import (
|
||||
"time"
|
||||
|
||||
"cursor/gen/aiserverv1"
|
||||
"cursor/internal/backend/server"
|
||||
"cursor/internal/logger"
|
||||
"cursor/internal/netproxy"
|
||||
legacyruntime "cursor/internal/runtime"
|
||||
@@ -87,7 +86,7 @@ func buildUpstreamRequest(reqCtx *RequestContext, body []byte, options ForwardOp
|
||||
}
|
||||
upstreamRequest.Host = reqCtx.TargetURL.Host
|
||||
|
||||
if reqCtx.Mode == server.ModeLocal && shouldRewriteHost(reqCtx.TargetURL.Hostname()) {
|
||||
if shouldRewriteHost(reqCtx.TargetURL.Hostname()) {
|
||||
auth := formatBearerAuthorization(legacyruntime.LocalRelayToken)
|
||||
if auth == "" {
|
||||
return nil, nil, legacyruntime.ErrInvalidSystemSetting
|
||||
|
||||
@@ -20,6 +20,13 @@ type SystemSettingService interface {
|
||||
ResolveModelAdapters(context.Context) ([]legacyruntime.ModelAdapterConfig, error)
|
||||
}
|
||||
|
||||
// AuthorizationProvider supplies the independent Cursor account used only by
|
||||
// official control-plane requests such as Plugins, Skills, and MCP registry.
|
||||
type AuthorizationProvider interface {
|
||||
Authorization(context.Context) (string, error)
|
||||
SignedIn() bool
|
||||
}
|
||||
|
||||
type HTTPClient interface {
|
||||
Do(req *http.Request) (*http.Response, error)
|
||||
}
|
||||
@@ -41,7 +48,6 @@ type RequestContext struct {
|
||||
Headers http.Header
|
||||
ContentType string
|
||||
RequestBody []byte
|
||||
Mode server.ExecutionMode
|
||||
Deps *Dependencies
|
||||
HTTPRequestID string
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user