mirror of
https://wget.la/https://github.com/leookun/cursor-byok
synced 2026-10-08 07:21:13 +08:00
v0.3.8
This commit is contained in:
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,356 @@
|
||||
package forwarder
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"connectrpc.com/connect"
|
||||
|
||||
"cursor/gen/aiserverv1"
|
||||
"cursor/gen/aiserverv1/aiserverv1connect"
|
||||
)
|
||||
|
||||
type usageLookupRecord struct {
|
||||
InputTokens int64
|
||||
OutputTokens int64
|
||||
CreatedAt time.Time
|
||||
}
|
||||
|
||||
const (
|
||||
dashboardServiceGetTokenUsageProcedure = "/aiserver.v1.DashboardService/GetTokenUsage"
|
||||
dashboardServiceGetGlassEarlyPreviewEnrollmentProcedure = "/aiserver.v1.DashboardService/GetGlassEarlyPreviewEnrollment"
|
||||
)
|
||||
|
||||
func newAIHandler(service *Service) http.Handler {
|
||||
mux := http.NewServeMux()
|
||||
mux.Handle(
|
||||
dashboardServiceGetTokenUsageProcedure,
|
||||
connect.NewUnaryHandler(dashboardServiceGetTokenUsageProcedure, service.GetTokenUsage),
|
||||
)
|
||||
mux.Handle(
|
||||
dashboardServiceGetGlassEarlyPreviewEnrollmentProcedure,
|
||||
connect.NewUnaryHandler(dashboardServiceGetGlassEarlyPreviewEnrollmentProcedure, service.GetGlassEarlyPreviewEnrollment),
|
||||
)
|
||||
mux.Handle(
|
||||
aiserverv1connect.AiServiceCountTokensProcedure,
|
||||
connect.NewUnaryHandler(aiserverv1connect.AiServiceCountTokensProcedure, service.CountTokens),
|
||||
)
|
||||
mux.Handle(
|
||||
aiserverv1connect.AiServiceGetThoughtAnnotationProcedure,
|
||||
connect.NewUnaryHandler(aiserverv1connect.AiServiceGetThoughtAnnotationProcedure, service.GetThoughtAnnotation),
|
||||
)
|
||||
mux.Handle(
|
||||
aiserverv1connect.AiServiceWriteGitCommitMessageProcedure,
|
||||
connect.NewUnaryHandler(aiserverv1connect.AiServiceWriteGitCommitMessageProcedure, service.WriteGitCommitMessage),
|
||||
)
|
||||
mux.Handle(
|
||||
aiserverv1connect.AiServiceCreateExperimentalIndexProcedure,
|
||||
connect.NewUnaryHandler(aiserverv1connect.AiServiceCreateExperimentalIndexProcedure, service.CreateExperimentalIndex),
|
||||
)
|
||||
mux.Handle(
|
||||
aiserverv1connect.AiServiceListExperimentalIndexFilesProcedure,
|
||||
connect.NewUnaryHandler(aiserverv1connect.AiServiceListExperimentalIndexFilesProcedure, service.ListExperimentalIndexFiles),
|
||||
)
|
||||
mux.Handle(
|
||||
aiserverv1connect.AiServiceListenExperimentalIndexProcedure,
|
||||
connect.NewServerStreamHandler(aiserverv1connect.AiServiceListenExperimentalIndexProcedure, service.ListenExperimentalIndex),
|
||||
)
|
||||
mux.Handle(
|
||||
aiserverv1connect.AiServiceRegisterFileToIndexProcedure,
|
||||
connect.NewUnaryHandler(aiserverv1connect.AiServiceRegisterFileToIndexProcedure, service.RegisterFileToIndex),
|
||||
)
|
||||
mux.Handle(
|
||||
aiserverv1connect.AiServiceSetupIndexDependenciesProcedure,
|
||||
connect.NewUnaryHandler(aiserverv1connect.AiServiceSetupIndexDependenciesProcedure, service.SetupIndexDependencies),
|
||||
)
|
||||
mux.Handle(
|
||||
aiserverv1connect.AiServiceComputeIndexTopoSortProcedure,
|
||||
connect.NewUnaryHandler(aiserverv1connect.AiServiceComputeIndexTopoSortProcedure, service.ComputeIndexTopoSort),
|
||||
)
|
||||
mux.Handle(
|
||||
aiserverv1connect.AiServiceDocumentationQueryProcedure,
|
||||
connect.NewUnaryHandler(aiserverv1connect.AiServiceDocumentationQueryProcedure, service.DocumentationQuery),
|
||||
)
|
||||
mux.Handle(
|
||||
aiserverv1connect.AiServiceAvailableDocsProcedure,
|
||||
connect.NewUnaryHandler(aiserverv1connect.AiServiceAvailableDocsProcedure, service.AvailableDocs),
|
||||
)
|
||||
mux.Handle(
|
||||
aiserverv1connect.AiServiceKnowledgeBaseAddProcedure,
|
||||
connect.NewUnaryHandler(aiserverv1connect.AiServiceKnowledgeBaseAddProcedure, service.KnowledgeBaseAdd),
|
||||
)
|
||||
mux.Handle(
|
||||
aiserverv1connect.AiServiceKnowledgeBaseListProcedure,
|
||||
connect.NewUnaryHandler(aiserverv1connect.AiServiceKnowledgeBaseListProcedure, service.KnowledgeBaseList),
|
||||
)
|
||||
mux.Handle(
|
||||
aiserverv1connect.AiServiceKnowledgeBaseRemoveProcedure,
|
||||
connect.NewUnaryHandler(aiserverv1connect.AiServiceKnowledgeBaseRemoveProcedure, service.KnowledgeBaseRemove),
|
||||
)
|
||||
mux.Handle(
|
||||
aiserverv1connect.AiServiceKnowledgeBaseUpdateProcedure,
|
||||
connect.NewUnaryHandler(aiserverv1connect.AiServiceKnowledgeBaseUpdateProcedure, service.KnowledgeBaseUpdate),
|
||||
)
|
||||
mux.Handle(
|
||||
aiserverv1connect.AiServiceFetchRelevantKnowledgeForConversationProcedure,
|
||||
connect.NewUnaryHandler(aiserverv1connect.AiServiceFetchRelevantKnowledgeForConversationProcedure, service.FetchRelevantKnowledgeForConversation),
|
||||
)
|
||||
mux.Handle("/", http.NotFoundHandler())
|
||||
return mux
|
||||
}
|
||||
|
||||
func (service *Service) GetThoughtAnnotation(_ context.Context, req *connect.Request[aiserverv1.GetThoughtAnnotationRequest]) (*connect.Response[aiserverv1.GetThoughtAnnotationResponse], error) {
|
||||
requestID := strings.TrimSpace(req.Msg.GetRequestId())
|
||||
thought, ok, err := service.lookupThoughtAnnotation(requestID)
|
||||
if err != nil {
|
||||
return nil, connect.NewError(connect.CodeInternal, err)
|
||||
}
|
||||
if !ok {
|
||||
return connect.NewResponse(&aiserverv1.GetThoughtAnnotationResponse{}), nil
|
||||
}
|
||||
return connect.NewResponse(&aiserverv1.GetThoughtAnnotationResponse{
|
||||
ThoughtAnnotation: &aiserverv1.AiThoughtAnnotation{
|
||||
RequestId: requestID,
|
||||
Thought: thought,
|
||||
},
|
||||
}), nil
|
||||
}
|
||||
|
||||
func (service *Service) GetTokenUsage(_ context.Context, req *connect.Request[aiserverv1.GetTokenUsageRequest]) (*connect.Response[aiserverv1.GetTokenUsageResponse], error) {
|
||||
record, ok, err := service.lookupUsageRecord(strings.TrimSpace(req.Msg.GetUsageUuid()))
|
||||
if err != nil {
|
||||
return nil, connect.NewError(connect.CodeInternal, err)
|
||||
}
|
||||
if !ok {
|
||||
return connect.NewResponse(&aiserverv1.GetTokenUsageResponse{}), nil
|
||||
}
|
||||
return connect.NewResponse(&aiserverv1.GetTokenUsageResponse{
|
||||
InputTokens: clampInt64ToInt32(record.InputTokens),
|
||||
OutputTokens: clampInt64ToInt32(record.OutputTokens),
|
||||
}), nil
|
||||
}
|
||||
|
||||
func (service *Service) GetGlassEarlyPreviewEnrollment(context.Context, *connect.Request[aiserverv1.GetGlassEarlyPreviewEnrollmentRequest]) (*connect.Response[aiserverv1.GetGlassEarlyPreviewEnrollmentResponse], error) {
|
||||
granted := true
|
||||
return connect.NewResponse(&aiserverv1.GetGlassEarlyPreviewEnrollmentResponse{
|
||||
Enabled: true,
|
||||
EnterpriseGlassSelfEnrollEligible: &granted,
|
||||
GlassAccessGranted: &granted,
|
||||
}), nil
|
||||
}
|
||||
|
||||
func (service *Service) CountTokens(_ context.Context, req *connect.Request[aiserverv1.CountTokensRequest]) (*connect.Response[aiserverv1.CountTokensResponse], error) {
|
||||
total := int64(0)
|
||||
details := make([]*aiserverv1.ContextItemTokenDetail, 0, len(req.Msg.GetContextItems()))
|
||||
for _, item := range req.Msg.GetContextItems() {
|
||||
count := estimateContextItemTokens(item)
|
||||
total += count
|
||||
details = append(details, &aiserverv1.ContextItemTokenDetail{
|
||||
RelativeWorkspacePath: contextItemRelativeWorkspacePath(item),
|
||||
Count: clampInt64ToInt32(count),
|
||||
LineCount: lineCountForContextItem(item),
|
||||
})
|
||||
}
|
||||
return connect.NewResponse(&aiserverv1.CountTokensResponse{
|
||||
Count: clampInt64ToInt32(total),
|
||||
TokenDetails: details,
|
||||
}), nil
|
||||
}
|
||||
|
||||
func (service *Service) lookupUsageRecord(usageUUID string) (usageLookupRecord, bool, error) {
|
||||
if service == nil {
|
||||
return usageLookupRecord{}, false, nil
|
||||
}
|
||||
if service.usageStore == nil {
|
||||
return usageLookupRecord{}, false, nil
|
||||
}
|
||||
item, ok, err := service.usageStore.LookupEvent(strings.TrimSpace(usageUUID))
|
||||
if err != nil || !ok {
|
||||
return usageLookupRecord{}, ok, err
|
||||
}
|
||||
return usageLookupRecord{
|
||||
InputTokens: item.InputTokens + item.CacheReadTokens + item.CacheWriteTokens,
|
||||
OutputTokens: item.OutputTokens,
|
||||
CreatedAt: item.At,
|
||||
}, true, nil
|
||||
}
|
||||
|
||||
func (service *Service) lookupThoughtAnnotation(requestID string) (string, bool, error) {
|
||||
if service == nil || service.store == nil {
|
||||
return "", false, nil
|
||||
}
|
||||
needle := strings.TrimSpace(requestID)
|
||||
if needle == "" {
|
||||
return "", false, nil
|
||||
}
|
||||
foundRequest := false
|
||||
foundThought := false
|
||||
latestThought := ""
|
||||
latestThoughtAt := time.Time{}
|
||||
conversationIDs, err := service.store.ListConversationIDs()
|
||||
if err != nil {
|
||||
return "", false, err
|
||||
}
|
||||
for _, conversationID := range conversationIDs {
|
||||
conversation, err := service.store.LoadConversation(conversationID)
|
||||
if err != nil {
|
||||
return "", false, err
|
||||
}
|
||||
if conversation == nil {
|
||||
continue
|
||||
}
|
||||
for _, entry := range conversation.Entries {
|
||||
if strings.TrimSpace(entry.RequestID) != needle {
|
||||
continue
|
||||
}
|
||||
foundRequest = true
|
||||
if strings.TrimSpace(entry.Kind) != "metadata" {
|
||||
continue
|
||||
}
|
||||
var payload metadataPayload
|
||||
if err := json.Unmarshal(entry.Payload, &payload); err != nil {
|
||||
continue
|
||||
}
|
||||
if strings.TrimSpace(payload.Type) != "thought_annotation" {
|
||||
continue
|
||||
}
|
||||
if strings.TrimSpace(readStringValue(payload.Value["kind"])) != "summary_completed" {
|
||||
continue
|
||||
}
|
||||
thought := strings.TrimSpace(readStringValue(payload.Value["thought"]))
|
||||
if thought == "" {
|
||||
continue
|
||||
}
|
||||
if !foundThought || entry.CreatedAt.After(latestThoughtAt) {
|
||||
foundThought = true
|
||||
latestThought = thought
|
||||
latestThoughtAt = entry.CreatedAt
|
||||
}
|
||||
}
|
||||
}
|
||||
if foundThought {
|
||||
return latestThought, true, nil
|
||||
}
|
||||
if foundRequest {
|
||||
return defaultSummaryCompletedThought, true, nil
|
||||
}
|
||||
return "", false, nil
|
||||
}
|
||||
|
||||
func lineCountForContextItem(item *aiserverv1.ContextItem) int32 {
|
||||
if item == nil {
|
||||
return 0
|
||||
}
|
||||
if chunk := item.GetFileChunk(); chunk != nil {
|
||||
return countTextLines(chunk.GetChunkContents())
|
||||
}
|
||||
if outline := item.GetOutlineChunk(); outline != nil {
|
||||
return lineCountForRange(outline.GetFullRange())
|
||||
}
|
||||
if selection := item.GetCmdKSelection(); selection != nil {
|
||||
return int32(len(selection.GetLines()))
|
||||
}
|
||||
if sparse := item.GetSparseFileChunk(); sparse != nil {
|
||||
return int32(len(sparse.GetLines()))
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
func contextItemRelativeWorkspacePath(item *aiserverv1.ContextItem) string {
|
||||
if item == nil {
|
||||
return ""
|
||||
}
|
||||
if chunk := item.GetFileChunk(); chunk != nil {
|
||||
return strings.TrimSpace(chunk.GetRelativeWorkspacePath())
|
||||
}
|
||||
if outline := item.GetOutlineChunk(); outline != nil {
|
||||
return strings.TrimSpace(outline.GetRelativeWorkspacePath())
|
||||
}
|
||||
if sparse := item.GetSparseFileChunk(); sparse != nil {
|
||||
return strings.TrimSpace(sparse.GetRelativeWorkspacePath())
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func lineCountForRange(lineRange *aiserverv1.LineRange) int32 {
|
||||
if lineRange == nil {
|
||||
return 0
|
||||
}
|
||||
start := lineRange.GetStartLineNumber()
|
||||
end := lineRange.GetEndLineNumberInclusive()
|
||||
if start <= 0 || end < start {
|
||||
return 0
|
||||
}
|
||||
return end - start + 1
|
||||
}
|
||||
|
||||
func countTextLines(text string) int32 {
|
||||
if strings.TrimSpace(text) == "" {
|
||||
return 0
|
||||
}
|
||||
return int32(strings.Count(text, "\n") + 1)
|
||||
}
|
||||
|
||||
func readInt64Value(value any) int64 {
|
||||
switch typed := value.(type) {
|
||||
case float64:
|
||||
return int64(typed)
|
||||
case float32:
|
||||
return int64(typed)
|
||||
case int:
|
||||
return int64(typed)
|
||||
case int32:
|
||||
return int64(typed)
|
||||
case int64:
|
||||
return typed
|
||||
case uint32:
|
||||
return int64(typed)
|
||||
case uint64:
|
||||
if typed > uint64(^uint32(0)>>1) {
|
||||
return int64(^uint32(0) >> 1)
|
||||
}
|
||||
return int64(typed)
|
||||
case json.Number:
|
||||
value, err := typed.Int64()
|
||||
if err == nil {
|
||||
return value
|
||||
}
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
func readStringValue(value any) string {
|
||||
switch typed := value.(type) {
|
||||
case string:
|
||||
return strings.TrimSpace(typed)
|
||||
default:
|
||||
return ""
|
||||
}
|
||||
}
|
||||
|
||||
func readBoolValue(value any) bool {
|
||||
switch typed := value.(type) {
|
||||
case bool:
|
||||
return typed
|
||||
case string:
|
||||
return strings.EqualFold(strings.TrimSpace(typed), "true")
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func readTimeValue(value any) time.Time {
|
||||
return parseRFC3339Time(readStringValue(value))
|
||||
}
|
||||
|
||||
func firstNonZeroTime(values ...time.Time) time.Time {
|
||||
for _, value := range values {
|
||||
if !value.IsZero() {
|
||||
return value.UTC()
|
||||
}
|
||||
}
|
||||
return time.Time{}
|
||||
}
|
||||
@@ -0,0 +1,149 @@
|
||||
package forwarder
|
||||
|
||||
import (
|
||||
"context"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
const appendSequenceRetention = 10 * time.Minute
|
||||
|
||||
type appendSequenceTracker struct {
|
||||
mu sync.Mutex
|
||||
states map[string]*appendSequenceState
|
||||
}
|
||||
|
||||
type appendSequenceState struct {
|
||||
mu sync.Mutex
|
||||
next int64
|
||||
processing bool
|
||||
ready chan struct{}
|
||||
updatedAt time.Time
|
||||
}
|
||||
|
||||
type appendSequenceTicket struct {
|
||||
state *appendSequenceState
|
||||
seq int64
|
||||
}
|
||||
|
||||
func newAppendSequenceTracker() *appendSequenceTracker {
|
||||
return &appendSequenceTracker{
|
||||
states: make(map[string]*appendSequenceState),
|
||||
}
|
||||
}
|
||||
|
||||
func (tracker *appendSequenceTracker) Acquire(ctx context.Context, requestID string, appendSeq int64) (appendSequenceTicket, bool, error) {
|
||||
if tracker == nil || strings.TrimSpace(requestID) == "" || appendSeq <= 0 {
|
||||
return appendSequenceTicket{}, false, nil
|
||||
}
|
||||
state := tracker.state(strings.TrimSpace(requestID))
|
||||
stale, err := state.acquire(ctx, appendSeq)
|
||||
if err != nil || stale {
|
||||
return appendSequenceTicket{}, stale, err
|
||||
}
|
||||
return appendSequenceTicket{
|
||||
state: state,
|
||||
seq: appendSeq,
|
||||
}, false, nil
|
||||
}
|
||||
|
||||
func (tracker *appendSequenceTracker) state(requestID string) *appendSequenceState {
|
||||
now := time.Now().UTC()
|
||||
cutoff := now.Add(-appendSequenceRetention)
|
||||
|
||||
tracker.mu.Lock()
|
||||
defer tracker.mu.Unlock()
|
||||
for key, state := range tracker.states {
|
||||
if state == nil || state.expired(cutoff) {
|
||||
delete(tracker.states, key)
|
||||
}
|
||||
}
|
||||
if state, ok := tracker.states[requestID]; ok && state != nil {
|
||||
state.touch(now)
|
||||
return state
|
||||
}
|
||||
state := &appendSequenceState{
|
||||
next: 1,
|
||||
ready: make(chan struct{}),
|
||||
updatedAt: now,
|
||||
}
|
||||
tracker.states[requestID] = state
|
||||
return state
|
||||
}
|
||||
|
||||
func (state *appendSequenceState) acquire(ctx context.Context, appendSeq int64) (bool, error) {
|
||||
for {
|
||||
state.mu.Lock()
|
||||
now := time.Now().UTC()
|
||||
if state.next <= 0 {
|
||||
state.next = 1
|
||||
}
|
||||
if state.ready == nil {
|
||||
state.ready = make(chan struct{})
|
||||
}
|
||||
state.updatedAt = now
|
||||
switch {
|
||||
case appendSeq < state.next:
|
||||
state.mu.Unlock()
|
||||
return true, nil
|
||||
case appendSeq == state.next && !state.processing:
|
||||
state.processing = true
|
||||
state.mu.Unlock()
|
||||
return false, nil
|
||||
default:
|
||||
ready := state.ready
|
||||
state.mu.Unlock()
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return false, ctx.Err()
|
||||
case <-ready:
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (state *appendSequenceState) Release(seq int64) {
|
||||
if state == nil || seq <= 0 {
|
||||
return
|
||||
}
|
||||
|
||||
state.mu.Lock()
|
||||
defer state.mu.Unlock()
|
||||
|
||||
if state.processing && state.next == seq {
|
||||
state.processing = false
|
||||
state.next++
|
||||
close(state.ready)
|
||||
state.ready = make(chan struct{})
|
||||
}
|
||||
state.updatedAt = time.Now().UTC()
|
||||
}
|
||||
|
||||
func (state *appendSequenceState) expired(cutoff time.Time) bool {
|
||||
if state == nil {
|
||||
return true
|
||||
}
|
||||
state.mu.Lock()
|
||||
defer state.mu.Unlock()
|
||||
if state.processing {
|
||||
return false
|
||||
}
|
||||
return !state.updatedAt.IsZero() && state.updatedAt.Before(cutoff)
|
||||
}
|
||||
|
||||
func (state *appendSequenceState) touch(now time.Time) {
|
||||
if state == nil {
|
||||
return
|
||||
}
|
||||
state.mu.Lock()
|
||||
state.updatedAt = now
|
||||
state.mu.Unlock()
|
||||
}
|
||||
|
||||
func (ticket appendSequenceTicket) Release() {
|
||||
if ticket.state == nil || ticket.seq <= 0 {
|
||||
return
|
||||
}
|
||||
ticket.state.Release(ticket.seq)
|
||||
}
|
||||
@@ -0,0 +1,310 @@
|
||||
package forwarder
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
// artifactRecorder 只缓存当前 provider call 的请求摘要,用于本轮内 usage/model 信息补齐。
|
||||
// 对话恢复事实只存在 state.json 与 context.json。
|
||||
type artifactRecorder struct {
|
||||
store *ConversationFileStore
|
||||
broker *StreamBroker
|
||||
debug *debugRecorder
|
||||
|
||||
mu sync.Mutex
|
||||
sessions map[string]artifactSession
|
||||
}
|
||||
|
||||
type artifactSession struct {
|
||||
conversationID string
|
||||
turnSeq int64
|
||||
requestPayload map[string]any
|
||||
summaryPayload map[string]any
|
||||
}
|
||||
|
||||
type requestArtifactPrefix struct {
|
||||
Provider string
|
||||
OpenAIEndpoint string
|
||||
Model string
|
||||
PromptTokensTotal int64
|
||||
ReplayMessageCount int
|
||||
CanonicalBodyHash string
|
||||
FrontierHash string
|
||||
FrontierPath string
|
||||
BreakpointCount int
|
||||
ExpectedCacheRead bool
|
||||
PreviousFrontierMatched bool
|
||||
}
|
||||
|
||||
func newArtifactRecorder(store *ConversationFileStore, broker *StreamBroker, debug *debugRecorder) *artifactRecorder {
|
||||
return &artifactRecorder{
|
||||
store: store,
|
||||
broker: broker,
|
||||
debug: debug,
|
||||
sessions: make(map[string]artifactSession),
|
||||
}
|
||||
}
|
||||
|
||||
func (recorder *artifactRecorder) RecordLLMRequest(requestID string, _ string, modelCallID string, payload map[string]any) (string, error) {
|
||||
session, err := recorder.ensureSession(requestID, modelCallID)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
session.requestPayload = cloneStringAnyMap(payload)
|
||||
recorder.mu.Lock()
|
||||
recorder.sessions[artifactSessionKey(requestID, modelCallID)] = session
|
||||
recorder.mu.Unlock()
|
||||
if prefix, ok, err := decodeRequestPrefixPayload(session.requestPayload); err == nil && ok && prefix != nil {
|
||||
recorder.persistLatestRequestPrefix(session.conversationID, requestID, modelCallID, prefix)
|
||||
}
|
||||
recorder.debug.LogProviderArtifact(context.Background(), requestID, session.conversationID, modelCallID, "llm_request", session.requestPayload)
|
||||
return "", nil
|
||||
}
|
||||
|
||||
func (recorder *artifactRecorder) AppendLLMResponseChunk(requestID string, runID string, modelCallID string, chunk string) (string, error) {
|
||||
session, err := recorder.ensureSession(requestID, modelCallID)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
recorder.debug.LogProviderArtifact(context.Background(), requestID, session.conversationID, modelCallID, "llm_response_chunk", map[string]any{
|
||||
"run_id": strings.TrimSpace(runID),
|
||||
"raw_chunk": chunk,
|
||||
"byte_len": len([]byte(chunk)),
|
||||
})
|
||||
return "", nil
|
||||
}
|
||||
|
||||
func (recorder *artifactRecorder) RecordLLMSummary(requestID string, _ string, modelCallID string, payload map[string]any) (string, error) {
|
||||
session, err := recorder.ensureSession(requestID, modelCallID)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
session.summaryPayload = cloneStringAnyMap(payload)
|
||||
recorder.mu.Lock()
|
||||
recorder.sessions[artifactSessionKey(requestID, modelCallID)] = session
|
||||
recorder.mu.Unlock()
|
||||
if prefix, ok, err := decodeRequestPrefixPayload(session.requestPayload); err == nil && ok && prefix != nil {
|
||||
if tokens := readInt64Value(session.summaryPayload["prompt_tokens_total"]); tokens > 0 {
|
||||
prefix.PromptTokensTotal = tokens
|
||||
}
|
||||
recorder.persistLatestRequestPrefix(session.conversationID, requestID, modelCallID, prefix)
|
||||
}
|
||||
recorder.debug.LogProviderArtifact(context.Background(), requestID, session.conversationID, modelCallID, "llm_summary", session.summaryPayload)
|
||||
return "", nil
|
||||
}
|
||||
|
||||
func (recorder *artifactRecorder) persistLatestRequestPrefix(conversationID string, requestID string, modelCallID string, prefix *requestArtifactPrefix) {
|
||||
if recorder == nil || recorder.store == nil || prefix == nil || strings.TrimSpace(conversationID) == "" || strings.TrimSpace(conversationID) == "unknown" {
|
||||
return
|
||||
}
|
||||
_, _ = recorder.store.UpdateConversationMeta(conversationID, func(item *ConversationFile) error {
|
||||
if item == nil {
|
||||
return nil
|
||||
}
|
||||
item.LatestRequestPrefix = &ConversationRequestPrefix{
|
||||
RequestID: strings.TrimSpace(requestID),
|
||||
ModelCallID: firstNonEmpty(strings.TrimSpace(modelCallID), strings.TrimSpace(requestID)),
|
||||
Provider: strings.TrimSpace(prefix.Provider),
|
||||
OpenAIEndpoint: strings.TrimSpace(prefix.OpenAIEndpoint),
|
||||
Model: strings.TrimSpace(prefix.Model),
|
||||
PromptTokensTotal: prefix.PromptTokensTotal,
|
||||
ReplayMessageCount: prefix.ReplayMessageCount,
|
||||
CanonicalBodyHash: strings.TrimSpace(prefix.CanonicalBodyHash),
|
||||
FrontierHash: strings.TrimSpace(prefix.FrontierHash),
|
||||
FrontierPath: strings.TrimSpace(prefix.FrontierPath),
|
||||
BreakpointCount: prefix.BreakpointCount,
|
||||
ExpectedCacheRead: prefix.ExpectedCacheRead,
|
||||
PreviousFrontierMatched: prefix.PreviousFrontierMatched,
|
||||
UpdatedAt: time.Now().UTC(),
|
||||
}
|
||||
return nil
|
||||
})
|
||||
}
|
||||
|
||||
func (recorder *artifactRecorder) ClearActiveArtifacts(requestID string, modelCallID string) {
|
||||
if recorder == nil {
|
||||
return
|
||||
}
|
||||
recorder.mu.Lock()
|
||||
delete(recorder.sessions, artifactSessionKey(requestID, modelCallID))
|
||||
recorder.mu.Unlock()
|
||||
}
|
||||
|
||||
func (recorder *artifactRecorder) ensureSession(requestID string, modelCallID string) (artifactSession, error) {
|
||||
if recorder == nil {
|
||||
return artifactSession{}, fmt.Errorf("artifact recorder is nil")
|
||||
}
|
||||
key := artifactSessionKey(requestID, modelCallID)
|
||||
recorder.mu.Lock()
|
||||
defer recorder.mu.Unlock()
|
||||
if session, ok := recorder.sessions[key]; ok {
|
||||
return session, nil
|
||||
}
|
||||
conversationID, turnSeq := recorder.resolveConversationContext(requestID)
|
||||
session := artifactSession{
|
||||
conversationID: conversationID,
|
||||
turnSeq: turnSeq,
|
||||
}
|
||||
recorder.sessions[key] = session
|
||||
return session, nil
|
||||
}
|
||||
|
||||
func artifactSessionKey(requestID string, modelCallID string) string {
|
||||
return strings.TrimSpace(requestID) + "::" + strings.TrimSpace(modelCallID)
|
||||
}
|
||||
|
||||
func (recorder *artifactRecorder) resolveConversationContext(requestID string) (string, int64) {
|
||||
if recorder == nil || recorder.broker == nil {
|
||||
return "unknown", 0
|
||||
}
|
||||
stream, ok := recorder.broker.Get(requestID)
|
||||
if !ok || stream == nil {
|
||||
return "unknown", 0
|
||||
}
|
||||
stream.mu.Lock()
|
||||
defer stream.mu.Unlock()
|
||||
return sanitizeArtifactName(firstNonEmpty(stream.ConversationID, "unknown")), stream.TurnSeq
|
||||
}
|
||||
|
||||
func decodeRequestPrefixPayload(payload map[string]any) (*requestArtifactPrefix, bool, error) {
|
||||
if len(payload) == 0 {
|
||||
return nil, false, nil
|
||||
}
|
||||
provider := strings.TrimSpace(readStringValue(payload["provider"]))
|
||||
model := strings.TrimSpace(readStringValue(payload["model"]))
|
||||
openAIEndpoint := strings.TrimSpace(readStringValue(payload["openai_endpoint"]))
|
||||
if provider == "" && model == "" && openAIEndpoint == "" {
|
||||
return nil, false, nil
|
||||
}
|
||||
replayMessageCount := replayMessageCountFromRequestPayload(payload)
|
||||
frontier := cacheFrontierFromRequestPayload(payload)
|
||||
return &requestArtifactPrefix{
|
||||
Provider: provider,
|
||||
OpenAIEndpoint: openAIEndpoint,
|
||||
Model: model,
|
||||
ReplayMessageCount: replayMessageCount,
|
||||
CanonicalBodyHash: frontier.CanonicalBodyHash,
|
||||
FrontierHash: frontier.FrontierHash,
|
||||
FrontierPath: frontier.FrontierPath,
|
||||
BreakpointCount: frontier.BreakpointCount,
|
||||
ExpectedCacheRead: frontier.ExpectedCacheRead,
|
||||
PreviousFrontierMatched: frontier.PreviousFrontierMatched,
|
||||
}, true, nil
|
||||
}
|
||||
|
||||
type requestCacheFrontierPayload struct {
|
||||
CanonicalBodyHash string
|
||||
FrontierHash string
|
||||
FrontierPath string
|
||||
BreakpointCount int
|
||||
ExpectedCacheRead bool
|
||||
PreviousFrontierMatched bool
|
||||
}
|
||||
|
||||
func cacheFrontierFromRequestPayload(payload map[string]any) requestCacheFrontierPayload {
|
||||
var frontier requestCacheFrontierPayload
|
||||
knobs, ok := payload["request_knobs"].(map[string]any)
|
||||
if !ok {
|
||||
return frontier
|
||||
}
|
||||
rawFrontier, ok := knobs["cache_frontier"].(map[string]any)
|
||||
if !ok {
|
||||
return frontier
|
||||
}
|
||||
frontier.CanonicalBodyHash = strings.TrimSpace(readStringValue(rawFrontier["canonical_body_hash"]))
|
||||
frontier.FrontierHash = strings.TrimSpace(readStringValue(rawFrontier["frontier_hash"]))
|
||||
frontier.FrontierPath = strings.TrimSpace(readStringValue(rawFrontier["frontier_path"]))
|
||||
frontier.BreakpointCount = int(readInt64Value(rawFrontier["breakpoint_count"]))
|
||||
frontier.ExpectedCacheRead = readBoolValue(rawFrontier["expected_cache_read"])
|
||||
frontier.PreviousFrontierMatched = readBoolValue(rawFrontier["previous_frontier_matched"])
|
||||
return frontier
|
||||
}
|
||||
|
||||
func replayMessageCountFromRequestPayload(payload map[string]any) int {
|
||||
switch messages := payload["messages_summary"].(type) {
|
||||
case []map[string]any:
|
||||
return replayMessageCountFromSummaryObjects(messages)
|
||||
case []any:
|
||||
items := make([]map[string]any, 0, len(messages))
|
||||
for _, item := range messages {
|
||||
message, ok := item.(map[string]any)
|
||||
if ok {
|
||||
items = append(items, message)
|
||||
}
|
||||
}
|
||||
return replayMessageCountFromSummaryObjects(items)
|
||||
default:
|
||||
return 0
|
||||
}
|
||||
}
|
||||
|
||||
func replayMessageCountFromSummaryObjects(messages []map[string]any) int {
|
||||
count := 0
|
||||
for _, message := range messages {
|
||||
if strings.TrimSpace(readStringValue(message["role"])) == "system" {
|
||||
continue
|
||||
}
|
||||
count++
|
||||
}
|
||||
return count
|
||||
}
|
||||
|
||||
func sanitizeArtifactName(value string) string {
|
||||
replacer := strings.NewReplacer("/", "_", "\\", "_", ":", "_", " ", "_")
|
||||
normalized := replacer.Replace(strings.TrimSpace(value))
|
||||
if normalized == "" {
|
||||
return "unknown"
|
||||
}
|
||||
return normalized
|
||||
}
|
||||
|
||||
func artifactConversationDir(historyRoot string, conversationID string) string {
|
||||
return filepath.Join(strings.TrimSpace(historyRoot), sanitizeArtifactName(conversationID))
|
||||
}
|
||||
|
||||
func openUniqueArtifactTempFile(path string) (*os.File, string, error) {
|
||||
dir := filepath.Dir(path)
|
||||
base := filepath.Base(path)
|
||||
for attempt := 0; attempt < 32; attempt++ {
|
||||
tempPath := filepath.Join(dir, fmt.Sprintf("%s.tmp-%d-%d", base, os.Getpid(), time.Now().UnixNano()))
|
||||
file, err := os.OpenFile(tempPath, os.O_CREATE|os.O_EXCL|os.O_WRONLY, 0o644)
|
||||
if err == nil {
|
||||
return file, tempPath, nil
|
||||
}
|
||||
if !os.IsExist(err) {
|
||||
return nil, "", err
|
||||
}
|
||||
}
|
||||
return nil, "", fmt.Errorf("exhausted temp file attempts")
|
||||
}
|
||||
|
||||
func renameArtifactTempFile(tempPath string, path string) error {
|
||||
attempts := 1
|
||||
delay := 10 * time.Millisecond
|
||||
if runtime.GOOS == "windows" {
|
||||
attempts = 12
|
||||
}
|
||||
var lastErr error
|
||||
for attempt := 0; attempt < attempts; attempt++ {
|
||||
if err := os.Rename(tempPath, path); err == nil {
|
||||
return nil
|
||||
} else {
|
||||
lastErr = err
|
||||
}
|
||||
if runtime.GOOS != "windows" || os.IsNotExist(lastErr) || attempt == attempts-1 {
|
||||
break
|
||||
}
|
||||
time.Sleep(delay)
|
||||
if delay < 200*time.Millisecond {
|
||||
delay *= 2
|
||||
}
|
||||
}
|
||||
return lastErr
|
||||
}
|
||||
@@ -0,0 +1,954 @@
|
||||
package forwarder
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"regexp"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"cursor/gen/agentv1"
|
||||
execbridge "cursor/internal/backend/agent/bridge/exec"
|
||||
runtimecore "cursor/internal/backend/agent/core"
|
||||
)
|
||||
|
||||
const awaitShellOutputLimit = 16 * 1024
|
||||
|
||||
const (
|
||||
backgroundShellStatusBackgrounded = "backgrounded"
|
||||
backgroundShellStatusRunning = "running"
|
||||
backgroundShellStatusCompleted = "completed"
|
||||
backgroundShellStatusRejected = "rejected"
|
||||
backgroundShellStatusPermissionDenied = "permission_denied"
|
||||
backgroundShellStatusTransportClosed = "transport_closed"
|
||||
backgroundShellStatusUnknown = "unknown"
|
||||
|
||||
backgroundShellActionSourceClient = "client"
|
||||
backgroundShellActionSourceLocalBackgrounded = "local_shell_backgrounded"
|
||||
)
|
||||
|
||||
type awaitShellArgs struct {
|
||||
ShellID string `json:"shell_id,omitempty"`
|
||||
TaskID string `json:"task_id,omitempty"`
|
||||
BlockUntilMS *int64 `json:"block_until_ms,omitempty"`
|
||||
Pattern string `json:"pattern,omitempty"`
|
||||
}
|
||||
|
||||
type awaitShellResult struct {
|
||||
ShellID string `json:"shell_id,omitempty"`
|
||||
Status string `json:"status"`
|
||||
Matched bool `json:"matched"`
|
||||
TimedOut bool `json:"timed_out"`
|
||||
ExitCode *int64 `json:"exit_code,omitempty"`
|
||||
Stdout string `json:"stdout,omitempty"`
|
||||
Stderr string `json:"stderr,omitempty"`
|
||||
StdoutOffset int `json:"stdout_offset,omitempty"`
|
||||
StderrOffset int `json:"stderr_offset,omitempty"`
|
||||
RuntimeMS uint64 `json:"runtime_ms,omitempty"`
|
||||
OutputLength uint64 `json:"output_length,omitempty"`
|
||||
RegexRequested bool `json:"regex_requested,omitempty"`
|
||||
RegexMatch *string `json:"regex_match,omitempty"`
|
||||
Message string `json:"message,omitempty"`
|
||||
}
|
||||
|
||||
func (service *Service) handleAwaitShellToolInvocation(stream *ActiveStream, invocation runtimecore.ToolInvocation) error {
|
||||
args, err := decodeAwaitShellArgs(invocation.ArgsJSON)
|
||||
if err != nil {
|
||||
result := awaitShellResult{Status: "error", Message: err.Error()}
|
||||
payload, encodeErr := json.Marshal(result)
|
||||
if encodeErr != nil {
|
||||
return encodeErr
|
||||
}
|
||||
return service.completeImmediateToolResult(stream, invocation, string(payload), buildAwaitShellToolCall(buildAwaitArgsFromAwaitShellArgs(args), buildAwaitShellProtoResult(result)))
|
||||
}
|
||||
result := service.awaitShellSnapshot(stream, args)
|
||||
payload, err := json.Marshal(result)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return service.completeImmediateToolResult(stream, invocation, string(payload), buildAwaitShellToolCall(buildAwaitArgsFromAwaitShellArgs(args), buildAwaitShellProtoResult(result)))
|
||||
}
|
||||
|
||||
func decodeAwaitShellArgs(raw []byte) (awaitShellArgs, error) {
|
||||
argsMap, err := runtimecore.DecodeArgsMap(raw)
|
||||
if err != nil {
|
||||
return awaitShellArgs{}, fmt.Errorf("decode AwaitShell args failed: %w", err)
|
||||
}
|
||||
result := awaitShellArgs{
|
||||
ShellID: strings.TrimSpace(runtimecore.ReadStringArg(argsMap, "shell_id")),
|
||||
TaskID: strings.TrimSpace(runtimecore.ReadStringArg(argsMap, "task_id")),
|
||||
Pattern: strings.TrimSpace(runtimecore.ReadStringArg(argsMap, "pattern")),
|
||||
}
|
||||
if result.ShellID == "" {
|
||||
result.ShellID = result.TaskID
|
||||
}
|
||||
if value, found, err := runtimecore.ReadInt64Arg(argsMap, "block_until_ms"); err != nil {
|
||||
return result, err
|
||||
} else if found {
|
||||
if value < 0 {
|
||||
value = 0
|
||||
}
|
||||
result.BlockUntilMS = &value
|
||||
}
|
||||
if result.BlockUntilMS != nil && *result.BlockUntilMS == 0 && strings.TrimSpace(result.ShellID) == "" {
|
||||
return result, fmt.Errorf("AwaitShell shell_id is required when block_until_ms is 0")
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func (service *Service) awaitShellSnapshot(stream *ActiveStream, args awaitShellArgs) awaitShellResult {
|
||||
blockUntilMS := int64(30000)
|
||||
if args.BlockUntilMS != nil {
|
||||
blockUntilMS = *args.BlockUntilMS
|
||||
}
|
||||
shellID := strings.TrimSpace(args.ShellID)
|
||||
if shellID == "" {
|
||||
return awaitShellResult{
|
||||
Status: "waited",
|
||||
TimedOut: false,
|
||||
Message: fmt.Sprintf("waited %dms", blockUntilMS),
|
||||
}
|
||||
}
|
||||
|
||||
service.refreshBackgroundShellFromTerminalFile(stream, shellID)
|
||||
|
||||
stream.mu.Lock()
|
||||
state, ok := stream.BackgroundShells[shellID]
|
||||
if !ok || state == nil {
|
||||
stream.mu.Unlock()
|
||||
return awaitShellResult{
|
||||
ShellID: shellID,
|
||||
Status: backgroundShellStatusUnknown,
|
||||
TimedOut: false,
|
||||
Message: "unknown or expired shell_id",
|
||||
}
|
||||
}
|
||||
stdoutStart := clampOffset(state.AwaitStdoutOffset, len(state.StdoutBuffer))
|
||||
stderrStart := clampOffset(state.AwaitStderrOffset, len(state.StderrBuffer))
|
||||
stdout := state.StdoutBuffer[stdoutStart:]
|
||||
stderr := state.StderrBuffer[stderrStart:]
|
||||
stdoutEnd := len(state.StdoutBuffer)
|
||||
stderrEnd := len(state.StderrBuffer)
|
||||
state.AwaitStdoutOffset = stdoutEnd
|
||||
state.AwaitStderrOffset = stderrEnd
|
||||
status := strings.TrimSpace(state.Status)
|
||||
if status == "" {
|
||||
status = backgroundShellStatusUnknown
|
||||
}
|
||||
var exitCode *int64
|
||||
if state.ExitCode != nil {
|
||||
value := int64(*state.ExitCode)
|
||||
exitCode = &value
|
||||
}
|
||||
createdAt := state.CreatedAt
|
||||
completedAt := state.CompletedAt
|
||||
combinedOutput := state.StdoutBuffer + "\n" + state.StderrBuffer
|
||||
stream.UpdatedAt = time.Now().UTC()
|
||||
stream.mu.Unlock()
|
||||
|
||||
matched, matchText, patternErr := awaitShellPatternMatched(args.Pattern, combinedOutput)
|
||||
message := ""
|
||||
if patternErr != nil {
|
||||
message = patternErr.Error()
|
||||
}
|
||||
timedOut := blockUntilMS > 0 && !matched && !isBackgroundShellTerminalStatus(status)
|
||||
stdout = truncateAwaitShellOutput(stdout)
|
||||
stderr = truncateAwaitShellOutput(stderr)
|
||||
return awaitShellResult{
|
||||
ShellID: shellID,
|
||||
Status: status,
|
||||
Matched: matched,
|
||||
TimedOut: timedOut,
|
||||
ExitCode: exitCode,
|
||||
Stdout: stdout,
|
||||
Stderr: stderr,
|
||||
StdoutOffset: stdoutEnd,
|
||||
StderrOffset: stderrEnd,
|
||||
RuntimeMS: backgroundShellRuntimeMS(createdAt, completedAt),
|
||||
OutputLength: uint64(len(combinedOutput)),
|
||||
RegexRequested: strings.TrimSpace(args.Pattern) != "",
|
||||
RegexMatch: matchText,
|
||||
Message: message,
|
||||
}
|
||||
}
|
||||
|
||||
func buildAwaitArgsFromAwaitShellArgs(args awaitShellArgs) *agentv1.AwaitArgs {
|
||||
awaitArgs := &agentv1.AwaitArgs{TaskId: strings.TrimSpace(firstNonEmpty(args.ShellID, args.TaskID))}
|
||||
if args.BlockUntilMS != nil {
|
||||
value := uint32(0)
|
||||
if *args.BlockUntilMS > 0 {
|
||||
value = uint32(*args.BlockUntilMS)
|
||||
}
|
||||
awaitArgs.BlockUntilMs = &value
|
||||
}
|
||||
if pattern := strings.TrimSpace(args.Pattern); pattern != "" {
|
||||
awaitArgs.Regex = &pattern
|
||||
}
|
||||
return awaitArgs
|
||||
}
|
||||
|
||||
func buildAwaitShellToolCall(args *agentv1.AwaitArgs, result *agentv1.AwaitResult) *agentv1.ToolCall {
|
||||
return &agentv1.ToolCall{
|
||||
Tool: &agentv1.ToolCall_AwaitToolCall{
|
||||
AwaitToolCall: &agentv1.AwaitToolCall{
|
||||
Args: args,
|
||||
Result: result,
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func buildAwaitShellProtoResult(result awaitShellResult) *agentv1.AwaitResult {
|
||||
switch strings.TrimSpace(result.Status) {
|
||||
case "error", backgroundShellStatusUnknown:
|
||||
message := firstNonEmpty(strings.TrimSpace(result.Message), "unknown or expired shell_id")
|
||||
return &agentv1.AwaitResult{Result: &agentv1.AwaitResult_Error{Error: &agentv1.AwaitError{Error: message}}}
|
||||
}
|
||||
if isBackgroundShellTerminalStatus(result.Status) || result.Matched {
|
||||
task := &agentv1.AwaitTaskComplete{
|
||||
TaskId: strings.TrimSpace(result.ShellID),
|
||||
RuntimeMs: result.RuntimeMS,
|
||||
OutputLength: result.OutputLength,
|
||||
RegexRequested: result.RegexRequested,
|
||||
RegexMatch: result.RegexMatch,
|
||||
}
|
||||
if result.ExitCode != nil {
|
||||
exitCode := int32(*result.ExitCode)
|
||||
task.ExitCode = &exitCode
|
||||
}
|
||||
return &agentv1.AwaitResult{Result: &agentv1.AwaitResult_Complete{Complete: task}}
|
||||
}
|
||||
task := &agentv1.AwaitTaskStillRunning{
|
||||
TaskId: strings.TrimSpace(result.ShellID),
|
||||
RuntimeMs: result.RuntimeMS,
|
||||
OutputLength: result.OutputLength,
|
||||
RegexRequested: result.RegexRequested,
|
||||
RegexMatch: result.RegexMatch,
|
||||
}
|
||||
return &agentv1.AwaitResult{Result: &agentv1.AwaitResult_StillRunning{StillRunning: task}}
|
||||
}
|
||||
|
||||
func backgroundShellRuntimeMS(createdAt time.Time, completedAt time.Time) uint64 {
|
||||
if createdAt.IsZero() {
|
||||
return 0
|
||||
}
|
||||
end := completedAt
|
||||
if end.IsZero() {
|
||||
end = time.Now().UTC()
|
||||
}
|
||||
if end.Before(createdAt) {
|
||||
return 0
|
||||
}
|
||||
return uint64(end.Sub(createdAt).Milliseconds())
|
||||
}
|
||||
|
||||
func awaitShellPatternMatched(pattern string, output string) (bool, *string, error) {
|
||||
trimmed := strings.TrimSpace(pattern)
|
||||
if trimmed == "" {
|
||||
return false, nil, nil
|
||||
}
|
||||
expr, err := regexp.Compile("(?m)" + trimmed)
|
||||
if err != nil {
|
||||
return false, nil, fmt.Errorf("invalid AwaitShell pattern: %w", err)
|
||||
}
|
||||
match := expr.FindString(output)
|
||||
if match == "" {
|
||||
return false, nil, nil
|
||||
}
|
||||
return true, &match, nil
|
||||
}
|
||||
|
||||
func truncateAwaitShellOutput(value string) string {
|
||||
if len(value) <= awaitShellOutputLimit {
|
||||
return value
|
||||
}
|
||||
return value[len(value)-awaitShellOutputLimit:]
|
||||
}
|
||||
|
||||
func clampOffset(offset int, length int) int {
|
||||
if offset < 0 {
|
||||
return 0
|
||||
}
|
||||
if offset > length {
|
||||
return length
|
||||
}
|
||||
return offset
|
||||
}
|
||||
|
||||
type terminalShellFileSnapshot struct {
|
||||
PID *uint32
|
||||
Command string
|
||||
CWD string
|
||||
StartedAt time.Time
|
||||
Output string
|
||||
ExitCode *int32
|
||||
EndedAt time.Time
|
||||
}
|
||||
|
||||
func (service *Service) refreshBackgroundShellFromTerminalFile(stream *ActiveStream, shellID string) {
|
||||
if stream == nil {
|
||||
return
|
||||
}
|
||||
stream.mu.Lock()
|
||||
terminalsFolder := strings.TrimSpace(stream.TerminalsFolder)
|
||||
stream.mu.Unlock()
|
||||
terminalSnapshot, ok := readTerminalShellFileSnapshot(terminalsFolder, shellID)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
now := time.Now().UTC()
|
||||
stream.mu.Lock()
|
||||
defer stream.mu.Unlock()
|
||||
state := ensureBackgroundShellStateLocked(stream, shellID, runtimecore.PendingExec{}, firstNonZeroTime(terminalSnapshot.StartedAt, now))
|
||||
if state == nil {
|
||||
return
|
||||
}
|
||||
if state.Command == "" {
|
||||
state.Command = terminalSnapshot.Command
|
||||
}
|
||||
if state.WorkingDirectory == "" {
|
||||
state.WorkingDirectory = terminalSnapshot.CWD
|
||||
}
|
||||
if state.PID == nil && terminalSnapshot.PID != nil {
|
||||
pid := *terminalSnapshot.PID
|
||||
state.PID = &pid
|
||||
}
|
||||
if state.CreatedAt.IsZero() {
|
||||
state.CreatedAt = firstNonZeroTime(terminalSnapshot.StartedAt, now)
|
||||
}
|
||||
if state.StdoutBuffer == "" && state.StderrBuffer == "" && terminalSnapshot.Output != "" {
|
||||
state.StdoutBuffer = terminalSnapshot.Output
|
||||
if !isBackgroundShellTerminalStatus(state.Status) {
|
||||
state.Status = backgroundShellStatusRunning
|
||||
}
|
||||
state.LastActivityAt = now
|
||||
}
|
||||
if !isBackgroundShellTerminalStatus(state.Status) && terminalSnapshot.ExitCode != nil {
|
||||
exitCode := *terminalSnapshot.ExitCode
|
||||
state.ExitCode = &exitCode
|
||||
state.Status = backgroundShellStatusCompleted
|
||||
state.CompletedAt = firstNonZeroTime(terminalSnapshot.EndedAt, now)
|
||||
state.LastActivityAt = state.CompletedAt
|
||||
}
|
||||
if strings.TrimSpace(state.Status) == "" {
|
||||
state.Status = backgroundShellStatusBackgrounded
|
||||
}
|
||||
if state.LastActivityAt.IsZero() {
|
||||
state.LastActivityAt = now
|
||||
}
|
||||
stream.UpdatedAt = now
|
||||
}
|
||||
|
||||
func readTerminalShellFileSnapshot(terminalsFolder string, shellID string) (terminalShellFileSnapshot, bool) {
|
||||
path, ok := terminalShellFilePath(terminalsFolder, shellID)
|
||||
if !ok {
|
||||
return terminalShellFileSnapshot{}, false
|
||||
}
|
||||
data, err := os.ReadFile(path)
|
||||
if err != nil || len(data) == 0 {
|
||||
return terminalShellFileSnapshot{}, false
|
||||
}
|
||||
return parseTerminalShellFileSnapshot(string(data)), true
|
||||
}
|
||||
|
||||
func terminalShellFilePath(terminalsFolder string, shellID string) (string, bool) {
|
||||
folder := filepath.Clean(strings.TrimSpace(terminalsFolder))
|
||||
if folder == "." || !filepath.IsAbs(folder) {
|
||||
return "", false
|
||||
}
|
||||
id, err := strconv.ParseUint(strings.TrimSpace(shellID), 10, 32)
|
||||
if err != nil {
|
||||
return "", false
|
||||
}
|
||||
filename := strconv.FormatUint(id, 10) + ".txt"
|
||||
return filepath.Join(folder, filename), true
|
||||
}
|
||||
|
||||
func parseTerminalShellFileSnapshot(raw string) terminalShellFileSnapshot {
|
||||
lines := strings.Split(strings.ReplaceAll(raw, "\r\n", "\n"), "\n")
|
||||
separators := terminalShellSeparatorIndexes(lines)
|
||||
snapshot := terminalShellFileSnapshot{}
|
||||
if len(separators) >= 2 {
|
||||
parseTerminalShellMetadata(lines[separators[0]+1:separators[1]], &snapshot)
|
||||
}
|
||||
outputStart := 0
|
||||
if len(separators) >= 2 {
|
||||
outputStart = separators[1] + 1
|
||||
}
|
||||
outputEnd := len(lines)
|
||||
for index := len(separators) - 2; index >= 1; index-- {
|
||||
block := lines[separators[index]+1 : separators[index+1]]
|
||||
if terminalMetadataBlockHasKey(block, "exit_code") || terminalMetadataBlockHasKey(block, "ended_at") {
|
||||
parseTerminalShellMetadata(block, &snapshot)
|
||||
outputEnd = separators[index]
|
||||
break
|
||||
}
|
||||
}
|
||||
if outputStart < outputEnd {
|
||||
snapshot.Output = strings.Join(lines[outputStart:outputEnd], "\n")
|
||||
snapshot.Output = strings.TrimSuffix(snapshot.Output, "\n")
|
||||
}
|
||||
return snapshot
|
||||
}
|
||||
|
||||
func terminalShellSeparatorIndexes(lines []string) []int {
|
||||
indexes := make([]int, 0, 4)
|
||||
for index, line := range lines {
|
||||
if strings.TrimSpace(line) == "---" {
|
||||
indexes = append(indexes, index)
|
||||
}
|
||||
}
|
||||
return indexes
|
||||
}
|
||||
|
||||
func parseTerminalShellMetadata(lines []string, snapshot *terminalShellFileSnapshot) {
|
||||
if snapshot == nil {
|
||||
return
|
||||
}
|
||||
for _, line := range lines {
|
||||
key, value, ok := strings.Cut(line, ":")
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
key = strings.TrimSpace(key)
|
||||
value = parseTerminalShellMetadataValue(value)
|
||||
switch key {
|
||||
case "pid":
|
||||
if parsed, err := strconv.ParseUint(value, 10, 32); err == nil {
|
||||
pid := uint32(parsed)
|
||||
snapshot.PID = &pid
|
||||
}
|
||||
case "cwd":
|
||||
snapshot.CWD = value
|
||||
case "command":
|
||||
snapshot.Command = value
|
||||
case "started_at":
|
||||
snapshot.StartedAt = parseTerminalShellTime(value)
|
||||
case "exit_code":
|
||||
if parsed, err := strconv.ParseInt(value, 10, 32); err == nil {
|
||||
exitCode := int32(parsed)
|
||||
snapshot.ExitCode = &exitCode
|
||||
}
|
||||
case "ended_at":
|
||||
snapshot.EndedAt = parseTerminalShellTime(value)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func parseTerminalShellMetadataValue(value string) string {
|
||||
trimmed := strings.TrimSpace(value)
|
||||
if unquoted, err := strconv.Unquote(trimmed); err == nil {
|
||||
return unquoted
|
||||
}
|
||||
return trimmed
|
||||
}
|
||||
|
||||
func parseTerminalShellTime(value string) time.Time {
|
||||
parsed, err := time.Parse(time.RFC3339Nano, strings.TrimSpace(value))
|
||||
if err != nil {
|
||||
return time.Time{}
|
||||
}
|
||||
return parsed.UTC()
|
||||
}
|
||||
|
||||
func terminalMetadataBlockHasKey(lines []string, key string) bool {
|
||||
for _, line := range lines {
|
||||
current, _, ok := strings.Cut(line, ":")
|
||||
if ok && strings.TrimSpace(current) == key {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func isBackgroundShellTerminalStatus(status string) bool {
|
||||
switch strings.TrimSpace(status) {
|
||||
case backgroundShellStatusCompleted, backgroundShellStatusRejected, backgroundShellStatusPermissionDenied, backgroundShellStatusTransportClosed:
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func (service *Service) observeBackgroundShellExecClientMessage(stream *ActiveStream, pending runtimecore.PendingExec, message *agentv1.ExecClientMessage) {
|
||||
if stream == nil || message == nil {
|
||||
return
|
||||
}
|
||||
now := time.Now().UTC()
|
||||
stream.mu.Lock()
|
||||
defer stream.mu.Unlock()
|
||||
service.observeBackgroundShellSpawnResultLocked(stream, pending, message.GetBackgroundShellSpawnResult(), now)
|
||||
service.observeForceBackgroundShellResultLocked(stream, pending, message.GetForceBackgroundShellResult(), now)
|
||||
}
|
||||
|
||||
func (service *Service) observeMissingBackgroundShellExecClientMessage(stream *ActiveStream, message *agentv1.ExecClientMessage) bool {
|
||||
if stream == nil || message == nil {
|
||||
return false
|
||||
}
|
||||
if message.GetBackgroundShellSpawnResult() == nil && message.GetForceBackgroundShellResult() == nil {
|
||||
return false
|
||||
}
|
||||
now := time.Now().UTC()
|
||||
stream.mu.Lock()
|
||||
defer stream.mu.Unlock()
|
||||
pending := runtimecore.PendingExec{
|
||||
MessageID: message.GetId(),
|
||||
ExecID: strings.TrimSpace(message.GetExecId()),
|
||||
ExecKind: "shell",
|
||||
}
|
||||
return service.observeBackgroundShellSpawnResultLocked(stream, pending, message.GetBackgroundShellSpawnResult(), now) || service.observeForceBackgroundShellResultLocked(stream, pending, message.GetForceBackgroundShellResult(), now)
|
||||
}
|
||||
|
||||
func (service *Service) observeBackgroundShellSpawnResultLocked(stream *ActiveStream, pending runtimecore.PendingExec, result *agentv1.BackgroundShellSpawnResult, now time.Time) bool {
|
||||
if stream == nil || result == nil {
|
||||
return false
|
||||
}
|
||||
if stream.BackgroundShells == nil {
|
||||
stream.BackgroundShells = make(map[string]*BackgroundShellState)
|
||||
}
|
||||
if stream.BackgroundShellsByMessageID == nil {
|
||||
stream.BackgroundShellsByMessageID = make(map[uint32]string)
|
||||
}
|
||||
if stream.BackgroundShellsByExecID == nil {
|
||||
stream.BackgroundShellsByExecID = make(map[string]string)
|
||||
}
|
||||
success := result.GetSuccess()
|
||||
if success != nil {
|
||||
shellID := strconv.FormatUint(uint64(success.GetShellId()), 10)
|
||||
state := stream.BackgroundShells[shellID]
|
||||
if state == nil {
|
||||
state = &BackgroundShellState{ShellID: shellID, CreatedAt: now}
|
||||
stream.BackgroundShells[shellID] = state
|
||||
}
|
||||
pid := success.Pid
|
||||
state.Command = firstNonEmpty(strings.TrimSpace(success.GetCommand()), state.Command)
|
||||
state.WorkingDirectory = firstNonEmpty(strings.TrimSpace(success.GetWorkingDirectory()), state.WorkingDirectory)
|
||||
state.PID = pid
|
||||
state.Status = backgroundShellStatusBackgrounded
|
||||
state.OriginalToolCallID = firstNonEmpty(strings.TrimSpace(pending.ToolCallID), state.OriginalToolCallID)
|
||||
state.OriginalExecID = firstNonEmpty(strings.TrimSpace(pending.ExecID), state.OriginalExecID)
|
||||
if state.OriginalMessageID == 0 {
|
||||
state.OriginalMessageID = pending.MessageID
|
||||
}
|
||||
if len(state.ArgsJSON) == 0 && len(pending.ArgsJSON) > 0 {
|
||||
state.ArgsJSON = append([]byte(nil), pending.ArgsJSON...)
|
||||
}
|
||||
state.ModelCallID = firstNonEmpty(strings.TrimSpace(pending.ModelCallID), state.ModelCallID)
|
||||
state.LastActivityAt = now
|
||||
if pending.MessageID != 0 {
|
||||
stream.BackgroundShellsByMessageID[pending.MessageID] = shellID
|
||||
}
|
||||
if strings.TrimSpace(pending.ExecID) != "" {
|
||||
stream.BackgroundShellsByExecID[strings.TrimSpace(pending.ExecID)] = shellID
|
||||
}
|
||||
stream.UpdatedAt = now
|
||||
return true
|
||||
}
|
||||
shellID := backgroundShellIDForMessageLocked(stream, pending.MessageID, pending.ExecID)
|
||||
if shellID == "" {
|
||||
return false
|
||||
}
|
||||
state := ensureBackgroundShellStateLocked(stream, shellID, pending, now)
|
||||
if state == nil {
|
||||
return false
|
||||
}
|
||||
switch {
|
||||
case result.GetRejected() != nil:
|
||||
state.Status = backgroundShellStatusRejected
|
||||
state.StderrBuffer += strings.TrimSpace(result.GetRejected().GetReason())
|
||||
case result.GetPermissionDenied() != nil:
|
||||
state.Status = backgroundShellStatusPermissionDenied
|
||||
state.StderrBuffer += strings.TrimSpace(result.GetPermissionDenied().GetError())
|
||||
case result.GetError() != nil:
|
||||
state.Status = backgroundShellStatusTransportClosed
|
||||
state.StderrBuffer += strings.TrimSpace(result.GetError().GetError())
|
||||
default:
|
||||
return false
|
||||
}
|
||||
state.LastActivityAt = now
|
||||
state.CompletedAt = now
|
||||
stream.UpdatedAt = now
|
||||
return true
|
||||
}
|
||||
|
||||
func (service *Service) observeForceBackgroundShellResultLocked(stream *ActiveStream, pending runtimecore.PendingExec, result *agentv1.ForceBackgroundShellResult, now time.Time) bool {
|
||||
if stream == nil || result == nil || result.GetShellResult() == nil {
|
||||
return false
|
||||
}
|
||||
shellResult := result.GetShellResult()
|
||||
shellIDValue := uint32(0)
|
||||
if success := shellResult.GetSuccess(); success != nil {
|
||||
shellIDValue = success.GetShellId()
|
||||
}
|
||||
if shellIDValue == 0 {
|
||||
return false
|
||||
}
|
||||
shellID := strconv.FormatUint(uint64(shellIDValue), 10)
|
||||
state := ensureBackgroundShellStateLocked(stream, shellID, pending, now)
|
||||
if state == nil {
|
||||
return false
|
||||
}
|
||||
state.Status = backgroundShellStatusBackgrounded
|
||||
state.LastActivityAt = now
|
||||
stream.UpdatedAt = now
|
||||
return true
|
||||
}
|
||||
|
||||
func (service *Service) observeShellExecClientMessage(stream *ActiveStream, pending runtimecore.PendingExec, message *agentv1.ExecClientMessage) {
|
||||
if stream == nil || message == nil || strings.TrimSpace(pending.ExecKind) != "shell" {
|
||||
return
|
||||
}
|
||||
shellStream := message.GetShellStream()
|
||||
if shellStream == nil {
|
||||
return
|
||||
}
|
||||
now := time.Now().UTC()
|
||||
stream.mu.Lock()
|
||||
defer stream.mu.Unlock()
|
||||
service.observeShellStreamLocked(stream, pending, shellStream, now)
|
||||
}
|
||||
|
||||
func (service *Service) observeMissingShellExecClientMessage(stream *ActiveStream, message *agentv1.ExecClientMessage) bool {
|
||||
if stream == nil || message == nil || message.GetShellStream() == nil {
|
||||
return false
|
||||
}
|
||||
now := time.Now().UTC()
|
||||
stream.mu.Lock()
|
||||
defer stream.mu.Unlock()
|
||||
shellID := backgroundShellIDForMessageLocked(stream, message.GetId(), message.GetExecId())
|
||||
if shellID == "" {
|
||||
return false
|
||||
}
|
||||
pending := runtimecore.PendingExec{
|
||||
MessageID: message.GetId(),
|
||||
ExecID: strings.TrimSpace(message.GetExecId()),
|
||||
ExecKind: "shell",
|
||||
}
|
||||
if state := stream.BackgroundShells[shellID]; state != nil {
|
||||
pending.ToolCallID = state.OriginalToolCallID
|
||||
pending.ArgsJSON = append([]byte(nil), state.ArgsJSON...)
|
||||
pending.ModelCallID = state.ModelCallID
|
||||
}
|
||||
service.observeShellStreamLocked(stream, pending, message.GetShellStream(), now)
|
||||
return true
|
||||
}
|
||||
|
||||
func (service *Service) observeShellStreamLocked(stream *ActiveStream, pending runtimecore.PendingExec, shellStream *agentv1.ShellStream, now time.Time) {
|
||||
if stream.BackgroundShells == nil {
|
||||
stream.BackgroundShells = make(map[string]*BackgroundShellState)
|
||||
}
|
||||
if stream.BackgroundShellsByMessageID == nil {
|
||||
stream.BackgroundShellsByMessageID = make(map[uint32]string)
|
||||
}
|
||||
if stream.BackgroundShellsByExecID == nil {
|
||||
stream.BackgroundShellsByExecID = make(map[string]string)
|
||||
}
|
||||
|
||||
shellID := backgroundShellIDForMessageLocked(stream, pending.MessageID, pending.ExecID)
|
||||
switch event := shellStream.GetEvent().(type) {
|
||||
case *agentv1.ShellStream_Backgrounded:
|
||||
shellID = strconv.FormatUint(uint64(event.Backgrounded.GetShellId()), 10)
|
||||
state := stream.BackgroundShells[shellID]
|
||||
if state == nil {
|
||||
state = &BackgroundShellState{ShellID: shellID, CreatedAt: now}
|
||||
stream.BackgroundShells[shellID] = state
|
||||
}
|
||||
state.Command = firstNonEmpty(strings.TrimSpace(event.Backgrounded.GetCommand()), state.Command)
|
||||
state.WorkingDirectory = firstNonEmpty(strings.TrimSpace(event.Backgrounded.GetWorkingDirectory()), state.WorkingDirectory)
|
||||
state.PID = event.Backgrounded.Pid
|
||||
state.Status = backgroundShellStatusBackgrounded
|
||||
state.OriginalToolCallID = strings.TrimSpace(pending.ToolCallID)
|
||||
state.OriginalExecID = strings.TrimSpace(pending.ExecID)
|
||||
state.OriginalMessageID = pending.MessageID
|
||||
state.ArgsJSON = append([]byte(nil), pending.ArgsJSON...)
|
||||
state.ModelCallID = strings.TrimSpace(pending.ModelCallID)
|
||||
state.LastActivityAt = now
|
||||
if pending.MessageID != 0 {
|
||||
stream.BackgroundShellsByMessageID[pending.MessageID] = shellID
|
||||
}
|
||||
if strings.TrimSpace(pending.ExecID) != "" {
|
||||
stream.BackgroundShellsByExecID[strings.TrimSpace(pending.ExecID)] = shellID
|
||||
}
|
||||
case *agentv1.ShellStream_Stdout:
|
||||
state := ensureBackgroundShellStateLocked(stream, shellID, pending, now)
|
||||
if state == nil {
|
||||
return
|
||||
}
|
||||
state.StdoutBuffer += execbridge.DecodeShellStdout(event.Stdout)
|
||||
state.Status = backgroundShellStatusRunning
|
||||
state.LastActivityAt = now
|
||||
case *agentv1.ShellStream_Stderr:
|
||||
state := ensureBackgroundShellStateLocked(stream, shellID, pending, now)
|
||||
if state == nil {
|
||||
return
|
||||
}
|
||||
state.StderrBuffer += event.Stderr.GetData()
|
||||
state.Status = backgroundShellStatusRunning
|
||||
state.LastActivityAt = now
|
||||
case *agentv1.ShellStream_Exit:
|
||||
state := ensureBackgroundShellStateLocked(stream, shellID, pending, now)
|
||||
if state == nil {
|
||||
return
|
||||
}
|
||||
exitCode := int32(event.Exit.GetCode())
|
||||
state.ExitCode = &exitCode
|
||||
state.WorkingDirectory = firstNonEmpty(strings.TrimSpace(event.Exit.GetCwd()), state.WorkingDirectory)
|
||||
state.Status = backgroundShellStatusCompleted
|
||||
state.LastActivityAt = now
|
||||
state.CompletedAt = now
|
||||
case *agentv1.ShellStream_Rejected:
|
||||
state := ensureBackgroundShellStateLocked(stream, shellID, pending, now)
|
||||
if state == nil {
|
||||
return
|
||||
}
|
||||
state.Status = backgroundShellStatusRejected
|
||||
state.StderrBuffer += strings.TrimSpace(event.Rejected.GetReason())
|
||||
state.LastActivityAt = now
|
||||
state.CompletedAt = now
|
||||
case *agentv1.ShellStream_PermissionDenied:
|
||||
state := ensureBackgroundShellStateLocked(stream, shellID, pending, now)
|
||||
if state == nil {
|
||||
return
|
||||
}
|
||||
state.Status = backgroundShellStatusPermissionDenied
|
||||
state.StderrBuffer += strings.TrimSpace(event.PermissionDenied.GetError())
|
||||
state.LastActivityAt = now
|
||||
state.CompletedAt = now
|
||||
}
|
||||
stream.UpdatedAt = now
|
||||
}
|
||||
|
||||
func observeBackgroundTaskCompletionAction(stream *ActiveStream, message *agentv1.AgentClientMessage) {
|
||||
if stream == nil || message == nil {
|
||||
return
|
||||
}
|
||||
action := message.GetConversationAction()
|
||||
if action == nil {
|
||||
return
|
||||
}
|
||||
item, ok := action.GetAction().(*agentv1.ConversationAction_BackgroundTaskCompletionAction)
|
||||
if !ok || item.BackgroundTaskCompletionAction == nil {
|
||||
return
|
||||
}
|
||||
now := time.Now().UTC()
|
||||
stream.mu.Lock()
|
||||
defer stream.mu.Unlock()
|
||||
for _, completion := range item.BackgroundTaskCompletionAction.GetCompletions() {
|
||||
observeBackgroundTaskCompletionLocked(stream, completion, now)
|
||||
}
|
||||
}
|
||||
|
||||
func observeBackgroundTaskCompletionLocked(stream *ActiveStream, completion *agentv1.BackgroundTaskCompletion, now time.Time) bool {
|
||||
if stream == nil || completion == nil || completion.GetKind() != agentv1.BackgroundTaskKind_BACKGROUND_TASK_KIND_SHELL {
|
||||
return false
|
||||
}
|
||||
shellID := strings.TrimSpace(completion.GetTaskId())
|
||||
if shellID == "" {
|
||||
return false
|
||||
}
|
||||
state := ensureBackgroundShellStateLocked(stream, shellID, runtimecore.PendingExec{}, now)
|
||||
if state == nil {
|
||||
return false
|
||||
}
|
||||
detail := strings.TrimSpace(completion.GetDetail())
|
||||
switch completion.GetStatus() {
|
||||
case agentv1.BackgroundTaskStatus_BACKGROUND_TASK_STATUS_SUCCESS:
|
||||
state.Status = backgroundShellStatusCompleted
|
||||
if detail != "" {
|
||||
state.StdoutBuffer = appendBackgroundShellBuffer(state.StdoutBuffer, detail)
|
||||
}
|
||||
case agentv1.BackgroundTaskStatus_BACKGROUND_TASK_STATUS_ERROR:
|
||||
state.Status = backgroundShellStatusTransportClosed
|
||||
if detail != "" {
|
||||
state.StderrBuffer = appendBackgroundShellBuffer(state.StderrBuffer, detail)
|
||||
}
|
||||
case agentv1.BackgroundTaskStatus_BACKGROUND_TASK_STATUS_ABORTED:
|
||||
state.Status = backgroundShellStatusTransportClosed
|
||||
if detail != "" {
|
||||
state.StderrBuffer = appendBackgroundShellBuffer(state.StderrBuffer, detail)
|
||||
}
|
||||
default:
|
||||
if completion.GetReason() == agentv1.BackgroundTaskCompletionReason_BACKGROUND_TASK_COMPLETION_REASON_TASK_PROGRESS {
|
||||
state.Status = backgroundShellStatusRunning
|
||||
} else {
|
||||
return false
|
||||
}
|
||||
}
|
||||
state.LastActivityAt = now
|
||||
if completion.GetReason() == agentv1.BackgroundTaskCompletionReason_BACKGROUND_TASK_COMPLETION_REASON_TASK_FINISHED || isBackgroundShellTerminalStatus(state.Status) {
|
||||
state.CompletedAt = now
|
||||
}
|
||||
stream.UpdatedAt = now
|
||||
return true
|
||||
}
|
||||
|
||||
func observeBackgroundShellAction(stream *ActiveStream, message *agentv1.AgentClientMessage) (string, bool) {
|
||||
if stream == nil || message == nil {
|
||||
return "", false
|
||||
}
|
||||
action := message.GetConversationAction()
|
||||
if action == nil {
|
||||
return "", false
|
||||
}
|
||||
item, ok := action.GetAction().(*agentv1.ConversationAction_BackgroundShellAction)
|
||||
if !ok || item.BackgroundShellAction == nil {
|
||||
return "", false
|
||||
}
|
||||
return recordBackgroundShellActionMemory(stream, item.BackgroundShellAction.GetToolCallId(), time.Now().UTC())
|
||||
}
|
||||
|
||||
func recordBackgroundShellActionMemory(stream *ActiveStream, toolCallID string, now time.Time) (string, bool) {
|
||||
trimmedToolCallID := strings.TrimSpace(toolCallID)
|
||||
if stream == nil || trimmedToolCallID == "" {
|
||||
return "", false
|
||||
}
|
||||
stream.mu.Lock()
|
||||
defer stream.mu.Unlock()
|
||||
return recordBackgroundShellActionLocked(stream, trimmedToolCallID, now)
|
||||
}
|
||||
|
||||
func recordBackgroundShellActionLocked(stream *ActiveStream, toolCallID string, now time.Time) (string, bool) {
|
||||
trimmedToolCallID := strings.TrimSpace(toolCallID)
|
||||
if stream == nil || trimmedToolCallID == "" {
|
||||
return "", false
|
||||
}
|
||||
if stream.BackgroundShellActions == nil {
|
||||
stream.BackgroundShellActions = make(map[string]time.Time)
|
||||
}
|
||||
if _, exists := stream.BackgroundShellActions[trimmedToolCallID]; exists {
|
||||
return trimmedToolCallID, false
|
||||
}
|
||||
stream.BackgroundShellActions[trimmedToolCallID] = now
|
||||
stream.UpdatedAt = now
|
||||
return trimmedToolCallID, true
|
||||
}
|
||||
|
||||
func newBackgroundShellActionMetadataEntry(turnSeq int64, requestID string, toolCallID string, source string) HistoryEntry {
|
||||
return newMetadataEntry(turnSeq, requestID, "background_shell_action", map[string]any{
|
||||
"tool_call_id": strings.TrimSpace(toolCallID),
|
||||
"source": strings.TrimSpace(source),
|
||||
})
|
||||
}
|
||||
|
||||
func backgroundTaskCompletionMetadataEntries(turnSeq int64, requestID string, message *agentv1.AgentClientMessage) []HistoryEntry {
|
||||
if message == nil || message.GetConversationAction() == nil {
|
||||
return nil
|
||||
}
|
||||
item, ok := message.GetConversationAction().GetAction().(*agentv1.ConversationAction_BackgroundTaskCompletionAction)
|
||||
if !ok || item.BackgroundTaskCompletionAction == nil {
|
||||
return nil
|
||||
}
|
||||
completions := item.BackgroundTaskCompletionAction.GetCompletions()
|
||||
if len(completions) == 0 {
|
||||
return nil
|
||||
}
|
||||
entries := make([]HistoryEntry, 0, len(completions))
|
||||
for _, completion := range completions {
|
||||
if completion == nil {
|
||||
continue
|
||||
}
|
||||
values := map[string]any{
|
||||
"task_id": strings.TrimSpace(completion.GetTaskId()),
|
||||
"kind": completion.GetKind().String(),
|
||||
"status": completion.GetStatus().String(),
|
||||
"reason": completion.GetReason().String(),
|
||||
}
|
||||
if title := strings.TrimSpace(completion.GetTitle()); title != "" {
|
||||
values["title"] = title
|
||||
}
|
||||
if detail := strings.TrimSpace(completion.GetDetail()); detail != "" {
|
||||
values["detail"] = detail
|
||||
}
|
||||
if outputPath := strings.TrimSpace(completion.GetOutputPath()); outputPath != "" {
|
||||
values["output_path"] = outputPath
|
||||
}
|
||||
if threadID := strings.TrimSpace(completion.GetThreadId()); threadID != "" {
|
||||
values["thread_id"] = threadID
|
||||
}
|
||||
entries = append(entries, newMetadataEntry(turnSeq, requestID, "background_task_completion_action", values))
|
||||
}
|
||||
return entries
|
||||
}
|
||||
|
||||
func shellToolCallIsBackgrounded(toolCall *agentv1.ToolCall) bool {
|
||||
if toolCall == nil {
|
||||
return false
|
||||
}
|
||||
shellToolCall := toolCall.GetShellToolCall()
|
||||
if shellToolCall == nil || shellToolCall.GetResult() == nil {
|
||||
return false
|
||||
}
|
||||
return shellToolCall.GetResult().GetIsBackground()
|
||||
}
|
||||
|
||||
func appendBackgroundShellBuffer(current string, value string) string {
|
||||
trimmed := strings.TrimSpace(value)
|
||||
if trimmed == "" {
|
||||
return current
|
||||
}
|
||||
if current == "" || strings.HasSuffix(current, "\n") {
|
||||
return current + trimmed
|
||||
}
|
||||
return current + "\n" + trimmed
|
||||
}
|
||||
|
||||
func ensureBackgroundShellStateLocked(stream *ActiveStream, shellID string, pending runtimecore.PendingExec, now time.Time) *BackgroundShellState {
|
||||
trimmedShellID := strings.TrimSpace(shellID)
|
||||
if trimmedShellID == "" {
|
||||
return nil
|
||||
}
|
||||
if stream.BackgroundShells == nil {
|
||||
stream.BackgroundShells = make(map[string]*BackgroundShellState)
|
||||
}
|
||||
if stream.BackgroundShellsByMessageID == nil {
|
||||
stream.BackgroundShellsByMessageID = make(map[uint32]string)
|
||||
}
|
||||
if stream.BackgroundShellsByExecID == nil {
|
||||
stream.BackgroundShellsByExecID = make(map[string]string)
|
||||
}
|
||||
state := stream.BackgroundShells[trimmedShellID]
|
||||
if state == nil {
|
||||
state = &BackgroundShellState{ShellID: trimmedShellID, Status: backgroundShellStatusRunning, CreatedAt: now}
|
||||
stream.BackgroundShells[trimmedShellID] = state
|
||||
}
|
||||
if pending.MessageID != 0 {
|
||||
stream.BackgroundShellsByMessageID[pending.MessageID] = trimmedShellID
|
||||
}
|
||||
if strings.TrimSpace(pending.ExecID) != "" {
|
||||
stream.BackgroundShellsByExecID[strings.TrimSpace(pending.ExecID)] = trimmedShellID
|
||||
}
|
||||
if state.OriginalToolCallID == "" {
|
||||
state.OriginalToolCallID = strings.TrimSpace(pending.ToolCallID)
|
||||
}
|
||||
if state.OriginalExecID == "" {
|
||||
state.OriginalExecID = strings.TrimSpace(pending.ExecID)
|
||||
}
|
||||
if state.OriginalMessageID == 0 {
|
||||
state.OriginalMessageID = pending.MessageID
|
||||
}
|
||||
if len(state.ArgsJSON) == 0 && len(pending.ArgsJSON) > 0 {
|
||||
state.ArgsJSON = append([]byte(nil), pending.ArgsJSON...)
|
||||
}
|
||||
if state.ModelCallID == "" {
|
||||
state.ModelCallID = strings.TrimSpace(pending.ModelCallID)
|
||||
}
|
||||
return state
|
||||
}
|
||||
|
||||
func backgroundShellIDForMessageLocked(stream *ActiveStream, messageID uint32, execID string) string {
|
||||
if stream == nil {
|
||||
return ""
|
||||
}
|
||||
if strings.TrimSpace(execID) != "" {
|
||||
if shellID := strings.TrimSpace(stream.BackgroundShellsByExecID[strings.TrimSpace(execID)]); shellID != "" {
|
||||
return shellID
|
||||
}
|
||||
}
|
||||
if messageID != 0 {
|
||||
return strings.TrimSpace(stream.BackgroundShellsByMessageID[messageID])
|
||||
}
|
||||
return ""
|
||||
}
|
||||
@@ -0,0 +1,442 @@
|
||||
// broker.go 负责 request 维度活动流的订阅、广播、取消和终态收口。
|
||||
package forwarder
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"strings"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
"cursor/gen/agentv1"
|
||||
runtimecore "cursor/internal/backend/agent/core"
|
||||
)
|
||||
|
||||
const subscriberSignalBufferSize = 1
|
||||
const orphanSubscriberGracePeriod = 30 * time.Second
|
||||
const terminalStreamRetentionPeriod = 30 * time.Second
|
||||
|
||||
type StreamBroker struct {
|
||||
mu sync.RWMutex
|
||||
streams map[string]*ActiveStream
|
||||
nextID atomic.Uint64
|
||||
}
|
||||
|
||||
// NewStreamBroker 创建活动流注册表。
|
||||
func NewStreamBroker() *StreamBroker {
|
||||
return &StreamBroker{
|
||||
streams: make(map[string]*ActiveStream),
|
||||
}
|
||||
}
|
||||
|
||||
// OpenStream 打开或复用指定 request 的活动流,并刷新其最新上下文。
|
||||
func (broker *StreamBroker) OpenStream(requestID string, conversationID string, turnSeq int64, modelID string, modelName string, mode agentv1.AgentMode, latestUserText string) (*ActiveStream, error) {
|
||||
normalizedRequestID := strings.TrimSpace(requestID)
|
||||
if normalizedRequestID == "" {
|
||||
return nil, nil
|
||||
}
|
||||
normalizedMode, err := validateSupportedActiveMode(mode)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
broker.mu.Lock()
|
||||
defer broker.mu.Unlock()
|
||||
if existing, ok := broker.streams[normalizedRequestID]; ok {
|
||||
existing.mu.Lock()
|
||||
existing.ConversationID = strings.TrimSpace(conversationID)
|
||||
existing.TurnSeq = turnSeq
|
||||
existing.ModelID = strings.TrimSpace(modelID)
|
||||
existing.ModelName = strings.TrimSpace(modelName)
|
||||
existing.Mode = normalizedMode
|
||||
existing.LatestUserText = strings.TrimSpace(latestUserText)
|
||||
if existing.Status == "" {
|
||||
existing.Status = StreamStatusCreated
|
||||
}
|
||||
if existing.PendingExecs == nil {
|
||||
existing.PendingExecs = make(map[string]runtimecore.PendingExec)
|
||||
}
|
||||
if existing.PendingInteractions == nil {
|
||||
existing.PendingInteractions = make(map[string]runtimecore.PendingInteraction)
|
||||
}
|
||||
if existing.PartialToolCallIDs == nil {
|
||||
existing.PartialToolCallIDs = make(map[string]struct{})
|
||||
}
|
||||
if existing.PatchEditQueues == nil {
|
||||
existing.PatchEditQueues = make(map[string][]queuedPatchEditOperation)
|
||||
}
|
||||
if existing.BackgroundShells == nil {
|
||||
existing.BackgroundShells = make(map[string]*BackgroundShellState)
|
||||
}
|
||||
if existing.BackgroundShellsByMessageID == nil {
|
||||
existing.BackgroundShellsByMessageID = make(map[uint32]string)
|
||||
}
|
||||
if existing.BackgroundShellsByExecID == nil {
|
||||
existing.BackgroundShellsByExecID = make(map[string]string)
|
||||
}
|
||||
if existing.BackgroundShellActions == nil {
|
||||
existing.BackgroundShellActions = make(map[string]time.Time)
|
||||
}
|
||||
existing.UpdatedAt = time.Now().UTC()
|
||||
existing.mu.Unlock()
|
||||
return existing, nil
|
||||
}
|
||||
now := time.Now().UTC()
|
||||
stream := &ActiveStream{
|
||||
RequestID: normalizedRequestID,
|
||||
ConversationID: strings.TrimSpace(conversationID),
|
||||
TurnSeq: turnSeq,
|
||||
ModelID: strings.TrimSpace(modelID),
|
||||
ModelName: strings.TrimSpace(modelName),
|
||||
Mode: normalizedMode,
|
||||
LatestUserText: strings.TrimSpace(latestUserText),
|
||||
Status: StreamStatusCreated,
|
||||
Backlog: make([]StreamEvent, 0, 64),
|
||||
Subscribers: make(map[string]*StreamSubscriber),
|
||||
PendingExecs: make(map[string]runtimecore.PendingExec),
|
||||
PendingInteractions: make(map[string]runtimecore.PendingInteraction),
|
||||
PartialToolCallIDs: make(map[string]struct{}),
|
||||
PatchEditQueues: make(map[string][]queuedPatchEditOperation),
|
||||
MCPToolServers: make(map[string]string),
|
||||
RecentCompletedExecs: make(map[uint32]time.Time),
|
||||
BackgroundShells: make(map[string]*BackgroundShellState),
|
||||
BackgroundShellsByMessageID: make(map[uint32]string),
|
||||
BackgroundShellsByExecID: make(map[string]string),
|
||||
BackgroundShellActions: make(map[string]time.Time),
|
||||
CreatedAt: now,
|
||||
UpdatedAt: now,
|
||||
}
|
||||
broker.streams[normalizedRequestID] = stream
|
||||
return stream, nil
|
||||
}
|
||||
|
||||
// Get 返回指定 request 对应的活动流句柄。
|
||||
func (broker *StreamBroker) Get(requestID string) (*ActiveStream, bool) {
|
||||
if broker == nil {
|
||||
return nil, false
|
||||
}
|
||||
broker.mu.RLock()
|
||||
defer broker.mu.RUnlock()
|
||||
stream, ok := broker.streams[strings.TrimSpace(requestID)]
|
||||
return stream, ok
|
||||
}
|
||||
|
||||
// Subscribe 为指定 request 注册一个新订阅者,并返回用于唤醒 backlog 消费的信号通道。
|
||||
func (broker *StreamBroker) Subscribe(requestID string) (string, <-chan struct{}, error) {
|
||||
normalizedRequestID := strings.TrimSpace(requestID)
|
||||
if normalizedRequestID == "" {
|
||||
return "", nil, fmt.Errorf("request_id is required")
|
||||
}
|
||||
stream, ok := broker.Get(normalizedRequestID)
|
||||
if !ok || stream == nil {
|
||||
// RunSSE 可能先于 BidiAppend 到达。此时先创建一个占位活动流,
|
||||
// 等待后续上行把真实 conversation/model/mode 信息补齐。
|
||||
var err error
|
||||
stream, err = broker.OpenStream(normalizedRequestID, "", 0, "", "", agentv1.AgentMode_AGENT_MODE_AGENT, "")
|
||||
if err != nil {
|
||||
return "", nil, err
|
||||
}
|
||||
}
|
||||
subscriberID := fmt.Sprintf("sub-%d", broker.nextID.Add(1))
|
||||
subscriber := &StreamSubscriber{Signal: make(chan struct{}, subscriberSignalBufferSize)}
|
||||
|
||||
stream.mu.Lock()
|
||||
broker.stopTerminalCleanupTimerLocked(stream)
|
||||
stream.Subscribers[subscriberID] = subscriber
|
||||
stream.UpdatedAt = time.Now().UTC()
|
||||
stream.mu.Unlock()
|
||||
|
||||
return subscriberID, subscriber.Signal, nil
|
||||
}
|
||||
|
||||
func (broker *StreamBroker) stopTerminalCleanupTimerLocked(stream *ActiveStream) {
|
||||
if stream == nil {
|
||||
return
|
||||
}
|
||||
stream.TerminalCleanupSeq.Add(1)
|
||||
if stream.TerminalCleanupTimer != nil {
|
||||
stream.TerminalCleanupTimer.Stop()
|
||||
stream.TerminalCleanupTimer = nil
|
||||
}
|
||||
}
|
||||
|
||||
// Unsubscribe 移除并关闭指定订阅者,并返回移除后的剩余订阅者数量。
|
||||
func (broker *StreamBroker) Unsubscribe(requestID string, subscriberID string) int {
|
||||
stream, ok := broker.Get(requestID)
|
||||
if !ok || stream == nil {
|
||||
return 0
|
||||
}
|
||||
remaining := 0
|
||||
stream.mu.Lock()
|
||||
if _, ok := stream.Subscribers[strings.TrimSpace(subscriberID)]; ok {
|
||||
delete(stream.Subscribers, strings.TrimSpace(subscriberID))
|
||||
}
|
||||
remaining = len(stream.Subscribers)
|
||||
stream.mu.Unlock()
|
||||
return remaining
|
||||
}
|
||||
|
||||
func (broker *StreamBroker) OtherConversationRequestIDs(conversationID string, keepRequestID string) []string {
|
||||
normalizedConversationID := strings.TrimSpace(conversationID)
|
||||
normalizedKeepRequestID := strings.TrimSpace(keepRequestID)
|
||||
if normalizedConversationID == "" {
|
||||
return nil
|
||||
}
|
||||
type requestStream struct {
|
||||
requestID string
|
||||
stream *ActiveStream
|
||||
}
|
||||
candidates := make([]requestStream, 0, 2)
|
||||
broker.mu.RLock()
|
||||
for requestID, stream := range broker.streams {
|
||||
if stream == nil || strings.TrimSpace(requestID) == normalizedKeepRequestID {
|
||||
continue
|
||||
}
|
||||
candidates = append(candidates, requestStream{
|
||||
requestID: requestID,
|
||||
stream: stream,
|
||||
})
|
||||
}
|
||||
broker.mu.RUnlock()
|
||||
requestIDs := make([]string, 0, 2)
|
||||
for _, candidate := range candidates {
|
||||
stream := candidate.stream
|
||||
stream.mu.Lock()
|
||||
sameConversation := strings.TrimSpace(stream.ConversationID) == normalizedConversationID
|
||||
status := stream.Status
|
||||
phase := stream.Phase
|
||||
stream.mu.Unlock()
|
||||
terminalPhase := phase == TurnPhaseCanceled || phase == TurnPhaseCompleted || phase == TurnPhaseFailed
|
||||
if !sameConversation || isTerminalStreamStatus(status) || terminalPhase {
|
||||
continue
|
||||
}
|
||||
requestIDs = append(requestIDs, candidate.requestID)
|
||||
}
|
||||
return requestIDs
|
||||
}
|
||||
|
||||
func (broker *StreamBroker) scheduleTerminalCleanup(requestID string) bool {
|
||||
stream, ok := broker.Get(requestID)
|
||||
if !ok || stream == nil {
|
||||
return false
|
||||
}
|
||||
stream.mu.Lock()
|
||||
defer stream.mu.Unlock()
|
||||
if len(stream.Subscribers) > 0 {
|
||||
broker.stopTerminalCleanupTimerLocked(stream)
|
||||
return false
|
||||
}
|
||||
if stream.Status != StreamStatusCompleted && stream.Status != StreamStatusCanceled && stream.Status != StreamStatusFailed {
|
||||
broker.stopTerminalCleanupTimerLocked(stream)
|
||||
return false
|
||||
}
|
||||
sequence := stream.TerminalCleanupSeq.Add(1)
|
||||
if stream.TerminalCleanupTimer != nil {
|
||||
stream.TerminalCleanupTimer.Stop()
|
||||
}
|
||||
stream.TerminalCleanupTimer = time.AfterFunc(terminalStreamRetentionPeriod, func() {
|
||||
broker.runScheduledTerminalCleanup(requestID, sequence)
|
||||
})
|
||||
stream.UpdatedAt = time.Now().UTC()
|
||||
return true
|
||||
}
|
||||
|
||||
func (broker *StreamBroker) runScheduledTerminalCleanup(requestID string, sequence uint64) {
|
||||
stream, ok := broker.Get(requestID)
|
||||
if !ok || stream == nil {
|
||||
return
|
||||
}
|
||||
stream.mu.Lock()
|
||||
if stream.TerminalCleanupSeq.Load() != sequence {
|
||||
stream.mu.Unlock()
|
||||
return
|
||||
}
|
||||
stream.TerminalCleanupTimer = nil
|
||||
if len(stream.Subscribers) > 0 {
|
||||
stream.mu.Unlock()
|
||||
return
|
||||
}
|
||||
if stream.Status != StreamStatusCompleted && stream.Status != StreamStatusCanceled && stream.Status != StreamStatusFailed {
|
||||
stream.mu.Unlock()
|
||||
return
|
||||
}
|
||||
stream.mu.Unlock()
|
||||
broker.RemoveIfIdle(requestID)
|
||||
}
|
||||
|
||||
// RemoveIfIdle 在没有订阅者时移除终态流,或移除仍为空壳的占位流。
|
||||
func (broker *StreamBroker) RemoveIfIdle(requestID string) bool {
|
||||
normalizedRequestID := strings.TrimSpace(requestID)
|
||||
if normalizedRequestID == "" {
|
||||
return false
|
||||
}
|
||||
broker.mu.Lock()
|
||||
defer broker.mu.Unlock()
|
||||
stream, ok := broker.streams[normalizedRequestID]
|
||||
if !ok || stream == nil {
|
||||
return false
|
||||
}
|
||||
stream.mu.Lock()
|
||||
subscriberCount := len(stream.Subscribers)
|
||||
isActive := stream.ProviderActive
|
||||
hasBacklog := len(stream.Backlog) > 0
|
||||
hasConversation := strings.TrimSpace(stream.ConversationID) != ""
|
||||
status := stream.Status
|
||||
if status == StreamStatusCompleted || status == StreamStatusCanceled || status == StreamStatusFailed {
|
||||
broker.stopTerminalCleanupTimerLocked(stream)
|
||||
}
|
||||
stream.mu.Unlock()
|
||||
if subscriberCount > 0 {
|
||||
return false
|
||||
}
|
||||
if status == StreamStatusCompleted || status == StreamStatusCanceled || status == StreamStatusFailed {
|
||||
delete(broker.streams, normalizedRequestID)
|
||||
return true
|
||||
}
|
||||
if isActive || hasBacklog || hasConversation {
|
||||
return false
|
||||
}
|
||||
delete(broker.streams, normalizedRequestID)
|
||||
return true
|
||||
}
|
||||
|
||||
// Publish 把一个事件写入 backlog,并唤醒当前所有订阅者读取 backlog。
|
||||
func (broker *StreamBroker) Publish(requestID string, event StreamEvent) error {
|
||||
stream, ok := broker.Get(requestID)
|
||||
if !ok || stream == nil {
|
||||
return fmt.Errorf("request is not active: %s", strings.TrimSpace(requestID))
|
||||
}
|
||||
stream.mu.Lock()
|
||||
if !event.End && isTerminalStreamStatus(stream.Status) {
|
||||
stream.mu.Unlock()
|
||||
return nil
|
||||
}
|
||||
stream.Backlog = append(stream.Backlog, event)
|
||||
stream.UpdatedAt = time.Now().UTC()
|
||||
subscribers := make([]*StreamSubscriber, 0, len(stream.Subscribers))
|
||||
for _, subscriber := range stream.Subscribers {
|
||||
subscribers = append(subscribers, subscriber)
|
||||
}
|
||||
stream.mu.Unlock()
|
||||
|
||||
for _, subscriber := range subscribers {
|
||||
if subscriber == nil {
|
||||
continue
|
||||
}
|
||||
select {
|
||||
case subscriber.Signal <- struct{}{}:
|
||||
default:
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// ReadFromCursor 返回从 cursor 开始尚未消费的 backlog 事件副本。
|
||||
func (broker *StreamBroker) ReadFromCursor(requestID string, cursor int) ([]StreamEvent, error) {
|
||||
stream, ok := broker.Get(requestID)
|
||||
if !ok || stream == nil {
|
||||
return nil, fmt.Errorf("request is not active: %s", strings.TrimSpace(requestID))
|
||||
}
|
||||
stream.mu.Lock()
|
||||
defer stream.mu.Unlock()
|
||||
if cursor < 0 {
|
||||
cursor = 0
|
||||
}
|
||||
if cursor >= len(stream.Backlog) {
|
||||
return nil, nil
|
||||
}
|
||||
return append([]StreamEvent(nil), stream.Backlog[cursor:]...), nil
|
||||
}
|
||||
|
||||
// Complete 把活动流标记为成功完成,并发布一个成功 endstream 事件。
|
||||
func (broker *StreamBroker) Complete(requestID string, terminalCode string, terminalMessage string) error {
|
||||
stream, ok := broker.Get(requestID)
|
||||
if !ok || stream == nil {
|
||||
return fmt.Errorf("request is not active: %s", strings.TrimSpace(requestID))
|
||||
}
|
||||
stream.mu.Lock()
|
||||
if stream.Status == StreamStatusCanceled || stream.Status == StreamStatusFailed || stream.Status == StreamStatusCompleted {
|
||||
stream.mu.Unlock()
|
||||
return nil
|
||||
}
|
||||
broker.stopTerminalCleanupTimerLocked(stream)
|
||||
stream.Status = StreamStatusCompleted
|
||||
subscriberCount := len(stream.Subscribers)
|
||||
stream.UpdatedAt = time.Now().UTC()
|
||||
stream.mu.Unlock()
|
||||
if err := broker.Publish(requestID, StreamEvent{
|
||||
End: true,
|
||||
TerminalErrorCode: strings.TrimSpace(terminalCode),
|
||||
TerminalErrorMessage: strings.TrimSpace(terminalMessage),
|
||||
}); err != nil {
|
||||
return err
|
||||
}
|
||||
if subscriberCount == 0 {
|
||||
broker.scheduleTerminalCleanup(requestID)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Fail 把活动流标记为失败,并发布一个失败 endstream 事件。
|
||||
func (broker *StreamBroker) Fail(requestID string, terminalCode string, terminalMessage string) error {
|
||||
stream, ok := broker.Get(requestID)
|
||||
if !ok || stream == nil {
|
||||
return fmt.Errorf("request is not active: %s", strings.TrimSpace(requestID))
|
||||
}
|
||||
stream.mu.Lock()
|
||||
broker.stopTerminalCleanupTimerLocked(stream)
|
||||
stream.Status = StreamStatusFailed
|
||||
subscriberCount := len(stream.Subscribers)
|
||||
stream.UpdatedAt = time.Now().UTC()
|
||||
stream.mu.Unlock()
|
||||
if err := broker.Publish(requestID, StreamEvent{
|
||||
End: true,
|
||||
TerminalErrorCode: strings.TrimSpace(terminalCode),
|
||||
TerminalErrorMessage: strings.TrimSpace(terminalMessage),
|
||||
}); err != nil {
|
||||
return err
|
||||
}
|
||||
if subscriberCount == 0 {
|
||||
broker.scheduleTerminalCleanup(requestID)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Cancel 主动取消活动流,并发布 canceled endstream。
|
||||
func (broker *StreamBroker) Cancel(requestID string, terminalMessage string) error {
|
||||
stream, ok := broker.Get(requestID)
|
||||
if !ok || stream == nil {
|
||||
return fmt.Errorf("request is not active: %s", strings.TrimSpace(requestID))
|
||||
}
|
||||
stream.mu.Lock()
|
||||
broker.stopTerminalCleanupTimerLocked(stream)
|
||||
if stream.ProviderCancel != nil {
|
||||
stream.ProviderCancel()
|
||||
stream.ProviderCancel = nil
|
||||
}
|
||||
stream.ProviderActive = false
|
||||
stream.Status = StreamStatusCanceled
|
||||
subscriberCount := len(stream.Subscribers)
|
||||
stream.UpdatedAt = time.Now().UTC()
|
||||
stream.mu.Unlock()
|
||||
if err := broker.Publish(requestID, StreamEvent{
|
||||
End: true,
|
||||
TerminalErrorCode: "canceled",
|
||||
TerminalErrorMessage: firstNonEmpty(strings.TrimSpace(terminalMessage), "[canceled] User aborted request"),
|
||||
}); err != nil {
|
||||
return err
|
||||
}
|
||||
if subscriberCount == 0 {
|
||||
broker.scheduleTerminalCleanup(requestID)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// firstNonEmpty 返回第一个非空白字符串。
|
||||
func firstNonEmpty(values ...string) string {
|
||||
for _, value := range values {
|
||||
if strings.TrimSpace(value) != "" {
|
||||
return strings.TrimSpace(value)
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
@@ -0,0 +1,147 @@
|
||||
package forwarder
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
runtimecore "cursor/internal/backend/agent/core"
|
||||
)
|
||||
|
||||
func (service *Service) replaceCheckpointConversation(stream *ActiveStream, conversation *ConversationFile) error {
|
||||
if stream == nil {
|
||||
return fmt.Errorf("active stream is required")
|
||||
}
|
||||
if conversation == nil {
|
||||
return fmt.Errorf("checkpoint conversation is required")
|
||||
}
|
||||
stream.mu.Lock()
|
||||
stream.CheckpointConversation = cloneConversationFile(conversation)
|
||||
stream.mu.Unlock()
|
||||
return nil
|
||||
}
|
||||
|
||||
func checkpointConversationInitialized(stream *ActiveStream) bool {
|
||||
if stream == nil {
|
||||
return false
|
||||
}
|
||||
stream.mu.Lock()
|
||||
defer stream.mu.Unlock()
|
||||
return stream.CheckpointConversation != nil
|
||||
}
|
||||
|
||||
func (service *Service) appendCheckpointEntries(stream *ActiveStream, entries []HistoryEntry) error {
|
||||
if len(entries) == 0 {
|
||||
return nil
|
||||
}
|
||||
if stream == nil {
|
||||
return fmt.Errorf("active stream is required")
|
||||
}
|
||||
stream.mu.Lock()
|
||||
defer stream.mu.Unlock()
|
||||
if stream.CheckpointConversation == nil {
|
||||
return fmt.Errorf("checkpoint conversation is not initialized")
|
||||
}
|
||||
appendEntriesInPlace(stream.CheckpointConversation, entries)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (service *Service) snapshotCheckpointConversation(stream *ActiveStream) (*ConversationFile, []runtimecore.PendingExec, []runtimecore.PendingInteraction, error) {
|
||||
if stream == nil {
|
||||
return nil, nil, nil, fmt.Errorf("active stream is required")
|
||||
}
|
||||
stream.mu.Lock()
|
||||
defer stream.mu.Unlock()
|
||||
if stream.CheckpointConversation == nil {
|
||||
return nil, nil, nil, fmt.Errorf("checkpoint conversation is not initialized")
|
||||
}
|
||||
conversation := cloneConversationFile(stream.CheckpointConversation)
|
||||
pendingExecs := make([]runtimecore.PendingExec, 0, len(stream.PendingExecs))
|
||||
for _, pending := range stream.PendingExecs {
|
||||
pendingExecs = append(pendingExecs, pending)
|
||||
}
|
||||
pendingInteractions := make([]runtimecore.PendingInteraction, 0, len(stream.PendingInteractions))
|
||||
for _, pending := range stream.PendingInteractions {
|
||||
pendingInteractions = append(pendingInteractions, pending)
|
||||
}
|
||||
return conversation, pendingExecs, pendingInteractions, nil
|
||||
}
|
||||
|
||||
func (service *Service) appendConversationEntries(stream *ActiveStream, conversationID string, entries []HistoryEntry) ([]HistoryEntry, error) {
|
||||
if len(entries) == 0 {
|
||||
return nil, nil
|
||||
}
|
||||
if stream == nil {
|
||||
return nil, fmt.Errorf("active stream is required")
|
||||
}
|
||||
stream.mu.Lock()
|
||||
if stream.CheckpointConversation == nil {
|
||||
stream.mu.Unlock()
|
||||
return nil, fmt.Errorf("checkpoint conversation is not initialized")
|
||||
}
|
||||
working := cloneConversationFile(stream.CheckpointConversation)
|
||||
if service.store != nil {
|
||||
persisted, assigned, err := service.store.AppendEntries(conversationID, resetEntrySequences(entries))
|
||||
if err != nil {
|
||||
stream.mu.Unlock()
|
||||
return nil, err
|
||||
}
|
||||
if persisted != nil {
|
||||
working = persisted
|
||||
} else {
|
||||
appendEntriesInPlace(working, assigned)
|
||||
}
|
||||
stream.CheckpointConversation = working
|
||||
stream.UpdatedAt = time.Now().UTC()
|
||||
stream.mu.Unlock()
|
||||
return assigned, nil
|
||||
}
|
||||
assigned := appendEntriesInPlace(working, entries)
|
||||
stream.CheckpointConversation = working
|
||||
stream.UpdatedAt = time.Now().UTC()
|
||||
stream.mu.Unlock()
|
||||
return assigned, nil
|
||||
}
|
||||
|
||||
func resetEntrySequences(entries []HistoryEntry) []HistoryEntry {
|
||||
if len(entries) == 0 {
|
||||
return nil
|
||||
}
|
||||
reset := make([]HistoryEntry, 0, len(entries))
|
||||
for _, entry := range entries {
|
||||
next := entry
|
||||
next.Seq = 0
|
||||
reset = append(reset, next)
|
||||
}
|
||||
return reset
|
||||
}
|
||||
|
||||
func firstWorkspacePath(paths []string) string {
|
||||
for _, path := range paths {
|
||||
if strings.TrimSpace(path) != "" {
|
||||
return strings.TrimSpace(path)
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (service *Service) updateConversationMetaAndCheckpoint(stream *ActiveStream, conversationID string, update func(*ConversationFile) error) (*ConversationFile, error) {
|
||||
if stream == nil {
|
||||
return nil, fmt.Errorf("active stream is required")
|
||||
}
|
||||
stream.mu.Lock()
|
||||
defer stream.mu.Unlock()
|
||||
if stream.CheckpointConversation == nil {
|
||||
return nil, fmt.Errorf("checkpoint conversation is not initialized")
|
||||
}
|
||||
conversation := cloneConversationFile(stream.CheckpointConversation)
|
||||
if err := update(conversation); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := service.syncConversationRecord(conversationID, conversation); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
stream.CheckpointConversation = conversation
|
||||
stream.UpdatedAt = time.Now().UTC()
|
||||
return cloneConversationFile(conversation), nil
|
||||
}
|
||||
@@ -0,0 +1,200 @@
|
||||
package forwarder
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"connectrpc.com/connect"
|
||||
"github.com/google/uuid"
|
||||
|
||||
"cursor/gen/aiserverv1"
|
||||
)
|
||||
|
||||
func (service *Service) CreateExperimentalIndex(_ context.Context, req *connect.Request[aiserverv1.CreateExperimentalIndexRequest]) (*connect.Response[aiserverv1.CreateExperimentalIndexResponse], error) {
|
||||
store, err := service.requireCodebaseIndexStore()
|
||||
if err != nil {
|
||||
return nil, connect.NewError(connect.CodeInternal, err)
|
||||
}
|
||||
record, err := store.EnsureIndexed(CodebaseIndexRecord{
|
||||
Repo: strings.TrimSpace(req.Msg.GetRepo()),
|
||||
TargetDir: strings.TrimSpace(req.Msg.GetTargetDir()),
|
||||
Files: append([]string(nil), req.Msg.GetFiles()...),
|
||||
})
|
||||
if err != nil {
|
||||
return nil, connect.NewError(connect.CodeInternal, err)
|
||||
}
|
||||
return connect.NewResponse(&aiserverv1.CreateExperimentalIndexResponse{IndexId: record.IndexID}), nil
|
||||
}
|
||||
|
||||
func (service *Service) ListExperimentalIndexFiles(_ context.Context, req *connect.Request[aiserverv1.ListExperimentalIndexFilesRequest]) (*connect.Response[aiserverv1.ListExperimentalIndexFilesResponse], error) {
|
||||
store, err := service.requireCodebaseIndexStore()
|
||||
if err != nil {
|
||||
return nil, connect.NewError(connect.CodeInternal, err)
|
||||
}
|
||||
indexID := strings.TrimSpace(req.Msg.GetIndexId())
|
||||
files, err := store.ListFiles(indexID)
|
||||
if err != nil {
|
||||
return nil, connect.NewError(connect.CodeInternal, err)
|
||||
}
|
||||
return connect.NewResponse(&aiserverv1.ListExperimentalIndexFilesResponse{
|
||||
IndexId: indexID,
|
||||
Files: codebaseIndexFileData(indexID, files),
|
||||
}), nil
|
||||
}
|
||||
|
||||
func (service *Service) ListenExperimentalIndex(ctx context.Context, req *connect.Request[aiserverv1.ListenExperimentalIndexRequest], stream *connect.ServerStream[aiserverv1.ListenExperimentalIndexResponse]) error {
|
||||
store, err := service.requireCodebaseIndexStore()
|
||||
if err != nil {
|
||||
return connect.NewError(connect.CodeInternal, err)
|
||||
}
|
||||
indexID := strings.TrimSpace(req.Msg.GetIndexId())
|
||||
if indexID == "" {
|
||||
return connect.NewError(connect.CodeInvalidArgument, fmt.Errorf("index_id is required"))
|
||||
}
|
||||
_, err = store.EnsureIndexed(CodebaseIndexRecord{IndexID: indexID})
|
||||
if err != nil {
|
||||
return connect.NewError(connect.CodeInternal, err)
|
||||
}
|
||||
subscription, err := store.Subscribe(indexID)
|
||||
if err != nil {
|
||||
return connect.NewError(connect.CodeInternal, err)
|
||||
}
|
||||
defer store.Unsubscribe(subscription)
|
||||
if err := stream.Send(&aiserverv1.ListenExperimentalIndexResponse{
|
||||
IndexId: indexID,
|
||||
Item: &aiserverv1.ListenExperimentalIndexResponse_Ready{
|
||||
Ready: &aiserverv1.ListenExperimentalIndexResponse_ReadyItem{
|
||||
IndexId: indexID,
|
||||
Request: &aiserverv1.ListenExperimentalIndexRequest{
|
||||
IndexId: indexID,
|
||||
},
|
||||
},
|
||||
},
|
||||
}); err != nil {
|
||||
return err
|
||||
}
|
||||
var lastEventID int64
|
||||
for {
|
||||
events, err := store.ListEvents(indexID, lastEventID)
|
||||
if err != nil {
|
||||
return connect.NewError(connect.CodeInternal, err)
|
||||
}
|
||||
for _, event := range events {
|
||||
if event.ID > lastEventID {
|
||||
lastEventID = event.ID
|
||||
}
|
||||
if event.Kind != codebaseIndexEventKindRegister {
|
||||
continue
|
||||
}
|
||||
if err := stream.Send(codebaseIndexRegisterEventResponse(indexID, event)); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return nil
|
||||
case <-subscription.Signal:
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (service *Service) RegisterFileToIndex(_ context.Context, req *connect.Request[aiserverv1.RegisterFileToIndexRequest]) (*connect.Response[aiserverv1.RequestReceivedResponse], error) {
|
||||
store, err := service.requireCodebaseIndexStore()
|
||||
if err != nil {
|
||||
return nil, connect.NewError(connect.CodeInternal, err)
|
||||
}
|
||||
reqUUID := uuid.NewString()
|
||||
_, _, err = store.RegisterFileWithRequest(req.Msg.GetIndexId(), reqUUID, CodebaseIndexFileRecord{
|
||||
WorkspaceRelativePath: strings.TrimSpace(req.Msg.GetWorkspaceRelativePath()),
|
||||
ContentHash: codebaseFileContentHash(strings.Join(req.Msg.GetContent(), "\n")),
|
||||
Stage: codebaseIndexFileStage,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, connect.NewError(connect.CodeInvalidArgument, err)
|
||||
}
|
||||
return connect.NewResponse(&aiserverv1.RequestReceivedResponse{ReqUuid: reqUUID}), nil
|
||||
}
|
||||
|
||||
func (service *Service) SetupIndexDependencies(_ context.Context, req *connect.Request[aiserverv1.SetupIndexDependenciesRequest]) (*connect.Response[aiserverv1.SetupIndexDependenciesResponse], error) {
|
||||
store, err := service.requireCodebaseIndexStore()
|
||||
if err != nil {
|
||||
return nil, connect.NewError(connect.CodeInternal, err)
|
||||
}
|
||||
if err := store.MarkDependenciesReady(req.Msg.GetIndexId()); err != nil {
|
||||
return nil, connect.NewError(connect.CodeInvalidArgument, err)
|
||||
}
|
||||
return connect.NewResponse(&aiserverv1.SetupIndexDependenciesResponse{}), nil
|
||||
}
|
||||
|
||||
func (service *Service) ComputeIndexTopoSort(_ context.Context, req *connect.Request[aiserverv1.ComputeIndexTopoSortRequest]) (*connect.Response[aiserverv1.ComputeIndexTopoSortResponse], error) {
|
||||
store, err := service.requireCodebaseIndexStore()
|
||||
if err != nil {
|
||||
return nil, connect.NewError(connect.CodeInternal, err)
|
||||
}
|
||||
if err := store.MarkTopoSortReady(req.Msg.GetIndexId()); err != nil {
|
||||
return nil, connect.NewError(connect.CodeInvalidArgument, err)
|
||||
}
|
||||
return connect.NewResponse(&aiserverv1.ComputeIndexTopoSortResponse{}), nil
|
||||
}
|
||||
|
||||
func (service *Service) requireCodebaseIndexStore() (*CodebaseIndexStore, error) {
|
||||
if service == nil || service.codebaseIndexStore == nil {
|
||||
return nil, fmt.Errorf("codebase index store is not initialized")
|
||||
}
|
||||
return service.codebaseIndexStore, nil
|
||||
}
|
||||
|
||||
func codebaseIndexFileData(indexID string, files []CodebaseIndexFileRecord) []*aiserverv1.IndexFileData {
|
||||
result := make([]*aiserverv1.IndexFileData, 0, len(files))
|
||||
for index, file := range files {
|
||||
result = append(result, codebaseIndexFileDatum(indexID, file, int32(index+1)))
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
func codebaseIndexFileDatum(indexID string, file CodebaseIndexFileRecord, fallbackOrder int32) *aiserverv1.IndexFileData {
|
||||
order := file.Order
|
||||
if order == 0 {
|
||||
order = fallbackOrder
|
||||
}
|
||||
return &aiserverv1.IndexFileData{
|
||||
IndexId: indexID,
|
||||
WorkspaceRelativePath: file.WorkspaceRelativePath,
|
||||
Stage: file.Stage,
|
||||
Order: order,
|
||||
}
|
||||
}
|
||||
|
||||
func codebaseIndexRegisterEventResponse(indexID string, event CodebaseIndexEventRecord) *aiserverv1.ListenExperimentalIndexResponse {
|
||||
fileData := codebaseIndexFileDatum(indexID, event.File, 1)
|
||||
return &aiserverv1.ListenExperimentalIndexResponse{
|
||||
IndexId: indexID,
|
||||
Item: &aiserverv1.ListenExperimentalIndexResponse_Register{
|
||||
Register: &aiserverv1.ListenExperimentalIndexResponse_RegisterItem{
|
||||
ReqUuid: strings.TrimSpace(event.ReqUUID),
|
||||
Request: &aiserverv1.RegisterFileToIndexRequest{
|
||||
IndexId: indexID,
|
||||
WorkspaceRelativePath: event.File.WorkspaceRelativePath,
|
||||
},
|
||||
Response: &aiserverv1.RegisterFileToIndexResponse{
|
||||
FileId: codebaseIndexFileID(indexID, event.File.WorkspaceRelativePath),
|
||||
RootContextNodeId: codebaseIndexFileID(indexID, event.File.WorkspaceRelativePath),
|
||||
FileData: fileData,
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func codebaseIndexFileID(indexID string, workspaceRelativePath string) string {
|
||||
trimmedIndexID := strings.TrimSpace(indexID)
|
||||
trimmedPath := strings.TrimSpace(workspaceRelativePath)
|
||||
if trimmedIndexID == "" || trimmedPath == "" {
|
||||
return ""
|
||||
}
|
||||
sum := sha256.Sum256([]byte(trimmedIndexID + "\x00" + trimmedPath))
|
||||
return "file_" + hex.EncodeToString(sum[:])[:24]
|
||||
}
|
||||
@@ -0,0 +1,606 @@
|
||||
package forwarder
|
||||
|
||||
import (
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sort"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/google/uuid"
|
||||
)
|
||||
|
||||
const (
|
||||
codebaseIndexStatusReady = "ready"
|
||||
codebaseIndexFileStage = "indexed"
|
||||
codebaseIndexEventKindRegister = "register"
|
||||
codebaseIndexSubscriberSignalCapacity = 1
|
||||
codebaseIndexMaxPersistedEvents = 2000
|
||||
codebaseIndexSourceRepositoryService = "repository_service"
|
||||
repositoryIndexStatusReady = "ready"
|
||||
repositoryIndexSyncStatusCompleted = "sync_completed"
|
||||
repositoryIndexCopyStatusCompleted = "copy_completed"
|
||||
)
|
||||
|
||||
type CodebaseSearchBackend interface {
|
||||
EnsureIndexed(record CodebaseIndexRecord) (CodebaseIndexRecord, error)
|
||||
RegisterFile(indexID string, file CodebaseIndexFileRecord) (CodebaseIndexFileRecord, error)
|
||||
ListFiles(indexID string) ([]CodebaseIndexFileRecord, error)
|
||||
}
|
||||
|
||||
type CodebaseIndexStore struct {
|
||||
mu sync.Mutex
|
||||
root string
|
||||
path string
|
||||
loaded bool
|
||||
nextEventID int64
|
||||
state codebaseIndexState
|
||||
subscribers map[string]map[string]chan struct{}
|
||||
}
|
||||
|
||||
type codebaseIndexState struct {
|
||||
SchemaVersion int `json:"schema_version"`
|
||||
Indexes map[string]CodebaseIndexRecord `json:"indexes"`
|
||||
Events []CodebaseIndexEventRecord `json:"events,omitempty"`
|
||||
}
|
||||
|
||||
type CodebaseIndexRecord struct {
|
||||
IndexID string `json:"index_id"`
|
||||
Repo string `json:"repo,omitempty"`
|
||||
TargetDir string `json:"target_dir,omitempty"`
|
||||
Files []string `json:"files,omitempty"`
|
||||
Status string `json:"status"`
|
||||
Source string `json:"source,omitempty"`
|
||||
DependenciesReady bool `json:"dependencies_ready,omitempty"`
|
||||
TopoSortReady bool `json:"topo_sort_ready,omitempty"`
|
||||
RepositoryState *RepositoryIndexStateRecord `json:"repository_state,omitempty"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
RegisteredFiles map[string]CodebaseIndexFileRecord `json:"registered_files,omitempty"`
|
||||
}
|
||||
|
||||
type RepositoryIndexStateRecord struct {
|
||||
CodebaseID string `json:"codebase_id"`
|
||||
Repo string `json:"repo,omitempty"`
|
||||
TargetDir string `json:"target_dir,omitempty"`
|
||||
Status string `json:"status"`
|
||||
CopyStatus string `json:"copy_status,omitempty"`
|
||||
SyncStatus string `json:"sync_status,omitempty"`
|
||||
CopyTaskHandle string `json:"copy_task_handle,omitempty"`
|
||||
PathKeyHash string `json:"path_key_hash,omitempty"`
|
||||
Source string `json:"source"`
|
||||
LastHandshakeAt time.Time `json:"last_handshake_at,omitempty"`
|
||||
LastSyncAt time.Time `json:"last_sync_at,omitempty"`
|
||||
LastCopyStatusAt time.Time `json:"last_copy_status_at,omitempty"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
}
|
||||
|
||||
type CodebaseIndexFileRecord struct {
|
||||
WorkspaceRelativePath string `json:"workspace_relative_path"`
|
||||
ContentHash string `json:"content_hash,omitempty"`
|
||||
Stage string `json:"stage"`
|
||||
Order int32 `json:"order,omitempty"`
|
||||
NodeCount int32 `json:"node_count,omitempty"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
}
|
||||
|
||||
type CodebaseIndexEventRecord struct {
|
||||
ID int64 `json:"id"`
|
||||
IndexID string `json:"index_id"`
|
||||
Kind string `json:"kind"`
|
||||
ReqUUID string `json:"req_uuid,omitempty"`
|
||||
File CodebaseIndexFileRecord `json:"file,omitempty"`
|
||||
At time.Time `json:"at"`
|
||||
}
|
||||
|
||||
type CodebaseIndexSubscription struct {
|
||||
IndexID string
|
||||
SubscriberID string
|
||||
Signal <-chan struct{}
|
||||
}
|
||||
|
||||
func NewCodebaseIndexStore(root string) *CodebaseIndexStore {
|
||||
trimmed := strings.TrimSpace(root)
|
||||
if trimmed == "" {
|
||||
trimmed = "codebase-index"
|
||||
}
|
||||
return &CodebaseIndexStore{
|
||||
root: trimmed,
|
||||
path: filepath.Join(trimmed, "index.json"),
|
||||
subscribers: make(map[string]map[string]chan struct{}),
|
||||
state: codebaseIndexState{
|
||||
SchemaVersion: 1,
|
||||
Indexes: make(map[string]CodebaseIndexRecord),
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func (store *CodebaseIndexStore) EnsureIndexed(record CodebaseIndexRecord) (CodebaseIndexRecord, error) {
|
||||
if store == nil {
|
||||
return CodebaseIndexRecord{}, fmt.Errorf("codebase index store is nil")
|
||||
}
|
||||
store.mu.Lock()
|
||||
defer store.mu.Unlock()
|
||||
if err := store.loadLocked(); err != nil {
|
||||
return CodebaseIndexRecord{}, err
|
||||
}
|
||||
now := time.Now().UTC()
|
||||
indexID := strings.TrimSpace(record.IndexID)
|
||||
if indexID == "" {
|
||||
indexID = stableCodebaseIndexID(record.Repo, record.TargetDir, record.Files)
|
||||
}
|
||||
existing, ok := store.state.Indexes[indexID]
|
||||
if ok {
|
||||
existing.Repo = firstNonEmptyCodebaseIndex(record.Repo, existing.Repo)
|
||||
existing.TargetDir = firstNonEmptyCodebaseIndex(record.TargetDir, existing.TargetDir)
|
||||
if len(record.Files) > 0 {
|
||||
existing.Files = append([]string(nil), record.Files...)
|
||||
}
|
||||
if strings.TrimSpace(record.Source) != "" {
|
||||
existing.Source = strings.TrimSpace(record.Source)
|
||||
}
|
||||
if record.RepositoryState != nil {
|
||||
existing.RepositoryState = mergeRepositoryIndexState(existing.RepositoryState, record.RepositoryState, now)
|
||||
}
|
||||
existing.Status = codebaseIndexStatusReady
|
||||
existing.UpdatedAt = now
|
||||
if existing.RegisteredFiles == nil {
|
||||
existing.RegisteredFiles = make(map[string]CodebaseIndexFileRecord)
|
||||
}
|
||||
store.state.Indexes[indexID] = existing
|
||||
return existing, store.saveLocked()
|
||||
}
|
||||
record.IndexID = indexID
|
||||
record.Status = codebaseIndexStatusReady
|
||||
record.Source = strings.TrimSpace(record.Source)
|
||||
record.CreatedAt = now
|
||||
record.UpdatedAt = now
|
||||
if record.RepositoryState != nil {
|
||||
record.RepositoryState = mergeRepositoryIndexState(nil, record.RepositoryState, now)
|
||||
}
|
||||
record.Files = append([]string(nil), record.Files...)
|
||||
record.RegisteredFiles = make(map[string]CodebaseIndexFileRecord)
|
||||
store.state.Indexes[indexID] = record
|
||||
return record, store.saveLocked()
|
||||
}
|
||||
|
||||
func (store *CodebaseIndexStore) EnsureRepositoryIndexed(state RepositoryIndexStateRecord) (CodebaseIndexRecord, error) {
|
||||
if store == nil {
|
||||
return CodebaseIndexRecord{}, fmt.Errorf("codebase index store is nil")
|
||||
}
|
||||
codebaseID := strings.TrimSpace(state.CodebaseID)
|
||||
if codebaseID == "" {
|
||||
return CodebaseIndexRecord{}, fmt.Errorf("codebase_id is required")
|
||||
}
|
||||
state.CodebaseID = codebaseID
|
||||
state.Source = codebaseIndexSourceRepositoryService
|
||||
if strings.TrimSpace(state.Status) == "" {
|
||||
state.Status = repositoryIndexStatusReady
|
||||
}
|
||||
return store.EnsureIndexed(CodebaseIndexRecord{
|
||||
IndexID: codebaseID,
|
||||
Repo: strings.TrimSpace(state.Repo),
|
||||
TargetDir: strings.TrimSpace(state.TargetDir),
|
||||
Source: codebaseIndexSourceRepositoryService,
|
||||
RepositoryState: &state,
|
||||
})
|
||||
}
|
||||
|
||||
func (store *CodebaseIndexStore) MarkRepositorySyncComplete(codebaseID string, pathKeyHash string) error {
|
||||
_, err := store.EnsureRepositoryIndexed(RepositoryIndexStateRecord{
|
||||
CodebaseID: strings.TrimSpace(codebaseID),
|
||||
Status: repositoryIndexStatusReady,
|
||||
SyncStatus: repositoryIndexSyncStatusCompleted,
|
||||
PathKeyHash: strings.TrimSpace(pathKeyHash),
|
||||
LastSyncAt: time.Now().UTC(),
|
||||
})
|
||||
return err
|
||||
}
|
||||
|
||||
func (store *CodebaseIndexStore) MarkRepositoryCopyComplete(codebaseID string, copyTaskHandle string) error {
|
||||
_, err := store.EnsureRepositoryIndexed(RepositoryIndexStateRecord{
|
||||
CodebaseID: strings.TrimSpace(codebaseID),
|
||||
Status: repositoryIndexStatusReady,
|
||||
CopyStatus: repositoryIndexCopyStatusCompleted,
|
||||
CopyTaskHandle: strings.TrimSpace(copyTaskHandle),
|
||||
LastCopyStatusAt: time.Now().UTC(),
|
||||
})
|
||||
return err
|
||||
}
|
||||
|
||||
func (store *CodebaseIndexStore) RegisterFile(indexID string, file CodebaseIndexFileRecord) (CodebaseIndexFileRecord, error) {
|
||||
registered, _, err := store.RegisterFileWithRequest(indexID, "", file)
|
||||
return registered, err
|
||||
}
|
||||
|
||||
func (store *CodebaseIndexStore) RegisterFileWithRequest(indexID string, reqUUID string, file CodebaseIndexFileRecord) (CodebaseIndexFileRecord, CodebaseIndexEventRecord, error) {
|
||||
if store == nil {
|
||||
return CodebaseIndexFileRecord{}, CodebaseIndexEventRecord{}, fmt.Errorf("codebase index store is nil")
|
||||
}
|
||||
store.mu.Lock()
|
||||
defer store.mu.Unlock()
|
||||
if err := store.loadLocked(); err != nil {
|
||||
return CodebaseIndexFileRecord{}, CodebaseIndexEventRecord{}, err
|
||||
}
|
||||
trimmedIndexID := strings.TrimSpace(indexID)
|
||||
if trimmedIndexID == "" {
|
||||
return CodebaseIndexFileRecord{}, CodebaseIndexEventRecord{}, fmt.Errorf("index_id is required")
|
||||
}
|
||||
record, ok := store.state.Indexes[trimmedIndexID]
|
||||
if !ok {
|
||||
now := time.Now().UTC()
|
||||
record = CodebaseIndexRecord{
|
||||
IndexID: trimmedIndexID,
|
||||
Status: codebaseIndexStatusReady,
|
||||
CreatedAt: now,
|
||||
UpdatedAt: now,
|
||||
RegisteredFiles: make(map[string]CodebaseIndexFileRecord),
|
||||
}
|
||||
}
|
||||
if record.RegisteredFiles == nil {
|
||||
record.RegisteredFiles = make(map[string]CodebaseIndexFileRecord)
|
||||
}
|
||||
path := strings.TrimSpace(file.WorkspaceRelativePath)
|
||||
if path == "" {
|
||||
return CodebaseIndexFileRecord{}, CodebaseIndexEventRecord{}, fmt.Errorf("workspace_relative_path is required")
|
||||
}
|
||||
now := time.Now().UTC()
|
||||
current, exists := record.RegisteredFiles[path]
|
||||
if exists {
|
||||
file.CreatedAt = current.CreatedAt
|
||||
} else {
|
||||
file.CreatedAt = now
|
||||
}
|
||||
file.WorkspaceRelativePath = path
|
||||
if strings.TrimSpace(file.Stage) == "" {
|
||||
file.Stage = codebaseIndexFileStage
|
||||
}
|
||||
if file.Order == 0 {
|
||||
if exists && current.Order != 0 {
|
||||
file.Order = current.Order
|
||||
} else {
|
||||
file.Order = int32(len(record.RegisteredFiles) + 1)
|
||||
}
|
||||
}
|
||||
file.UpdatedAt = now
|
||||
record.RegisteredFiles[path] = file
|
||||
record.Status = codebaseIndexStatusReady
|
||||
record.UpdatedAt = now
|
||||
store.state.Indexes[trimmedIndexID] = record
|
||||
event := CodebaseIndexEventRecord{
|
||||
ID: store.nextCodebaseIndexEventIDLocked(),
|
||||
IndexID: trimmedIndexID,
|
||||
Kind: codebaseIndexEventKindRegister,
|
||||
ReqUUID: strings.TrimSpace(reqUUID),
|
||||
File: file,
|
||||
At: now,
|
||||
}
|
||||
store.state.Events = append(store.state.Events, event)
|
||||
store.compactEventsLocked()
|
||||
if err := store.saveLocked(); err != nil {
|
||||
return CodebaseIndexFileRecord{}, CodebaseIndexEventRecord{}, err
|
||||
}
|
||||
store.notifyIndexSubscribersLocked(trimmedIndexID)
|
||||
return file, event, nil
|
||||
}
|
||||
|
||||
func (store *CodebaseIndexStore) ListFiles(indexID string) ([]CodebaseIndexFileRecord, error) {
|
||||
if store == nil {
|
||||
return nil, fmt.Errorf("codebase index store is nil")
|
||||
}
|
||||
store.mu.Lock()
|
||||
defer store.mu.Unlock()
|
||||
if err := store.loadLocked(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
record, ok := store.state.Indexes[strings.TrimSpace(indexID)]
|
||||
if !ok || len(record.RegisteredFiles) == 0 {
|
||||
return nil, nil
|
||||
}
|
||||
files := make([]CodebaseIndexFileRecord, 0, len(record.RegisteredFiles))
|
||||
for _, file := range record.RegisteredFiles {
|
||||
files = append(files, file)
|
||||
}
|
||||
sort.Slice(files, func(i, j int) bool {
|
||||
if files[i].Order != files[j].Order {
|
||||
return files[i].Order < files[j].Order
|
||||
}
|
||||
return files[i].WorkspaceRelativePath < files[j].WorkspaceRelativePath
|
||||
})
|
||||
return files, nil
|
||||
}
|
||||
|
||||
func (store *CodebaseIndexStore) ListEvents(indexID string, afterEventID int64) ([]CodebaseIndexEventRecord, error) {
|
||||
if store == nil {
|
||||
return nil, fmt.Errorf("codebase index store is nil")
|
||||
}
|
||||
store.mu.Lock()
|
||||
defer store.mu.Unlock()
|
||||
if err := store.loadLocked(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
trimmedIndexID := strings.TrimSpace(indexID)
|
||||
if trimmedIndexID == "" {
|
||||
return nil, fmt.Errorf("index_id is required")
|
||||
}
|
||||
events := make([]CodebaseIndexEventRecord, 0)
|
||||
for _, event := range store.state.Events {
|
||||
if event.ID <= afterEventID || event.IndexID != trimmedIndexID {
|
||||
continue
|
||||
}
|
||||
events = append(events, event)
|
||||
}
|
||||
return events, nil
|
||||
}
|
||||
|
||||
func (store *CodebaseIndexStore) Subscribe(indexID string) (CodebaseIndexSubscription, error) {
|
||||
if store == nil {
|
||||
return CodebaseIndexSubscription{}, fmt.Errorf("codebase index store is nil")
|
||||
}
|
||||
store.mu.Lock()
|
||||
defer store.mu.Unlock()
|
||||
if err := store.loadLocked(); err != nil {
|
||||
return CodebaseIndexSubscription{}, err
|
||||
}
|
||||
trimmedIndexID := strings.TrimSpace(indexID)
|
||||
if trimmedIndexID == "" {
|
||||
return CodebaseIndexSubscription{}, fmt.Errorf("index_id is required")
|
||||
}
|
||||
if store.subscribers == nil {
|
||||
store.subscribers = make(map[string]map[string]chan struct{})
|
||||
}
|
||||
subscriberID := uuid.NewString()
|
||||
signal := make(chan struct{}, codebaseIndexSubscriberSignalCapacity)
|
||||
if store.subscribers[trimmedIndexID] == nil {
|
||||
store.subscribers[trimmedIndexID] = make(map[string]chan struct{})
|
||||
}
|
||||
store.subscribers[trimmedIndexID][subscriberID] = signal
|
||||
return CodebaseIndexSubscription{
|
||||
IndexID: trimmedIndexID,
|
||||
SubscriberID: subscriberID,
|
||||
Signal: signal,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (store *CodebaseIndexStore) Unsubscribe(subscription CodebaseIndexSubscription) {
|
||||
if store == nil {
|
||||
return
|
||||
}
|
||||
store.mu.Lock()
|
||||
defer store.mu.Unlock()
|
||||
indexID := strings.TrimSpace(subscription.IndexID)
|
||||
subscriberID := strings.TrimSpace(subscription.SubscriberID)
|
||||
if indexID == "" || subscriberID == "" {
|
||||
return
|
||||
}
|
||||
byIndex := store.subscribers[indexID]
|
||||
if byIndex == nil {
|
||||
return
|
||||
}
|
||||
delete(byIndex, subscriberID)
|
||||
if len(byIndex) == 0 {
|
||||
delete(store.subscribers, indexID)
|
||||
}
|
||||
}
|
||||
|
||||
func (store *CodebaseIndexStore) MarkDependenciesReady(indexID string) error {
|
||||
return store.updateIndex(indexID, func(record *CodebaseIndexRecord) {
|
||||
record.DependenciesReady = true
|
||||
})
|
||||
}
|
||||
|
||||
func (store *CodebaseIndexStore) MarkTopoSortReady(indexID string) error {
|
||||
return store.updateIndex(indexID, func(record *CodebaseIndexRecord) {
|
||||
record.TopoSortReady = true
|
||||
})
|
||||
}
|
||||
|
||||
func (store *CodebaseIndexStore) updateIndex(indexID string, update func(*CodebaseIndexRecord)) error {
|
||||
if store == nil {
|
||||
return fmt.Errorf("codebase index store is nil")
|
||||
}
|
||||
store.mu.Lock()
|
||||
defer store.mu.Unlock()
|
||||
if err := store.loadLocked(); err != nil {
|
||||
return err
|
||||
}
|
||||
trimmedIndexID := strings.TrimSpace(indexID)
|
||||
if trimmedIndexID == "" {
|
||||
return fmt.Errorf("index_id is required")
|
||||
}
|
||||
now := time.Now().UTC()
|
||||
record, ok := store.state.Indexes[trimmedIndexID]
|
||||
if !ok {
|
||||
record = CodebaseIndexRecord{
|
||||
IndexID: trimmedIndexID,
|
||||
Status: codebaseIndexStatusReady,
|
||||
CreatedAt: now,
|
||||
RegisteredFiles: make(map[string]CodebaseIndexFileRecord),
|
||||
}
|
||||
}
|
||||
if record.RegisteredFiles == nil {
|
||||
record.RegisteredFiles = make(map[string]CodebaseIndexFileRecord)
|
||||
}
|
||||
update(&record)
|
||||
record.Status = codebaseIndexStatusReady
|
||||
record.UpdatedAt = now
|
||||
store.state.Indexes[trimmedIndexID] = record
|
||||
return store.saveLocked()
|
||||
}
|
||||
|
||||
func (store *CodebaseIndexStore) loadLocked() error {
|
||||
if store.loaded {
|
||||
return nil
|
||||
}
|
||||
state := codebaseIndexState{SchemaVersion: 1, Indexes: make(map[string]CodebaseIndexRecord)}
|
||||
data, err := os.ReadFile(store.path)
|
||||
if err != nil {
|
||||
if os.IsNotExist(err) {
|
||||
store.state = state
|
||||
store.nextEventID = 1
|
||||
store.loaded = true
|
||||
return nil
|
||||
}
|
||||
return err
|
||||
}
|
||||
if len(data) != 0 {
|
||||
if err := json.Unmarshal(data, &state); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
if state.SchemaVersion <= 0 {
|
||||
state.SchemaVersion = 1
|
||||
}
|
||||
if state.Indexes == nil {
|
||||
state.Indexes = make(map[string]CodebaseIndexRecord)
|
||||
}
|
||||
var maxEventID int64
|
||||
for id, record := range state.Indexes {
|
||||
if record.RegisteredFiles == nil {
|
||||
record.RegisteredFiles = make(map[string]CodebaseIndexFileRecord)
|
||||
state.Indexes[id] = record
|
||||
}
|
||||
}
|
||||
for _, event := range state.Events {
|
||||
if event.ID > maxEventID {
|
||||
maxEventID = event.ID
|
||||
}
|
||||
}
|
||||
store.state = state
|
||||
store.nextEventID = maxEventID + 1
|
||||
store.loaded = true
|
||||
return nil
|
||||
}
|
||||
|
||||
func (store *CodebaseIndexStore) saveLocked() error {
|
||||
if err := os.MkdirAll(store.root, 0o755); err != nil {
|
||||
return err
|
||||
}
|
||||
data, err := json.MarshalIndent(store.state, "", " ")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
tmp, err := os.CreateTemp(store.root, ".index-*.json.tmp")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
tmpPath := tmp.Name()
|
||||
defer func() {
|
||||
_ = os.Remove(tmpPath)
|
||||
}()
|
||||
if _, err := tmp.Write(data); err != nil {
|
||||
_ = tmp.Close()
|
||||
return err
|
||||
}
|
||||
if err := tmp.Close(); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := os.Chmod(tmpPath, 0o600); err != nil {
|
||||
return err
|
||||
}
|
||||
return os.Rename(tmpPath, store.path)
|
||||
}
|
||||
|
||||
func (store *CodebaseIndexStore) nextCodebaseIndexEventIDLocked() int64 {
|
||||
if store.nextEventID <= 0 {
|
||||
store.nextEventID = 1
|
||||
}
|
||||
eventID := store.nextEventID
|
||||
store.nextEventID++
|
||||
return eventID
|
||||
}
|
||||
|
||||
func (store *CodebaseIndexStore) notifyIndexSubscribersLocked(indexID string) {
|
||||
if store.subscribers == nil {
|
||||
return
|
||||
}
|
||||
for _, signal := range store.subscribers[strings.TrimSpace(indexID)] {
|
||||
select {
|
||||
case signal <- struct{}{}:
|
||||
default:
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (store *CodebaseIndexStore) compactEventsLocked() {
|
||||
if len(store.state.Events) <= codebaseIndexMaxPersistedEvents {
|
||||
return
|
||||
}
|
||||
store.state.Events = append([]CodebaseIndexEventRecord(nil), store.state.Events[len(store.state.Events)-codebaseIndexMaxPersistedEvents:]...)
|
||||
}
|
||||
|
||||
func stableCodebaseIndexID(repo string, targetDir string, files []string) string {
|
||||
parts := append([]string{strings.TrimSpace(repo), strings.TrimSpace(targetDir)}, files...)
|
||||
normalized := make([]string, 0, len(parts))
|
||||
for _, part := range parts {
|
||||
trimmed := strings.TrimSpace(part)
|
||||
if trimmed != "" {
|
||||
normalized = append(normalized, trimmed)
|
||||
}
|
||||
}
|
||||
if len(normalized) == 0 {
|
||||
return "idx_local"
|
||||
}
|
||||
sort.Strings(normalized)
|
||||
sum := sha256.Sum256([]byte(strings.Join(normalized, "\x00")))
|
||||
return "idx_" + hex.EncodeToString(sum[:])[:24]
|
||||
}
|
||||
|
||||
func codebaseFileContentHash(content string) string {
|
||||
if content == "" {
|
||||
return ""
|
||||
}
|
||||
sum := sha256.Sum256([]byte(content))
|
||||
return hex.EncodeToString(sum[:])
|
||||
}
|
||||
|
||||
func firstNonEmptyCodebaseIndex(values ...string) string {
|
||||
for _, value := range values {
|
||||
trimmed := strings.TrimSpace(value)
|
||||
if trimmed != "" {
|
||||
return trimmed
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func mergeRepositoryIndexState(existing *RepositoryIndexStateRecord, next *RepositoryIndexStateRecord, now time.Time) *RepositoryIndexStateRecord {
|
||||
if next == nil {
|
||||
return existing
|
||||
}
|
||||
merged := RepositoryIndexStateRecord{}
|
||||
if existing != nil {
|
||||
merged = *existing
|
||||
}
|
||||
merged.CodebaseID = firstNonEmptyCodebaseIndex(next.CodebaseID, merged.CodebaseID)
|
||||
merged.Repo = firstNonEmptyCodebaseIndex(next.Repo, merged.Repo)
|
||||
merged.TargetDir = firstNonEmptyCodebaseIndex(next.TargetDir, merged.TargetDir)
|
||||
merged.Status = firstNonEmptyCodebaseIndex(next.Status, merged.Status, repositoryIndexStatusReady)
|
||||
merged.CopyStatus = firstNonEmptyCodebaseIndex(next.CopyStatus, merged.CopyStatus)
|
||||
merged.SyncStatus = firstNonEmptyCodebaseIndex(next.SyncStatus, merged.SyncStatus)
|
||||
merged.CopyTaskHandle = firstNonEmptyCodebaseIndex(next.CopyTaskHandle, merged.CopyTaskHandle)
|
||||
merged.PathKeyHash = firstNonEmptyCodebaseIndex(next.PathKeyHash, merged.PathKeyHash)
|
||||
merged.Source = firstNonEmptyCodebaseIndex(next.Source, merged.Source, codebaseIndexSourceRepositoryService)
|
||||
merged.LastHandshakeAt = firstNonZeroCodebaseIndexTime(next.LastHandshakeAt, merged.LastHandshakeAt)
|
||||
merged.LastSyncAt = firstNonZeroCodebaseIndexTime(next.LastSyncAt, merged.LastSyncAt)
|
||||
merged.LastCopyStatusAt = firstNonZeroCodebaseIndexTime(next.LastCopyStatusAt, merged.LastCopyStatusAt)
|
||||
merged.CreatedAt = firstNonZeroCodebaseIndexTime(merged.CreatedAt, next.CreatedAt, now)
|
||||
merged.UpdatedAt = now
|
||||
return &merged
|
||||
}
|
||||
|
||||
func firstNonZeroCodebaseIndexTime(values ...time.Time) time.Time {
|
||||
for _, value := range values {
|
||||
if !value.IsZero() {
|
||||
return value.UTC()
|
||||
}
|
||||
}
|
||||
return time.Time{}
|
||||
}
|
||||
@@ -0,0 +1,306 @@
|
||||
package forwarder
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log"
|
||||
"strings"
|
||||
|
||||
"connectrpc.com/connect"
|
||||
"github.com/google/uuid"
|
||||
"google.golang.org/protobuf/encoding/protojson"
|
||||
|
||||
"cursor/gen/agentv1"
|
||||
"cursor/gen/aiserverv1"
|
||||
modeladapter "cursor/internal/backend/agent/model"
|
||||
promptassets "cursor/prompt"
|
||||
)
|
||||
|
||||
const (
|
||||
commitMessageDiffTotalLimit = 120_000
|
||||
commitMessageSingleDiffLimit = 40_000
|
||||
commitMessagePreviousCommitLimit = 12
|
||||
commitMessageExplicitContextLimit = 20_000
|
||||
commitMessageMaxOutputTokens = 512
|
||||
commitMessageGeneratedRequestIDPrefix = "commit-message-"
|
||||
)
|
||||
|
||||
var errCommitMessageToolInvocation = errors.New("commit message generation must not invoke tools")
|
||||
|
||||
// WriteGitCommitMessage handles Cursor's SCM "Generate Commit Message" action.
|
||||
func (service *Service) WriteGitCommitMessage(ctx context.Context, req *connect.Request[aiserverv1.WriteGitCommitMessageRequest]) (*connect.Response[aiserverv1.WriteGitCommitMessageResponse], error) {
|
||||
if service == nil {
|
||||
return nil, connect.NewError(connect.CodeInternal, fmt.Errorf("forwarder service is nil"))
|
||||
}
|
||||
requestID := commitMessageGeneratedRequestIDPrefix + uuid.NewString()
|
||||
recorder, err := newCommitMessageLogRecorder(service.commitMessageHistoryRoot(), requestID)
|
||||
if err != nil {
|
||||
return nil, connect.NewError(connect.CodeInternal, err)
|
||||
}
|
||||
if req == nil || req.Msg == nil {
|
||||
return nil, commitMessageConnectError(recorder, connect.CodeInvalidArgument, fmt.Errorf("write git commit message request is required"))
|
||||
}
|
||||
if err := recordCommitMessageIncomingRequest(recorder, requestID, req.Msg); err != nil {
|
||||
return nil, connect.NewError(connect.CodeInternal, err)
|
||||
}
|
||||
diffs := truncateCommitMessageDiffs(req.Msg.GetDiffs(), commitMessageDiffTotalLimit, commitMessageSingleDiffLimit)
|
||||
if len(diffs) == 0 {
|
||||
return nil, commitMessageConnectError(recorder, connect.CodeInvalidArgument, fmt.Errorf("diffs are required"))
|
||||
}
|
||||
if service.provider == nil {
|
||||
return nil, commitMessageConnectError(recorder, connect.CodeInternal, fmt.Errorf("provider gateway is not initialized"))
|
||||
}
|
||||
messages, err := buildCommitMessagePrompt(req.Msg, diffs)
|
||||
if err != nil {
|
||||
return nil, commitMessageConnectError(recorder, connect.CodeInternal, err)
|
||||
}
|
||||
modelCallID := requestID + "-model"
|
||||
modelID, modelSource, lastAgentModelHash := service.resolveCommitMessageModelID(ctx)
|
||||
accumulated := ""
|
||||
artifactPaths := &modeladapter.LLMArtifactPaths{}
|
||||
err = service.provider.StartStream(ctx, ProviderRequest{
|
||||
RequestID: requestID,
|
||||
RunID: requestID,
|
||||
ModelCallID: modelCallID,
|
||||
ModelID: modelID,
|
||||
Mode: agentv1.AgentMode_AGENT_MODE_AGENT,
|
||||
Messages: messages,
|
||||
Tools: nil,
|
||||
MaxTokens: commitMessageMaxOutputTokens,
|
||||
CompileSummary: fmt.Sprintf("generate commit message diffs=%d previous_commits=%d model_source=%s last_agent_model_hash=%s", len(diffs), len(req.Msg.GetPreviousCommitMessages()), modelSource, lastAgentModelHash),
|
||||
Observer: recorder,
|
||||
ArtifactPaths: artifactPaths,
|
||||
}, func(event modeladapter.ModelEvent) error {
|
||||
switch event.Kind {
|
||||
case modeladapter.ModelEventKindTextDelta:
|
||||
accumulated += event.Text
|
||||
return nil
|
||||
case modeladapter.ModelEventKindThinkingDelta, modeladapter.ModelEventKindThinkingCompleted, modeladapter.ModelEventKindTurnFinished:
|
||||
return nil
|
||||
case modeladapter.ModelEventKindToolLikeCompleted, modeladapter.ModelEventKindPartialToolCall, modeladapter.ModelEventKindToolCallDelta:
|
||||
return errCommitMessageToolInvocation
|
||||
case modeladapter.ModelEventKindProviderError:
|
||||
if event.Err != nil {
|
||||
return providerTerminalError{cause: event.Err}
|
||||
}
|
||||
return providerTerminalError{cause: fmt.Errorf("provider error")}
|
||||
default:
|
||||
return nil
|
||||
}
|
||||
})
|
||||
if err != nil {
|
||||
if errors.Is(err, errCommitMessageToolInvocation) {
|
||||
return nil, commitMessageConnectError(recorder, connect.CodeInternal, errCommitMessageToolInvocation)
|
||||
}
|
||||
return nil, commitMessageConnectError(recorder, connect.CodeUnknown, err)
|
||||
}
|
||||
commitMessage := cleanGeneratedCommitMessage(accumulated)
|
||||
if firstCommitMessageLine(commitMessage) == "" {
|
||||
return nil, commitMessageConnectError(recorder, connect.CodeInternal, fmt.Errorf("generated commit message is empty"))
|
||||
}
|
||||
if _, err := recorder.appendEvent("final_response", map[string]any{
|
||||
"request_id": requestID,
|
||||
"model_call_id": modelCallID,
|
||||
"model_source": modelSource,
|
||||
"model_id": modelID,
|
||||
"raw_text": accumulated,
|
||||
"commit_message": commitMessage,
|
||||
"artifact_request": artifactPaths.RequestPath,
|
||||
"artifact_sse": artifactPaths.ResponsePath,
|
||||
"artifact_summary": artifactPaths.SummaryPath,
|
||||
}); err != nil {
|
||||
log.Printf("forwarder failed to write commit message final log request_id=%s error=%v", requestID, err)
|
||||
}
|
||||
return connect.NewResponse(&aiserverv1.WriteGitCommitMessageResponse{
|
||||
CommitMessage: commitMessage,
|
||||
}), nil
|
||||
}
|
||||
|
||||
func buildCommitMessagePrompt(req *aiserverv1.WriteGitCommitMessageRequest, diffs []string) ([]modeladapter.Message, error) {
|
||||
system, err := promptassets.ReadCommitPrompt()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if strings.TrimSpace(system) == "" {
|
||||
return nil, fmt.Errorf("commit prompt asset is empty")
|
||||
}
|
||||
sections := []string{"Generate a Git commit message for the following changes."}
|
||||
previous := truncateCommitMessagePreviousCommits(req.GetPreviousCommitMessages(), commitMessagePreviousCommitLimit)
|
||||
if len(previous) > 0 {
|
||||
sections = append(sections, "Recent commit messages:\n"+strings.Join(previous, "\n"))
|
||||
}
|
||||
if contextJSON, err := marshalCommitMessageExplicitContext(req.GetExplicitContext()); err != nil {
|
||||
return nil, err
|
||||
} else if contextJSON != "" {
|
||||
sections = append(sections, "Explicit context:\n"+contextJSON)
|
||||
}
|
||||
diffSections := make([]string, 0, len(diffs))
|
||||
for index, diff := range diffs {
|
||||
diffSections = append(diffSections, fmt.Sprintf("--- Diff %d ---\n%s", index+1, diff))
|
||||
}
|
||||
sections = append(sections, "Diffs:\n"+strings.Join(diffSections, "\n\n"))
|
||||
return []modeladapter.Message{
|
||||
{Role: "system", Content: system},
|
||||
{Role: "user", Content: strings.Join(sections, "\n\n")},
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (service *Service) commitMessageHistoryRoot() string {
|
||||
if service == nil || service.store == nil {
|
||||
return ""
|
||||
}
|
||||
return strings.TrimSpace(service.store.HistoryDir())
|
||||
}
|
||||
|
||||
func (service *Service) resolveCommitMessageModelID(ctx context.Context) (string, string, string) {
|
||||
if service == nil || service.modelMemory == nil || service.resolver == nil {
|
||||
return "", "default_fallback", ""
|
||||
}
|
||||
hash := strings.TrimSpace(service.modelMemory.LastAgentModelHash())
|
||||
if hash == "" {
|
||||
return "", "default_fallback", ""
|
||||
}
|
||||
channel, err := service.resolver.SelectChannelForModel(ctx, hash)
|
||||
if err != nil || channel == nil || strings.TrimSpace(channel.ID) != hash {
|
||||
if err != nil {
|
||||
log.Printf("forwarder commit message ignored invalid last agent model hash=%s error=%v", hash, err)
|
||||
}
|
||||
return "", "default_fallback", hash
|
||||
}
|
||||
return hash, "last_agent_model_hash", hash
|
||||
}
|
||||
|
||||
func recordCommitMessageIncomingRequest(recorder *commitMessageLogRecorder, requestID string, request *aiserverv1.WriteGitCommitMessageRequest) error {
|
||||
_, err := recorder.appendEvent("incoming_request", map[string]any{
|
||||
"request_id": requestID,
|
||||
})
|
||||
return err
|
||||
}
|
||||
|
||||
func commitMessageConnectError(recorder *commitMessageLogRecorder, code connect.Code, err error) error {
|
||||
if err == nil {
|
||||
err = fmt.Errorf("commit message generation failed")
|
||||
}
|
||||
if _, logErr := recorder.appendEvent("error", map[string]any{
|
||||
"code": code.String(),
|
||||
"error": err.Error(),
|
||||
}); logErr != nil {
|
||||
log.Printf("forwarder failed to write commit message error log error=%v original_error=%v", logErr, err)
|
||||
}
|
||||
return connect.NewError(code, err)
|
||||
}
|
||||
|
||||
func truncateCommitMessageDiffs(input []string, totalLimit int, singleLimit int) []string {
|
||||
result := make([]string, 0, len(input))
|
||||
remaining := totalLimit
|
||||
for _, raw := range input {
|
||||
diff := strings.TrimSpace(raw)
|
||||
if diff == "" || remaining <= 0 {
|
||||
continue
|
||||
}
|
||||
limit := singleLimit
|
||||
if remaining < limit {
|
||||
limit = remaining
|
||||
}
|
||||
truncated := truncateCommitMessageText(diff, limit)
|
||||
if truncated == "" {
|
||||
continue
|
||||
}
|
||||
result = append(result, truncated)
|
||||
remaining -= len([]rune(truncated))
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
func truncateCommitMessagePreviousCommits(input []string, limit int) []string {
|
||||
if limit <= 0 {
|
||||
return nil
|
||||
}
|
||||
result := make([]string, 0, limit)
|
||||
for _, raw := range input {
|
||||
value := strings.TrimSpace(raw)
|
||||
if value == "" {
|
||||
continue
|
||||
}
|
||||
result = append(result, "- "+value)
|
||||
if len(result) >= limit {
|
||||
break
|
||||
}
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
func marshalCommitMessageExplicitContext(contextValue *aiserverv1.ExplicitContext) (string, error) {
|
||||
if contextValue == nil {
|
||||
return "", nil
|
||||
}
|
||||
payload, err := protojson.MarshalOptions{UseProtoNames: true}.Marshal(contextValue)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("marshal explicit context failed: %w", err)
|
||||
}
|
||||
return truncateCommitMessageText(string(payload), commitMessageExplicitContextLimit), nil
|
||||
}
|
||||
|
||||
func truncateCommitMessageText(value string, limit int) string {
|
||||
trimmed := strings.TrimSpace(value)
|
||||
if trimmed == "" || limit <= 0 {
|
||||
return ""
|
||||
}
|
||||
runes := []rune(trimmed)
|
||||
if len(runes) <= limit {
|
||||
return trimmed
|
||||
}
|
||||
return strings.TrimSpace(string(runes[:limit])) + "\n...[truncated]"
|
||||
}
|
||||
|
||||
func cleanGeneratedCommitMessage(value string) string {
|
||||
result := strings.TrimSpace(value)
|
||||
result = stripCommitMessageCodeFence(result)
|
||||
prefixes := []string{
|
||||
"commit message:",
|
||||
"git commit message:",
|
||||
"message:",
|
||||
}
|
||||
for {
|
||||
lower := strings.ToLower(strings.TrimSpace(result))
|
||||
matched := ""
|
||||
for _, prefix := range prefixes {
|
||||
if strings.HasPrefix(lower, prefix) {
|
||||
matched = prefix
|
||||
break
|
||||
}
|
||||
}
|
||||
if matched == "" {
|
||||
break
|
||||
}
|
||||
result = strings.TrimSpace(result[len(matched):])
|
||||
}
|
||||
return strings.TrimSpace(result)
|
||||
}
|
||||
|
||||
func stripCommitMessageCodeFence(value string) string {
|
||||
trimmed := strings.TrimSpace(value)
|
||||
if !strings.HasPrefix(trimmed, "```") {
|
||||
return trimmed
|
||||
}
|
||||
lines := strings.Split(trimmed, "\n")
|
||||
if len(lines) < 2 {
|
||||
return trimmed
|
||||
}
|
||||
lines = lines[1:]
|
||||
if len(lines) > 0 && strings.HasPrefix(strings.TrimSpace(lines[len(lines)-1]), "```") {
|
||||
lines = lines[:len(lines)-1]
|
||||
}
|
||||
return strings.TrimSpace(strings.Join(lines, "\n"))
|
||||
}
|
||||
|
||||
func firstCommitMessageLine(value string) string {
|
||||
for _, line := range strings.Split(value, "\n") {
|
||||
line = strings.TrimSpace(line)
|
||||
if line != "" {
|
||||
return line
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
@@ -0,0 +1,23 @@
|
||||
package forwarder
|
||||
|
||||
type commitMessageLogRecorder struct{}
|
||||
|
||||
func newCommitMessageLogRecorder(_ string, _ string) (*commitMessageLogRecorder, error) {
|
||||
return &commitMessageLogRecorder{}, nil
|
||||
}
|
||||
|
||||
func (recorder *commitMessageLogRecorder) RecordLLMRequest(_ string, _ string, _ string, _ map[string]any) (string, error) {
|
||||
return "", nil
|
||||
}
|
||||
|
||||
func (recorder *commitMessageLogRecorder) AppendLLMResponseChunk(_ string, _ string, _ string, _ string) (string, error) {
|
||||
return "", nil
|
||||
}
|
||||
|
||||
func (recorder *commitMessageLogRecorder) RecordLLMSummary(_ string, _ string, _ string, _ map[string]any) (string, error) {
|
||||
return "", nil
|
||||
}
|
||||
|
||||
func (recorder *commitMessageLogRecorder) appendEvent(_ string, _ map[string]any) (string, error) {
|
||||
return "", nil
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,255 @@
|
||||
// compiler.go 负责把固定 prompt、自然历史和 tool catalog 编译成 provider 请求。
|
||||
package forwarder
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"cursor/gen/agentv1"
|
||||
modeladapter "cursor/internal/backend/agent/model"
|
||||
promptassets "cursor/prompt"
|
||||
)
|
||||
|
||||
type PromptCompiler interface {
|
||||
Compile(conversation *ConversationFile, mode agentv1.AgentMode, latestUserText string, modelName string) (CompiledConversation, error)
|
||||
DerivePromptContexts(conversation *ConversationFile, mode agentv1.AgentMode, latestUserText string) ([]PromptContextMessage, error)
|
||||
}
|
||||
|
||||
type DefaultPromptCompiler struct {
|
||||
projector *HistoryProjector
|
||||
catalog ToolCatalog
|
||||
reminders ReminderInjector
|
||||
rules *UserRuleStore
|
||||
}
|
||||
|
||||
// NewPromptCompiler 创建默认 prompt 编译器。
|
||||
func NewPromptCompiler(projector *HistoryProjector, catalog ToolCatalog, reminders ReminderInjector, rules *UserRuleStore) *DefaultPromptCompiler {
|
||||
return &DefaultPromptCompiler{
|
||||
projector: projector,
|
||||
catalog: catalog,
|
||||
reminders: reminders,
|
||||
rules: rules,
|
||||
}
|
||||
}
|
||||
|
||||
// Compile 生成当前 turn 应发送给 provider 的消息和工具集合。
|
||||
func (compiler *DefaultPromptCompiler) Compile(conversation *ConversationFile, mode agentv1.AgentMode, latestUserText string, modelName string) (CompiledConversation, error) {
|
||||
if compiler == nil || compiler.projector == nil || compiler.catalog == nil {
|
||||
return CompiledConversation{}, fmt.Errorf("prompt compiler dependencies are not initialized")
|
||||
}
|
||||
normalizedMode, err := validateSupportedActiveMode(mode)
|
||||
if err != nil {
|
||||
return CompiledConversation{}, err
|
||||
}
|
||||
subagentTypeName := ""
|
||||
if conversation != nil {
|
||||
subagentTypeName = conversation.SubagentTypeName
|
||||
}
|
||||
assetMode, err := promptAssetModeForConversation(normalizedMode, subagentTypeName)
|
||||
if err != nil {
|
||||
return CompiledConversation{}, err
|
||||
}
|
||||
systemPrompt, err := promptassets.ReadPrompt(assetMode)
|
||||
if err != nil {
|
||||
return CompiledConversation{}, err
|
||||
}
|
||||
tools, _, err := compiler.catalog.Load(normalizedMode, subagentTypeName)
|
||||
if err != nil {
|
||||
return CompiledConversation{}, err
|
||||
}
|
||||
replayMessages, err := compiler.projector.ProjectPromptReplay(conversation)
|
||||
if err != nil {
|
||||
return CompiledConversation{}, err
|
||||
}
|
||||
sharedRulesPrompt := ""
|
||||
sharedRuleCount := 0
|
||||
sharedRuleTotal := 0
|
||||
if compiler.rules != nil && normalizedMode != agentv1.AgentMode_AGENT_MODE_DEBUG {
|
||||
sharedRulesPrompt, sharedRuleTotal, sharedRuleCount, err = compiler.rules.BuildSystemPromptSection()
|
||||
if err != nil {
|
||||
return CompiledConversation{}, err
|
||||
}
|
||||
}
|
||||
messages := make([]modeladapter.Message, 0, len(replayMessages)+1)
|
||||
systemParts := []string{sanitizePromptAsset(systemPrompt, modelName)}
|
||||
if strings.TrimSpace(sharedRulesPrompt) != "" {
|
||||
systemParts = append(systemParts, sharedRulesPrompt)
|
||||
}
|
||||
systemText := strings.TrimSpace(strings.Join(filterNonEmpty(systemParts), "\n\n"))
|
||||
if systemText != "" {
|
||||
messages = append(messages, modeladapter.Message{
|
||||
Role: "system",
|
||||
Content: systemText,
|
||||
})
|
||||
}
|
||||
stableReplayCount, err := compiler.stableReplayMessageCount(conversation, replayMessages)
|
||||
if err != nil {
|
||||
return CompiledConversation{}, err
|
||||
}
|
||||
messages = append(messages, replayMessages...)
|
||||
return CompiledConversation{
|
||||
Mode: normalizedMode,
|
||||
Messages: messages,
|
||||
StableMessageCount: stableReplayCount,
|
||||
Tools: tools,
|
||||
CompileSummary: fmt.Sprintf("mode=%s asset_mode=%s child=%t messages=%d tools=%d shared_rules_total=%d shared_rules_deduped=%d", normalizedMode.String(), string(assetMode), isChildConversationSubagentTypeName(subagentTypeName), len(messages), len(tools), sharedRuleTotal, sharedRuleCount),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (compiler *DefaultPromptCompiler) DerivePromptContexts(conversation *ConversationFile, mode agentv1.AgentMode, latestUserText string) ([]PromptContextMessage, error) {
|
||||
if compiler == nil || compiler.projector == nil || compiler.catalog == nil || compiler.reminders == nil {
|
||||
return nil, fmt.Errorf("prompt compiler dependencies are not initialized")
|
||||
}
|
||||
normalizedMode, err := validateSupportedActiveMode(mode)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
subagentTypeName := ""
|
||||
if conversation != nil {
|
||||
subagentTypeName = conversation.SubagentTypeName
|
||||
}
|
||||
_, toolNames, err := compiler.catalog.Load(normalizedMode, subagentTypeName)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
replayMessages, err := compiler.projector.ProjectPromptReplay(conversation)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
structuredStatePromptContexts, structuredStateTailMessages, err := buildStructuredStatePromptContexts(conversation)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
promptReminders := compiler.reminders.Inject(normalizedMode, conversation, replayMessages, latestUserText, toolNames)
|
||||
candidates := make([]PromptContextMessage, 0, len(structuredStatePromptContexts)+len(structuredStateTailMessages)+len(promptReminders.PromptContexts)+len(promptReminders.TailMessages))
|
||||
candidates = append(candidates, structuredStatePromptContexts...)
|
||||
for _, message := range structuredStateTailMessages {
|
||||
candidates = append(candidates, newPromptContextMessage(promptContextSourceStructuredTodoReminder, message, true))
|
||||
}
|
||||
candidates = append(candidates, promptReminders.PromptContexts...)
|
||||
for index, message := range promptReminders.TailMessages {
|
||||
candidates = append(candidates, newPromptContextMessage(fmt.Sprintf("tail_reminder/%d", index), message, true))
|
||||
}
|
||||
for index := range candidates {
|
||||
candidates[index].Persist = true
|
||||
}
|
||||
return filterCurrentTurnPromptContexts(conversation, candidates), nil
|
||||
}
|
||||
|
||||
func (compiler *DefaultPromptCompiler) stableReplayMessageCount(conversation *ConversationFile, replayMessages []modeladapter.Message) (int, error) {
|
||||
if compiler == nil || compiler.projector == nil || conversation == nil || len(replayMessages) == 0 {
|
||||
return 0, nil
|
||||
}
|
||||
currentTurnSeq := conversation.CurrentTurnSeq
|
||||
if currentTurnSeq <= 0 {
|
||||
currentTurnSeq = conversation.NextTurnSeq - 1
|
||||
}
|
||||
stableCount := 0
|
||||
if currentTurnSeq > 0 {
|
||||
stableConversation := cloneConversationFile(conversation)
|
||||
stableConversation.Entries = stableReplayEntriesBeforeTurn(conversation.Entries, currentTurnSeq)
|
||||
stableMessages, err := compiler.projector.ProjectPromptReplay(stableConversation)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
stableCount = len(stableMessages)
|
||||
}
|
||||
if requestPrefixReplayCount := replayMessageCountFromRequestPrefix(conversation); requestPrefixReplayCount > stableCount {
|
||||
stableCount = requestPrefixReplayCount
|
||||
}
|
||||
if stableCount > len(replayMessages) {
|
||||
return len(replayMessages), nil
|
||||
}
|
||||
return stableCount, nil
|
||||
}
|
||||
|
||||
func replayMessageCountFromRequestPrefix(conversation *ConversationFile) int {
|
||||
if conversation == nil || conversation.LatestRequestPrefix == nil {
|
||||
return 0
|
||||
}
|
||||
requestID := strings.TrimSpace(conversation.CurrentRequestID)
|
||||
if requestID == "" || strings.TrimSpace(conversation.LatestRequestPrefix.RequestID) != requestID {
|
||||
return 0
|
||||
}
|
||||
return conversation.LatestRequestPrefix.ReplayMessageCount
|
||||
}
|
||||
|
||||
func stableReplayEntriesBeforeTurn(entries []HistoryEntry, currentTurnSeq int64) []HistoryEntry {
|
||||
if len(entries) == 0 || currentTurnSeq <= 0 {
|
||||
return nil
|
||||
}
|
||||
filtered := make([]HistoryEntry, 0, len(entries))
|
||||
for _, entry := range entries {
|
||||
if entry.TurnSeq > 0 && entry.TurnSeq >= currentTurnSeq {
|
||||
continue
|
||||
}
|
||||
filtered = append(filtered, entry)
|
||||
}
|
||||
return filtered
|
||||
}
|
||||
|
||||
// filterNonEmpty 过滤掉空白字符串,便于安全拼接 system prompt 片段。
|
||||
func filterNonEmpty(items []string) []string {
|
||||
filtered := make([]string, 0, len(items))
|
||||
for _, item := range items {
|
||||
if strings.TrimSpace(item) != "" {
|
||||
filtered = append(filtered, strings.TrimSpace(item))
|
||||
}
|
||||
}
|
||||
return filtered
|
||||
}
|
||||
|
||||
func mergeAdjacentPlainUserMessages(messages []modeladapter.Message) []modeladapter.Message {
|
||||
if len(messages) == 0 {
|
||||
return nil
|
||||
}
|
||||
merged := make([]modeladapter.Message, 0, len(messages))
|
||||
for _, message := range messages {
|
||||
text, ok := plainUserMessageText(message)
|
||||
if !ok {
|
||||
merged = append(merged, message)
|
||||
continue
|
||||
}
|
||||
if len(merged) == 0 {
|
||||
message.Role = "user"
|
||||
message.Content = text
|
||||
merged = append(merged, message)
|
||||
continue
|
||||
}
|
||||
last := &merged[len(merged)-1]
|
||||
if lastText, ok := plainUserMessageText(*last); ok {
|
||||
last.Role = "user"
|
||||
last.Content = lastText + "\n\n" + text
|
||||
continue
|
||||
}
|
||||
message.Role = "user"
|
||||
message.Content = text
|
||||
merged = append(merged, message)
|
||||
}
|
||||
return merged
|
||||
}
|
||||
|
||||
func plainUserMessageText(message modeladapter.Message) (string, bool) {
|
||||
if strings.TrimSpace(message.Role) != "user" {
|
||||
return "", false
|
||||
}
|
||||
if len(message.ContentParts) > 0 || len(message.ToolCalls) > 0 {
|
||||
return "", false
|
||||
}
|
||||
if strings.TrimSpace(message.ToolCallID) != "" || strings.TrimSpace(message.Name) != "" {
|
||||
return "", false
|
||||
}
|
||||
if strings.TrimSpace(message.ReasoningContent) != "" ||
|
||||
strings.TrimSpace(message.ReasoningSignature) != "" ||
|
||||
strings.TrimSpace(message.ReasoningSignatureSource) != "" ||
|
||||
strings.TrimSpace(message.OpenAIResponsesReasoningID) != "" ||
|
||||
strings.TrimSpace(message.OpenAIResponsesReasoningStatus) != "" ||
|
||||
len(message.OpenAIResponsesReasoningSummary) > 0 {
|
||||
return "", false
|
||||
}
|
||||
text := strings.TrimSpace(message.Content)
|
||||
if text == "" {
|
||||
return "", false
|
||||
}
|
||||
return text, true
|
||||
}
|
||||
@@ -0,0 +1,36 @@
|
||||
package forwarder
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
runtimecore "cursor/internal/backend/agent/core"
|
||||
)
|
||||
|
||||
func (service *Service) sanitizeCreatePlanInvocationForCurrentPlan(stream *ActiveStream, invocation runtimecore.ToolInvocation) (runtimecore.ToolInvocation, error) {
|
||||
if strings.TrimSpace(invocation.ToolName) != "CreatePlan" {
|
||||
return invocation, nil
|
||||
}
|
||||
args, err := runtimecore.DecodeCreatePlanArgsJSON(invocation.ArgsJSON)
|
||||
if err != nil {
|
||||
return invocation, newRecoverableToolInvocationError(fmt.Errorf("decode CreatePlan args failed: %w", err))
|
||||
}
|
||||
if strings.TrimSpace(args.GetName()) == "" {
|
||||
return invocation, nil
|
||||
}
|
||||
conversation, _, _, err := service.snapshotCheckpointConversation(stream)
|
||||
if err != nil {
|
||||
return invocation, err
|
||||
}
|
||||
if !hasCurrentPlan(conversation) {
|
||||
return invocation, nil
|
||||
}
|
||||
args.Name = ""
|
||||
sanitized, err := json.Marshal(args)
|
||||
if err != nil {
|
||||
return invocation, newRecoverableToolInvocationError(fmt.Errorf("encode sanitized CreatePlan args failed: %w", err))
|
||||
}
|
||||
invocation.ArgsJSON = sanitized
|
||||
return invocation, nil
|
||||
}
|
||||
@@ -0,0 +1,394 @@
|
||||
package forwarder
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"google.golang.org/protobuf/encoding/protojson"
|
||||
"google.golang.org/protobuf/proto"
|
||||
|
||||
"cursor/gen/agentv1"
|
||||
)
|
||||
|
||||
type debugLogConfig interface {
|
||||
IsObservabilityLogEnabled(context.Context) bool
|
||||
}
|
||||
|
||||
type debugRecorder struct {
|
||||
historyRoot string
|
||||
broker *StreamBroker
|
||||
config debugLogConfig
|
||||
mu sync.Mutex
|
||||
}
|
||||
|
||||
func newDebugRecorder(historyRoot string, broker *StreamBroker, config debugLogConfig) *debugRecorder {
|
||||
return &debugRecorder{
|
||||
historyRoot: strings.TrimSpace(historyRoot),
|
||||
broker: broker,
|
||||
config: config,
|
||||
}
|
||||
}
|
||||
|
||||
func (recorder *debugRecorder) enabled(ctx context.Context) bool {
|
||||
if recorder == nil || recorder.config == nil {
|
||||
return false
|
||||
}
|
||||
if ctx == nil {
|
||||
ctx = context.Background()
|
||||
}
|
||||
return recorder.config.IsObservabilityLogEnabled(ctx)
|
||||
}
|
||||
|
||||
func (recorder *debugRecorder) LogBidiRaw(ctx context.Context, requestID string, conversationID string, appendSeqno int64, dataHex string, status string, extra map[string]any) {
|
||||
if !recorder.enabled(ctx) {
|
||||
return
|
||||
}
|
||||
event := recorder.baseEvent("bidi_raw", requestID, conversationID)
|
||||
event["direction"] = "client_to_backend"
|
||||
event["procedure"] = "/aiserver.v1.BidiService/BidiAppend"
|
||||
event["append_seqno"] = appendSeqno
|
||||
event["status"] = strings.TrimSpace(status)
|
||||
event["data_hex"] = dataHex
|
||||
event["data_len"] = len(dataHex)
|
||||
for key, value := range extra {
|
||||
event[key] = value
|
||||
}
|
||||
recorder.appendJSONL(ctx, requestID, conversationID, "bidi.raw.jsonl", event)
|
||||
}
|
||||
|
||||
func (recorder *debugRecorder) LogBidiDecoded(ctx context.Context, requestID string, conversationID string, appendSeqno int64, clientKind string, message *agentv1.AgentClientMessage, intent InboundIntent, extra map[string]any) {
|
||||
if !recorder.enabled(ctx) {
|
||||
return
|
||||
}
|
||||
event := recorder.baseEvent("bidi_decoded", requestID, conversationID)
|
||||
event["schema_version"] = 2
|
||||
event["append_seqno"] = appendSeqno
|
||||
event["client_kind"] = strings.TrimSpace(clientKind)
|
||||
event["message_case"] = agentClientMessageCase(message)
|
||||
event["message"] = protoJSONDebugPayload(message)
|
||||
event["intent"] = inboundIntentDebugPayload(intent)
|
||||
if requestedModel := requestedModelDebugPayload(message); requestedModel != nil {
|
||||
event["requested_model"] = requestedModel
|
||||
}
|
||||
if actionCase := conversationActionCase(message); actionCase != "" {
|
||||
event["conversation_action"] = actionCase
|
||||
}
|
||||
for key, value := range extra {
|
||||
event[key] = value
|
||||
}
|
||||
recorder.appendJSONL(ctx, requestID, firstNonEmpty(conversationID, intent.ConversationID), "bidi.decoded.jsonl", event)
|
||||
}
|
||||
|
||||
func (recorder *debugRecorder) LogRuntime(ctx context.Context, requestID string, conversationID string, eventName string, fields map[string]any) {
|
||||
if !recorder.enabled(ctx) {
|
||||
return
|
||||
}
|
||||
event := recorder.baseEvent("runtime", requestID, conversationID)
|
||||
event["event"] = strings.TrimSpace(eventName)
|
||||
for key, value := range fields {
|
||||
event[key] = value
|
||||
}
|
||||
recorder.appendJSONL(ctx, requestID, conversationID, "runtime.jsonl", event)
|
||||
}
|
||||
|
||||
func (recorder *debugRecorder) LogRunSSE(ctx context.Context, requestID string, conversationID string, eventName string, fields map[string]any) {
|
||||
if !recorder.enabled(ctx) {
|
||||
return
|
||||
}
|
||||
event := recorder.baseEvent("runsse", requestID, conversationID)
|
||||
event["event"] = strings.TrimSpace(eventName)
|
||||
for key, value := range fields {
|
||||
event[key] = value
|
||||
}
|
||||
recorder.appendJSONL(ctx, requestID, conversationID, "runsse.jsonl", event)
|
||||
}
|
||||
|
||||
func (recorder *debugRecorder) LogProvider(ctx context.Context, requestID string, conversationID string, eventName string, fields map[string]any) {
|
||||
if !recorder.enabled(ctx) {
|
||||
return
|
||||
}
|
||||
event := recorder.baseEvent("provider", requestID, conversationID)
|
||||
event["event"] = strings.TrimSpace(eventName)
|
||||
for key, value := range fields {
|
||||
event[key] = value
|
||||
}
|
||||
recorder.appendJSONL(ctx, requestID, conversationID, "provider.jsonl", event)
|
||||
}
|
||||
|
||||
func (recorder *debugRecorder) LogProviderArtifact(ctx context.Context, requestID string, conversationID string, modelCallID string, eventName string, payload map[string]any) {
|
||||
if !recorder.enabled(ctx) {
|
||||
return
|
||||
}
|
||||
recorder.LogProvider(ctx, requestID, conversationID, eventName, map[string]any{
|
||||
"model_call_id": strings.TrimSpace(modelCallID),
|
||||
"payload": payload,
|
||||
})
|
||||
}
|
||||
|
||||
func (recorder *debugRecorder) baseEvent(layer string, requestID string, conversationID string) map[string]any {
|
||||
resolvedConversationID := firstNonEmpty(strings.TrimSpace(conversationID), recorder.conversationIDForRequest(requestID))
|
||||
return map[string]any{
|
||||
"schema_version": 1,
|
||||
"at": time.Now().UTC().Format(time.RFC3339Nano),
|
||||
"layer": strings.TrimSpace(layer),
|
||||
"request_id": strings.TrimSpace(requestID),
|
||||
"conversation_id": resolvedConversationID,
|
||||
}
|
||||
}
|
||||
|
||||
func (recorder *debugRecorder) appendJSONL(ctx context.Context, requestID string, conversationID string, filename string, event map[string]any) {
|
||||
if !recorder.enabled(ctx) || len(event) == 0 {
|
||||
return
|
||||
}
|
||||
dir := recorder.debugDir(requestID, conversationID)
|
||||
if strings.TrimSpace(dir) == "" {
|
||||
return
|
||||
}
|
||||
payload, err := json.Marshal(event)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
recorder.mu.Lock()
|
||||
defer recorder.mu.Unlock()
|
||||
if err := os.MkdirAll(dir, 0o755); err != nil {
|
||||
return
|
||||
}
|
||||
file, err := os.OpenFile(filepath.Join(dir, filename), os.O_CREATE|os.O_WRONLY|os.O_APPEND, 0o644)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
defer file.Close()
|
||||
_, _ = file.Write(append(payload, '\n'))
|
||||
}
|
||||
|
||||
func (recorder *debugRecorder) debugDir(requestID string, conversationID string) string {
|
||||
if recorder == nil || strings.TrimSpace(recorder.historyRoot) == "" {
|
||||
return ""
|
||||
}
|
||||
conversationID = firstNonEmpty(strings.TrimSpace(conversationID), recorder.conversationIDForRequest(requestID))
|
||||
if conversationID != "" && conversationID != "unknown" {
|
||||
return filepath.Join(recorder.historyRoot, sanitizeArtifactName(conversationID), "debug")
|
||||
}
|
||||
requestID = firstNonEmpty(strings.TrimSpace(requestID), "unknown")
|
||||
return filepath.Join(recorder.historyRoot, "_debug", "orphan", sanitizeArtifactName(requestID))
|
||||
}
|
||||
|
||||
func (recorder *debugRecorder) conversationIDForRequest(requestID string) string {
|
||||
if recorder == nil || recorder.broker == nil {
|
||||
return ""
|
||||
}
|
||||
stream, ok := recorder.broker.Get(requestID)
|
||||
if !ok || stream == nil {
|
||||
return ""
|
||||
}
|
||||
stream.mu.Lock()
|
||||
defer stream.mu.Unlock()
|
||||
return strings.TrimSpace(stream.ConversationID)
|
||||
}
|
||||
|
||||
func agentClientMessageCase(message *agentv1.AgentClientMessage) string {
|
||||
if message == nil {
|
||||
return ""
|
||||
}
|
||||
switch message.GetMessage().(type) {
|
||||
case *agentv1.AgentClientMessage_RunRequest:
|
||||
return "run_request"
|
||||
case *agentv1.AgentClientMessage_PrewarmRequest:
|
||||
return "prewarm_request"
|
||||
case *agentv1.AgentClientMessage_ConversationAction:
|
||||
return "conversation_action"
|
||||
case *agentv1.AgentClientMessage_ExecClientMessage:
|
||||
return "exec_client_message"
|
||||
case *agentv1.AgentClientMessage_ExecClientControlMessage:
|
||||
return "exec_client_control_message"
|
||||
case *agentv1.AgentClientMessage_InteractionResponse:
|
||||
return "interaction_response"
|
||||
case *agentv1.AgentClientMessage_ClientHeartbeat:
|
||||
return "client_heartbeat"
|
||||
case *agentv1.AgentClientMessage_KvClientMessage:
|
||||
return "kv_client_message"
|
||||
default:
|
||||
return fmt.Sprintf("%T", message.GetMessage())
|
||||
}
|
||||
}
|
||||
|
||||
func agentServerMessageCase(message *agentv1.AgentServerMessage) string {
|
||||
if message == nil {
|
||||
return ""
|
||||
}
|
||||
switch message.GetMessage().(type) {
|
||||
case *agentv1.AgentServerMessage_InteractionUpdate:
|
||||
return "interaction_update"
|
||||
case *agentv1.AgentServerMessage_ExecServerMessage:
|
||||
return "exec_server_message"
|
||||
case *agentv1.AgentServerMessage_ExecServerControlMessage:
|
||||
return "exec_server_control_message"
|
||||
case *agentv1.AgentServerMessage_ConversationCheckpointUpdate:
|
||||
return "conversation_checkpoint_update"
|
||||
case *agentv1.AgentServerMessage_KvServerMessage:
|
||||
return "kv_server_message"
|
||||
case *agentv1.AgentServerMessage_InteractionQuery:
|
||||
return "interaction_query"
|
||||
default:
|
||||
return fmt.Sprintf("%T", message.GetMessage())
|
||||
}
|
||||
}
|
||||
|
||||
func conversationActionCase(message *agentv1.AgentClientMessage) string {
|
||||
if message == nil {
|
||||
return ""
|
||||
}
|
||||
action := message.GetConversationAction()
|
||||
if action == nil && message.GetRunRequest() != nil {
|
||||
action = message.GetRunRequest().GetAction()
|
||||
}
|
||||
if action == nil {
|
||||
return ""
|
||||
}
|
||||
return conversationActionKind(action)
|
||||
}
|
||||
|
||||
func requestedModelDebugPayload(message *agentv1.AgentClientMessage) map[string]any {
|
||||
if message == nil {
|
||||
return nil
|
||||
}
|
||||
if runRequest := message.GetRunRequest(); runRequest != nil {
|
||||
return requestedModelPayload(runRequest.GetRequestedModel())
|
||||
}
|
||||
if prewarm := message.GetPrewarmRequest(); prewarm != nil {
|
||||
return requestedModelPayload(prewarm.GetRequestedModel())
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func requestedModelPayload(model *agentv1.RequestedModel) map[string]any {
|
||||
if model == nil {
|
||||
return nil
|
||||
}
|
||||
parameters := make([]map[string]string, 0, len(model.GetParameters()))
|
||||
for _, parameter := range model.GetParameters() {
|
||||
if parameter == nil {
|
||||
continue
|
||||
}
|
||||
parameters = append(parameters, map[string]string{
|
||||
"id": parameter.GetId(),
|
||||
"value": parameter.GetValue(),
|
||||
})
|
||||
}
|
||||
return map[string]any{
|
||||
"model_id": strings.TrimSpace(model.GetModelId()),
|
||||
"max_mode": model.GetMaxMode(),
|
||||
"built_in_model": model.GetBuiltInModel(),
|
||||
"is_variant_string_representation": model.GetIsVariantStringRepresentation(),
|
||||
"parameters": parameters,
|
||||
}
|
||||
}
|
||||
|
||||
func protoJSONDebugPayload(message proto.Message) any {
|
||||
if message == nil {
|
||||
return nil
|
||||
}
|
||||
payload, err := protojson.MarshalOptions{
|
||||
UseProtoNames: true,
|
||||
EmitUnpopulated: false,
|
||||
}.Marshal(message)
|
||||
if err != nil {
|
||||
return map[string]any{"marshal_error": err.Error()}
|
||||
}
|
||||
var decoded any
|
||||
if err := json.Unmarshal(payload, &decoded); err != nil {
|
||||
return string(payload)
|
||||
}
|
||||
return decoded
|
||||
}
|
||||
|
||||
func inboundIntentDebugPayload(intent InboundIntent) map[string]any {
|
||||
payload := map[string]any{
|
||||
"kind": strings.TrimSpace(intent.Kind),
|
||||
"request_id": strings.TrimSpace(intent.RequestID),
|
||||
"conversation_id": strings.TrimSpace(intent.ConversationID),
|
||||
"model_id": strings.TrimSpace(intent.ModelID),
|
||||
"model_name": strings.TrimSpace(intent.ModelName),
|
||||
"thinking_effort": strings.TrimSpace(intent.ThinkingEffort),
|
||||
"mode": intent.Mode.String(),
|
||||
"has_explicit_mode": intent.HasExplicitMode,
|
||||
"mode_source": string(intent.ModeSource),
|
||||
"starts_run": intent.StartsRun,
|
||||
"subagent_type_name": strings.TrimSpace(intent.SubagentTypeName),
|
||||
"cancel_reason": strings.TrimSpace(intent.CancelReason),
|
||||
"prewarm": intent.Prewarm,
|
||||
}
|
||||
if len(intent.SubagentModelOverrides) > 0 {
|
||||
payload["subagent_model_overrides"] = subagentModelOverrideSummaries(intent.SubagentModelOverrides)
|
||||
payload["subagent_model_override_count"] = len(intent.SubagentModelOverrides)
|
||||
}
|
||||
if intent.ClientMessage != nil {
|
||||
payload["client_message"] = protoJSONDebugPayload(intent.ClientMessage)
|
||||
}
|
||||
if intent.ConversationState != nil {
|
||||
payload["conversation_state"] = protoJSONDebugPayload(intent.ConversationState)
|
||||
}
|
||||
if intent.UserMessage != nil {
|
||||
payload["user_message"] = protoJSONDebugPayload(intent.UserMessage)
|
||||
}
|
||||
if intent.RequestContext != nil {
|
||||
payload["request_context"] = protoJSONDebugPayload(intent.RequestContext)
|
||||
}
|
||||
if strings.TrimSpace(intent.IgnoredReason) != "" {
|
||||
payload["ignored_reason"] = strings.TrimSpace(intent.IgnoredReason)
|
||||
payload["ignored_empty_resume"] = strings.TrimSpace(intent.IgnoredReason) == "empty_resume_without_pending_continuation"
|
||||
}
|
||||
if intent.ExecClientMessage != nil {
|
||||
payload["exec_client_message"] = protoJSONDebugPayload(intent.ExecClientMessage)
|
||||
}
|
||||
if intent.ExecClientControlMessage != nil {
|
||||
payload["exec_client_control_message"] = protoJSONDebugPayload(intent.ExecClientControlMessage)
|
||||
}
|
||||
if intent.InteractionResponse != nil {
|
||||
payload["interaction_response"] = protoJSONDebugPayload(intent.InteractionResponse)
|
||||
}
|
||||
if intent.KVClientMessage != nil {
|
||||
payload["kv_client_message"] = protoJSONDebugPayload(intent.KVClientMessage)
|
||||
}
|
||||
return payload
|
||||
}
|
||||
|
||||
func conversationActionKind(action *agentv1.ConversationAction) string {
|
||||
if action == nil {
|
||||
return ""
|
||||
}
|
||||
switch action.GetAction().(type) {
|
||||
case *agentv1.ConversationAction_UserMessageAction:
|
||||
return "user_message_action"
|
||||
case *agentv1.ConversationAction_ResumeAction:
|
||||
return "resume_action"
|
||||
case *agentv1.ConversationAction_CancelAction:
|
||||
return "cancel_action"
|
||||
case *agentv1.ConversationAction_SummarizeAction:
|
||||
return "summarize_action"
|
||||
case *agentv1.ConversationAction_ShellCommandAction:
|
||||
return "shell_command_action"
|
||||
case *agentv1.ConversationAction_StartPlanAction:
|
||||
return "start_plan_action"
|
||||
case *agentv1.ConversationAction_ExecutePlanAction:
|
||||
return "execute_plan_action"
|
||||
case *agentv1.ConversationAction_AsyncAskQuestionCompletionAction:
|
||||
return "async_ask_question_completion_action"
|
||||
case *agentv1.ConversationAction_CancelSubagentAction:
|
||||
return "cancel_subagent_action"
|
||||
case *agentv1.ConversationAction_BackgroundTaskCompletionAction:
|
||||
return "background_task_completion_action"
|
||||
case *agentv1.ConversationAction_BackgroundShellAction:
|
||||
return "background_shell_action"
|
||||
case *agentv1.ConversationAction_BackgroundSubagentAction:
|
||||
return "background_subagent_action"
|
||||
default:
|
||||
return fmt.Sprintf("%T", action.GetAction())
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,210 @@
|
||||
package forwarder
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"connectrpc.com/connect"
|
||||
|
||||
"cursor/gen/aiserverv1"
|
||||
)
|
||||
|
||||
func (service *Service) AvailableDocs(_ context.Context, req *connect.Request[aiserverv1.AvailableDocsRequest]) (*connect.Response[aiserverv1.AvailableDocsResponse], error) {
|
||||
store, err := service.requireDocsIndexStore()
|
||||
if err != nil {
|
||||
return nil, connect.NewError(connect.CodeInternal, err)
|
||||
}
|
||||
records, err := store.List("", 0)
|
||||
if err != nil {
|
||||
return nil, connect.NewError(connect.CodeInternal, err)
|
||||
}
|
||||
records = filterDocsIndexRecords(records, req.Msg)
|
||||
records, err = appendMissingAdditionalDocs(store, records, req.Msg)
|
||||
if err != nil {
|
||||
return nil, connect.NewError(connect.CodeInternal, err)
|
||||
}
|
||||
docs := make([]*aiserverv1.DocumentationInfo, 0, len(records))
|
||||
for _, record := range records {
|
||||
docs = append(docs, docsIndexDocumentationInfo(record))
|
||||
}
|
||||
return connect.NewResponse(&aiserverv1.AvailableDocsResponse{Docs: docs}), nil
|
||||
}
|
||||
|
||||
func (service *Service) DocumentationQuery(_ context.Context, req *connect.Request[aiserverv1.DocumentationQueryRequest]) (*connect.Response[aiserverv1.DocumentationQueryResponse], error) {
|
||||
store, err := service.requireDocsIndexStore()
|
||||
if err != nil {
|
||||
return nil, connect.NewError(connect.CodeInternal, err)
|
||||
}
|
||||
identifier := strings.TrimSpace(req.Msg.GetDocIdentifier())
|
||||
if identifier == "" {
|
||||
return connect.NewResponse(&aiserverv1.DocumentationQueryResponse{
|
||||
Status: aiserverv1.DocumentationQueryResponse_STATUS_NOT_FOUND,
|
||||
}), nil
|
||||
}
|
||||
record, ok, err := store.Get(identifier)
|
||||
if err != nil {
|
||||
return nil, connect.NewError(connect.CodeInternal, err)
|
||||
}
|
||||
if !ok {
|
||||
return connect.NewResponse(&aiserverv1.DocumentationQueryResponse{
|
||||
DocIdentifier: identifier,
|
||||
Status: aiserverv1.DocumentationQueryResponse_STATUS_NOT_FOUND,
|
||||
}), nil
|
||||
}
|
||||
return connect.NewResponse(&aiserverv1.DocumentationQueryResponse{
|
||||
DocIdentifier: record.Identifier,
|
||||
DocName: record.Title,
|
||||
DocChunks: docsIndexChunks(record, req.Msg.GetTopK()),
|
||||
Status: aiserverv1.DocumentationQueryResponse_STATUS_SUCCESS,
|
||||
}), nil
|
||||
}
|
||||
|
||||
func (service *Service) FetchRelevantKnowledgeForConversation(_ context.Context, req *connect.Request[aiserverv1.FetchRelevantKnowledgeForConversationRequest]) (*connect.Response[aiserverv1.FetchRelevantKnowledgeForConversationResponse], error) {
|
||||
store, err := service.requireDocsIndexStore()
|
||||
if err != nil {
|
||||
return nil, connect.NewError(connect.CodeInternal, err)
|
||||
}
|
||||
gitOrigin := strings.TrimSpace(req.Msg.GetGitOrigin())
|
||||
if gitOrigin == "" {
|
||||
return connect.NewResponse(&aiserverv1.FetchRelevantKnowledgeForConversationResponse{}), nil
|
||||
}
|
||||
records, err := store.List(gitOrigin, req.Msg.GetLimit())
|
||||
if err != nil {
|
||||
return nil, connect.NewError(connect.CodeInternal, err)
|
||||
}
|
||||
items := make([]*aiserverv1.ConversationMessage_KnowledgeItem, 0, len(records))
|
||||
for _, record := range records {
|
||||
items = append(items, &aiserverv1.ConversationMessage_KnowledgeItem{
|
||||
Title: record.Title,
|
||||
Knowledge: docsIndexKnowledgeText(record),
|
||||
KnowledgeId: record.Identifier,
|
||||
IsGenerated: false,
|
||||
})
|
||||
}
|
||||
return connect.NewResponse(&aiserverv1.FetchRelevantKnowledgeForConversationResponse{KnowledgeItems: items}), nil
|
||||
}
|
||||
|
||||
func (service *Service) requireDocsIndexStore() (*DocsIndexStore, error) {
|
||||
if service == nil || service.docsIndexStore == nil {
|
||||
return nil, fmt.Errorf("docs index store is not initialized")
|
||||
}
|
||||
return service.docsIndexStore, nil
|
||||
}
|
||||
|
||||
func filterDocsIndexRecords(records []DocsIndexRecord, req *aiserverv1.AvailableDocsRequest) []DocsIndexRecord {
|
||||
if req == nil || req.GetGetAll() || (req.GetPartialUrl() == "" && req.GetPartialDocName() == "" && len(req.GetAdditionalDocIdentifiers()) == 0) {
|
||||
return records
|
||||
}
|
||||
partialURL := strings.ToLower(strings.TrimSpace(req.GetPartialUrl()))
|
||||
partialName := strings.ToLower(strings.TrimSpace(req.GetPartialDocName()))
|
||||
allowed := make(map[string]struct{})
|
||||
for _, identifier := range req.GetAdditionalDocIdentifiers() {
|
||||
trimmed := strings.TrimSpace(identifier)
|
||||
if trimmed != "" {
|
||||
allowed[trimmed] = struct{}{}
|
||||
}
|
||||
}
|
||||
filtered := make([]DocsIndexRecord, 0, len(records))
|
||||
for _, record := range records {
|
||||
_, explicitlyAllowed := allowed[record.Identifier]
|
||||
matchesURL := partialURL != "" && strings.Contains(strings.ToLower(record.URL), partialURL)
|
||||
matchesName := partialName != "" && strings.Contains(strings.ToLower(record.Title), partialName)
|
||||
if explicitlyAllowed || matchesURL || matchesName {
|
||||
filtered = append(filtered, record)
|
||||
}
|
||||
}
|
||||
return filtered
|
||||
}
|
||||
|
||||
func appendMissingAdditionalDocs(store *DocsIndexStore, records []DocsIndexRecord, req *aiserverv1.AvailableDocsRequest) ([]DocsIndexRecord, error) {
|
||||
if req == nil {
|
||||
return records, nil
|
||||
}
|
||||
seen := make(map[string]struct{}, len(records))
|
||||
for _, record := range records {
|
||||
seen[record.Identifier] = struct{}{}
|
||||
}
|
||||
for _, identifier := range req.GetAdditionalDocIdentifiers() {
|
||||
trimmed := strings.TrimSpace(identifier)
|
||||
if trimmed == "" {
|
||||
continue
|
||||
}
|
||||
if _, ok := seen[trimmed]; ok {
|
||||
continue
|
||||
}
|
||||
record := DocsIndexRecord{
|
||||
ID: trimmed,
|
||||
Identifier: trimmed,
|
||||
Title: trimmed,
|
||||
Status: docsIndexStatusIndexed,
|
||||
Source: docsIndexSourceAdditional,
|
||||
}
|
||||
if store != nil {
|
||||
persisted, err := store.Upsert(record)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
record = persisted
|
||||
}
|
||||
records = append(records, record)
|
||||
seen[trimmed] = struct{}{}
|
||||
}
|
||||
return records, nil
|
||||
}
|
||||
|
||||
func docsIndexDocumentationInfo(record DocsIndexRecord) *aiserverv1.DocumentationInfo {
|
||||
return &aiserverv1.DocumentationInfo{
|
||||
DocIdentifier: record.Identifier,
|
||||
Metadata: &aiserverv1.DocumentationMetadata{
|
||||
PrefixUrl: record.URL,
|
||||
DocName: record.Title,
|
||||
TruePrefixUrl: record.URL,
|
||||
Public: false,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func docsIndexChunks(record DocsIndexRecord, topK uint32) []*aiserverv1.DocumentationChunk {
|
||||
content := strings.TrimSpace(record.Content)
|
||||
if content == "" {
|
||||
return nil
|
||||
}
|
||||
if topK == 0 {
|
||||
topK = 1
|
||||
}
|
||||
return []*aiserverv1.DocumentationChunk{
|
||||
{
|
||||
DocName: record.Title,
|
||||
PageUrl: record.URL,
|
||||
DocumentationChunk: content,
|
||||
Score: 1,
|
||||
PageTitle: record.Title,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func docsIndexKnowledgeText(record DocsIndexRecord) string {
|
||||
parts := make([]string, 0, 3)
|
||||
if strings.TrimSpace(record.URL) != "" {
|
||||
parts = append(parts, "URL: "+strings.TrimSpace(record.URL))
|
||||
}
|
||||
if strings.TrimSpace(record.Title) != "" {
|
||||
parts = append(parts, "Title: "+strings.TrimSpace(record.Title))
|
||||
}
|
||||
if strings.TrimSpace(record.Content) != "" {
|
||||
parts = append(parts, strings.TrimSpace(record.Content))
|
||||
}
|
||||
return strings.Join(parts, "\n")
|
||||
}
|
||||
|
||||
func docsIndexURLCandidate(value string) string {
|
||||
trimmed := strings.TrimSpace(value)
|
||||
if strings.HasPrefix(trimmed, "http://") || strings.HasPrefix(trimmed, "https://") {
|
||||
fields := strings.Fields(trimmed)
|
||||
if len(fields) > 0 {
|
||||
return fields[0]
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
@@ -0,0 +1,241 @@
|
||||
package forwarder
|
||||
|
||||
import (
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sort"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
const (
|
||||
docsIndexStatusIndexed = "indexed"
|
||||
docsIndexSourceLocal = "local_docs"
|
||||
docsIndexSourceAdditional = "additional_docs"
|
||||
)
|
||||
|
||||
type DocsIndexStore struct {
|
||||
mu sync.Mutex
|
||||
root string
|
||||
path string
|
||||
loaded bool
|
||||
state docsIndexState
|
||||
}
|
||||
|
||||
type docsIndexState struct {
|
||||
SchemaVersion int `json:"schema_version"`
|
||||
Docs map[string]DocsIndexRecord `json:"docs"`
|
||||
}
|
||||
|
||||
type DocsIndexRecord struct {
|
||||
ID string `json:"id"`
|
||||
Identifier string `json:"identifier"`
|
||||
Title string `json:"title"`
|
||||
URL string `json:"url,omitempty"`
|
||||
Content string `json:"content,omitempty"`
|
||||
GitOrigin string `json:"git_origin,omitempty"`
|
||||
Status string `json:"status"`
|
||||
Source string `json:"source"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
LastCrawlAt time.Time `json:"last_crawl_at,omitempty"`
|
||||
LastIndexAt time.Time `json:"last_index_at,omitempty"`
|
||||
}
|
||||
|
||||
func NewDocsIndexStore(root string) *DocsIndexStore {
|
||||
trimmed := strings.TrimSpace(root)
|
||||
if trimmed == "" {
|
||||
trimmed = "docs-index"
|
||||
}
|
||||
return &DocsIndexStore{
|
||||
root: trimmed,
|
||||
path: filepath.Join(trimmed, "index.json"),
|
||||
state: docsIndexState{
|
||||
SchemaVersion: 1,
|
||||
Docs: make(map[string]DocsIndexRecord),
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func (store *DocsIndexStore) Upsert(record DocsIndexRecord) (DocsIndexRecord, error) {
|
||||
if store == nil {
|
||||
return DocsIndexRecord{}, fmt.Errorf("docs index store is nil")
|
||||
}
|
||||
store.mu.Lock()
|
||||
defer store.mu.Unlock()
|
||||
if err := store.loadLocked(); err != nil {
|
||||
return DocsIndexRecord{}, err
|
||||
}
|
||||
now := time.Now().UTC()
|
||||
identifier := strings.TrimSpace(record.Identifier)
|
||||
if identifier == "" {
|
||||
identifier = stableDocsIdentifier(record.Title, record.URL, record.Content, record.GitOrigin)
|
||||
}
|
||||
existing, ok := store.state.Docs[identifier]
|
||||
if ok {
|
||||
existing.Title = firstNonEmptyDocs(record.Title, existing.Title, identifier)
|
||||
existing.URL = firstNonEmptyDocs(record.URL, existing.URL)
|
||||
existing.Content = firstNonEmptyDocs(record.Content, existing.Content)
|
||||
existing.GitOrigin = firstNonEmptyDocs(record.GitOrigin, existing.GitOrigin)
|
||||
existing.Status = docsIndexStatusIndexed
|
||||
existing.Source = firstNonEmptyDocs(record.Source, existing.Source, docsIndexSourceLocal)
|
||||
existing.UpdatedAt = now
|
||||
existing.LastCrawlAt = now
|
||||
existing.LastIndexAt = now
|
||||
store.state.Docs[identifier] = existing
|
||||
return existing, store.saveLocked()
|
||||
}
|
||||
record.ID = firstNonEmptyDocs(record.ID, identifier)
|
||||
record.Identifier = identifier
|
||||
record.Title = firstNonEmptyDocs(record.Title, identifier)
|
||||
record.Status = docsIndexStatusIndexed
|
||||
record.Source = firstNonEmptyDocs(record.Source, docsIndexSourceLocal)
|
||||
record.CreatedAt = now
|
||||
record.UpdatedAt = now
|
||||
record.LastCrawlAt = now
|
||||
record.LastIndexAt = now
|
||||
store.state.Docs[identifier] = record
|
||||
return record, store.saveLocked()
|
||||
}
|
||||
|
||||
func (store *DocsIndexStore) List(gitOrigin string, limit int32) ([]DocsIndexRecord, error) {
|
||||
if store == nil {
|
||||
return nil, fmt.Errorf("docs index store is nil")
|
||||
}
|
||||
store.mu.Lock()
|
||||
defer store.mu.Unlock()
|
||||
if err := store.loadLocked(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
origin := strings.TrimSpace(gitOrigin)
|
||||
records := make([]DocsIndexRecord, 0, len(store.state.Docs))
|
||||
for _, record := range store.state.Docs {
|
||||
if origin != "" && strings.TrimSpace(record.GitOrigin) != origin {
|
||||
continue
|
||||
}
|
||||
records = append(records, record)
|
||||
}
|
||||
sort.SliceStable(records, func(i, j int) bool {
|
||||
if records[i].UpdatedAt.Equal(records[j].UpdatedAt) {
|
||||
return records[i].Identifier < records[j].Identifier
|
||||
}
|
||||
return records[i].UpdatedAt.After(records[j].UpdatedAt)
|
||||
})
|
||||
if limit > 0 && len(records) > int(limit) {
|
||||
records = records[:limit]
|
||||
}
|
||||
return records, nil
|
||||
}
|
||||
|
||||
func (store *DocsIndexStore) Get(identifier string) (DocsIndexRecord, bool, error) {
|
||||
if store == nil {
|
||||
return DocsIndexRecord{}, false, fmt.Errorf("docs index store is nil")
|
||||
}
|
||||
store.mu.Lock()
|
||||
defer store.mu.Unlock()
|
||||
if err := store.loadLocked(); err != nil {
|
||||
return DocsIndexRecord{}, false, err
|
||||
}
|
||||
record, ok := store.state.Docs[strings.TrimSpace(identifier)]
|
||||
return record, ok, nil
|
||||
}
|
||||
|
||||
func (store *DocsIndexStore) Remove(identifier string) error {
|
||||
if store == nil {
|
||||
return fmt.Errorf("docs index store is nil")
|
||||
}
|
||||
store.mu.Lock()
|
||||
defer store.mu.Unlock()
|
||||
if err := store.loadLocked(); err != nil {
|
||||
return err
|
||||
}
|
||||
delete(store.state.Docs, strings.TrimSpace(identifier))
|
||||
return store.saveLocked()
|
||||
}
|
||||
|
||||
func (store *DocsIndexStore) loadLocked() error {
|
||||
if store.loaded {
|
||||
return nil
|
||||
}
|
||||
state := docsIndexState{SchemaVersion: 1, Docs: make(map[string]DocsIndexRecord)}
|
||||
data, err := os.ReadFile(store.path)
|
||||
if err != nil {
|
||||
if os.IsNotExist(err) {
|
||||
store.state = state
|
||||
store.loaded = true
|
||||
return nil
|
||||
}
|
||||
return err
|
||||
}
|
||||
if len(data) != 0 {
|
||||
if err := json.Unmarshal(data, &state); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
if state.SchemaVersion <= 0 {
|
||||
state.SchemaVersion = 1
|
||||
}
|
||||
if state.Docs == nil {
|
||||
state.Docs = make(map[string]DocsIndexRecord)
|
||||
}
|
||||
store.state = state
|
||||
store.loaded = true
|
||||
return nil
|
||||
}
|
||||
|
||||
func (store *DocsIndexStore) saveLocked() error {
|
||||
if err := os.MkdirAll(store.root, 0o755); err != nil {
|
||||
return err
|
||||
}
|
||||
data, err := json.MarshalIndent(store.state, "", " ")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
tmp, err := os.CreateTemp(store.root, ".index-*.json.tmp")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
tmpPath := tmp.Name()
|
||||
defer func() { _ = os.Remove(tmpPath) }()
|
||||
if _, err := tmp.Write(data); err != nil {
|
||||
_ = tmp.Close()
|
||||
return err
|
||||
}
|
||||
if err := tmp.Close(); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := os.Chmod(tmpPath, 0o600); err != nil {
|
||||
return err
|
||||
}
|
||||
return os.Rename(tmpPath, store.path)
|
||||
}
|
||||
|
||||
func stableDocsIdentifier(values ...string) string {
|
||||
parts := make([]string, 0, len(values))
|
||||
for _, value := range values {
|
||||
trimmed := strings.TrimSpace(value)
|
||||
if trimmed != "" {
|
||||
parts = append(parts, trimmed)
|
||||
}
|
||||
}
|
||||
if len(parts) == 0 {
|
||||
return "doc_local"
|
||||
}
|
||||
sum := sha256.Sum256([]byte(strings.Join(parts, "\x00")))
|
||||
return "doc_" + hex.EncodeToString(sum[:])[:24]
|
||||
}
|
||||
|
||||
func firstNonEmptyDocs(values ...string) string {
|
||||
for _, value := range values {
|
||||
trimmed := strings.TrimSpace(value)
|
||||
if trimmed != "" {
|
||||
return trimmed
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
@@ -0,0 +1,491 @@
|
||||
package forwarder
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"strings"
|
||||
"unicode/utf8"
|
||||
|
||||
"github.com/sergi/go-diff/diffmatchpatch"
|
||||
"google.golang.org/protobuf/encoding/protojson"
|
||||
|
||||
"cursor/gen/agentv1"
|
||||
runtimecore "cursor/internal/backend/agent/core"
|
||||
)
|
||||
|
||||
const editReadBinaryContentLimit = 32 * 1024
|
||||
|
||||
type editComputation struct {
|
||||
BeforeContent string
|
||||
AfterContent string
|
||||
DiffString string
|
||||
LinesAdded int32
|
||||
LinesRemoved int32
|
||||
Message string
|
||||
}
|
||||
|
||||
type editResultPayload struct {
|
||||
BeforeContent string
|
||||
AfterContent string
|
||||
DiffString string
|
||||
LinesAdded int32
|
||||
LinesRemoved int32
|
||||
Message string
|
||||
}
|
||||
|
||||
type editMatch struct {
|
||||
Start int
|
||||
End int
|
||||
}
|
||||
|
||||
type normalizedEditContent struct {
|
||||
Text string
|
||||
OriginalOffsets []int
|
||||
}
|
||||
|
||||
func computeEditDiff(beforeContent string, afterContent string) (string, int32, int32) {
|
||||
dmp := diffmatchpatch.New()
|
||||
charsBefore, charsAfter, lineArray := dmp.DiffLinesToChars(beforeContent, afterContent)
|
||||
diffs := dmp.DiffMain(charsBefore, charsAfter, false)
|
||||
diffs = dmp.DiffCharsToLines(diffs, lineArray)
|
||||
|
||||
linesAdded := int32(0)
|
||||
linesRemoved := int32(0)
|
||||
for _, diff := range diffs {
|
||||
switch diff.Type {
|
||||
case diffmatchpatch.DiffInsert:
|
||||
linesAdded += countChangedLines(diff.Text)
|
||||
case diffmatchpatch.DiffDelete:
|
||||
linesRemoved += countChangedLines(diff.Text)
|
||||
}
|
||||
}
|
||||
return dmp.PatchToText(dmp.PatchMake(beforeContent, diffs)), linesAdded, linesRemoved
|
||||
}
|
||||
|
||||
func normalizeEditContentLineEndings(content string) normalizedEditContent {
|
||||
var builder strings.Builder
|
||||
builder.Grow(len(content))
|
||||
offsets := make([]int, 1, len(content)+1)
|
||||
offsets[0] = 0
|
||||
for index := 0; index < len(content); {
|
||||
if content[index] == '\r' && index+1 < len(content) && content[index+1] == '\n' {
|
||||
builder.WriteByte('\n')
|
||||
index += 2
|
||||
offsets = append(offsets, index)
|
||||
continue
|
||||
}
|
||||
if content[index] == '\r' {
|
||||
builder.WriteByte('\n')
|
||||
index++
|
||||
offsets = append(offsets, index)
|
||||
continue
|
||||
}
|
||||
builder.WriteByte(content[index])
|
||||
index++
|
||||
offsets = append(offsets, index)
|
||||
}
|
||||
return normalizedEditContent{
|
||||
Text: builder.String(),
|
||||
OriginalOffsets: offsets,
|
||||
}
|
||||
}
|
||||
|
||||
func normalizeLineEndingsToLF(value string) string {
|
||||
if !strings.ContainsAny(value, "\r\n") {
|
||||
return value
|
||||
}
|
||||
normalized := strings.ReplaceAll(value, "\r\n", "\n")
|
||||
return strings.ReplaceAll(normalized, "\r", "\n")
|
||||
}
|
||||
|
||||
func detectPreferredFileLineEnding(content string) string {
|
||||
return dominantLineEnding(content, "\n")
|
||||
}
|
||||
|
||||
func preferredMatchLineEnding(content string, match editMatch, defaultLineEnding string) string {
|
||||
return dominantLineEnding(content[match.Start:match.End], defaultLineEnding)
|
||||
}
|
||||
|
||||
func normalizeReplacementLineEndings(newString string, lineEnding string) string {
|
||||
if !strings.ContainsAny(newString, "\r\n") {
|
||||
return newString
|
||||
}
|
||||
normalized := normalizeLineEndingsToLF(newString)
|
||||
switch lineEnding {
|
||||
case "\r\n":
|
||||
return strings.ReplaceAll(normalized, "\n", "\r\n")
|
||||
case "\r":
|
||||
return strings.ReplaceAll(normalized, "\n", "\r")
|
||||
default:
|
||||
return normalized
|
||||
}
|
||||
}
|
||||
|
||||
func dominantLineEnding(content string, fallback string) string {
|
||||
if fallback == "" {
|
||||
fallback = "\n"
|
||||
}
|
||||
counts := map[string]int{
|
||||
"\n": 0,
|
||||
"\r\n": 0,
|
||||
"\r": 0,
|
||||
}
|
||||
first := ""
|
||||
for index := 0; index < len(content); {
|
||||
switch {
|
||||
case content[index] == '\r' && index+1 < len(content) && content[index+1] == '\n':
|
||||
counts["\r\n"]++
|
||||
if first == "" {
|
||||
first = "\r\n"
|
||||
}
|
||||
index += 2
|
||||
case content[index] == '\r':
|
||||
counts["\r"]++
|
||||
if first == "" {
|
||||
first = "\r"
|
||||
}
|
||||
index++
|
||||
case content[index] == '\n':
|
||||
counts["\n"]++
|
||||
if first == "" {
|
||||
first = "\n"
|
||||
}
|
||||
index++
|
||||
default:
|
||||
index++
|
||||
}
|
||||
}
|
||||
best := ""
|
||||
bestCount := 0
|
||||
for _, candidate := range []string{"\r\n", "\n", "\r"} {
|
||||
if counts[candidate] > bestCount {
|
||||
best = candidate
|
||||
bestCount = counts[candidate]
|
||||
}
|
||||
}
|
||||
if bestCount == 0 {
|
||||
return fallback
|
||||
}
|
||||
for _, candidate := range []string{"\r\n", "\n", "\r"} {
|
||||
if counts[candidate] == bestCount && candidate == first {
|
||||
return candidate
|
||||
}
|
||||
}
|
||||
return best
|
||||
}
|
||||
|
||||
func countChangedLines(text string) int32 {
|
||||
if text == "" {
|
||||
return 0
|
||||
}
|
||||
count := int32(strings.Count(text, "\n"))
|
||||
if !strings.HasSuffix(text, "\n") {
|
||||
count++
|
||||
}
|
||||
return count
|
||||
}
|
||||
|
||||
func buildCompletedEditToolCall(path string, result *agentv1.EditResult) *agentv1.ToolCall {
|
||||
args := &agentv1.EditArgs{Path: strings.TrimSpace(path)}
|
||||
return &agentv1.ToolCall{
|
||||
Tool: &agentv1.ToolCall_EditToolCall{
|
||||
EditToolCall: &agentv1.EditToolCall{
|
||||
Args: args,
|
||||
Result: result,
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func buildSuccessfulEditResult(path string, beforeContent string, afterContent string, diffString string, linesAdded int32, linesRemoved int32, message string) *agentv1.EditResult {
|
||||
success := &agentv1.EditSuccess{
|
||||
Path: strings.TrimSpace(path),
|
||||
AfterFullFileContent: afterContent,
|
||||
}
|
||||
if beforeContent != "" {
|
||||
success.BeforeFullFileContent = stringPtr(beforeContent)
|
||||
}
|
||||
if diffString != "" {
|
||||
success.DiffString = stringPtr(diffString)
|
||||
}
|
||||
if linesAdded != 0 {
|
||||
success.LinesAdded = int32Ptr(linesAdded)
|
||||
}
|
||||
if linesRemoved != 0 {
|
||||
success.LinesRemoved = int32Ptr(linesRemoved)
|
||||
}
|
||||
if message != "" {
|
||||
success.Message = stringPtr(message)
|
||||
}
|
||||
return &agentv1.EditResult{
|
||||
Result: &agentv1.EditResult_Success{Success: success},
|
||||
}
|
||||
}
|
||||
|
||||
func buildFinalEditSuccessResult(path string, afterContent string, payload editResultPayload) *agentv1.EditResult {
|
||||
diffString := payload.DiffString
|
||||
linesAdded := payload.LinesAdded
|
||||
linesRemoved := payload.LinesRemoved
|
||||
if afterContent != payload.AfterContent {
|
||||
diffString, linesAdded, linesRemoved = computeEditDiff(payload.BeforeContent, afterContent)
|
||||
}
|
||||
return buildSuccessfulEditResult(path, payload.BeforeContent, afterContent, diffString, linesAdded, linesRemoved, payload.Message)
|
||||
}
|
||||
|
||||
func buildEditErrorResult(path string, errorMessage string) *agentv1.EditResult {
|
||||
trimmedPath := strings.TrimSpace(path)
|
||||
trimmedError := strings.TrimSpace(errorMessage)
|
||||
return &agentv1.EditResult{
|
||||
Result: &agentv1.EditResult_Error{
|
||||
Error: &agentv1.EditError{
|
||||
Path: trimmedPath,
|
||||
Error: trimmedError,
|
||||
ModelVisibleError: stringPtr(trimmedError),
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func decodeReadDataAsEditableText(data []byte) (string, bool, string) {
|
||||
if len(data) > editReadBinaryContentLimit {
|
||||
return "", false, fmt.Sprintf("Read returned binary data larger than %d bytes", editReadBinaryContentLimit)
|
||||
}
|
||||
if !utf8.Valid(data) {
|
||||
return "", false, "Read returned binary data that is not valid UTF-8 text"
|
||||
}
|
||||
return string(data), true, ""
|
||||
}
|
||||
|
||||
func buildEditResultFromReadResult(path string, result *agentv1.ReadResult) *agentv1.EditResult {
|
||||
if result == nil {
|
||||
return buildEditErrorResult(path, "read result missing")
|
||||
}
|
||||
switch item := result.GetResult().(type) {
|
||||
case *agentv1.ReadResult_FileNotFound:
|
||||
return &agentv1.EditResult{
|
||||
Result: &agentv1.EditResult_FileNotFound{
|
||||
FileNotFound: &agentv1.EditFileNotFound{Path: firstNonEmpty(item.FileNotFound.GetPath(), path)},
|
||||
},
|
||||
}
|
||||
case *agentv1.ReadResult_PermissionDenied:
|
||||
return &agentv1.EditResult{
|
||||
Result: &agentv1.EditResult_ReadPermissionDenied{
|
||||
ReadPermissionDenied: &agentv1.EditReadPermissionDenied{Path: firstNonEmpty(item.PermissionDenied.GetPath(), path)},
|
||||
},
|
||||
}
|
||||
case *agentv1.ReadResult_Rejected:
|
||||
return &agentv1.EditResult{
|
||||
Result: &agentv1.EditResult_Rejected{
|
||||
Rejected: &agentv1.EditRejected{
|
||||
Path: firstNonEmpty(item.Rejected.GetPath(), path),
|
||||
Reason: item.Rejected.GetReason(),
|
||||
},
|
||||
},
|
||||
}
|
||||
case *agentv1.ReadResult_InvalidFile:
|
||||
return buildEditErrorResult(firstNonEmpty(item.InvalidFile.GetPath(), path), item.InvalidFile.GetReason())
|
||||
case *agentv1.ReadResult_Error:
|
||||
return buildEditErrorResult(firstNonEmpty(item.Error.GetPath(), path), item.Error.GetError())
|
||||
case *agentv1.ReadResult_Success:
|
||||
return buildEditErrorResult(firstNonEmpty(item.Success.GetPath(), path), "read result did not include editable text content")
|
||||
default:
|
||||
return buildEditErrorResult(path, "unknown read result")
|
||||
}
|
||||
}
|
||||
|
||||
func buildEditResultFromWriteResult(path string, result *agentv1.WriteResult) *agentv1.EditResult {
|
||||
if result == nil {
|
||||
return buildEditErrorResult(path, "write result missing")
|
||||
}
|
||||
switch item := result.GetResult().(type) {
|
||||
case *agentv1.WriteResult_PermissionDenied:
|
||||
return &agentv1.EditResult{
|
||||
Result: &agentv1.EditResult_WritePermissionDenied{
|
||||
WritePermissionDenied: &agentv1.EditWritePermissionDenied{
|
||||
Path: firstNonEmpty(item.PermissionDenied.GetPath(), path),
|
||||
Error: item.PermissionDenied.GetError(),
|
||||
},
|
||||
},
|
||||
}
|
||||
case *agentv1.WriteResult_Rejected:
|
||||
return &agentv1.EditResult{
|
||||
Result: &agentv1.EditResult_Rejected{
|
||||
Rejected: &agentv1.EditRejected{
|
||||
Path: firstNonEmpty(item.Rejected.GetPath(), path),
|
||||
Reason: item.Rejected.GetReason(),
|
||||
},
|
||||
},
|
||||
}
|
||||
case *agentv1.WriteResult_NoSpace:
|
||||
return buildEditErrorResult(firstNonEmpty(item.NoSpace.GetPath(), path), "no space left")
|
||||
case *agentv1.WriteResult_Error:
|
||||
return buildEditErrorResult(firstNonEmpty(item.Error.GetPath(), path), item.Error.GetError())
|
||||
case *agentv1.WriteResult_Success:
|
||||
afterContent := item.Success.GetFileContentAfterWrite()
|
||||
return buildSuccessfulEditResult(firstNonEmpty(item.Success.GetPath(), path), "", afterContent, "", 0, 0, "")
|
||||
default:
|
||||
return buildEditErrorResult(path, "unknown write result")
|
||||
}
|
||||
}
|
||||
|
||||
func extractReadContentForEdit(result *agentv1.ReadResult) (string, bool) {
|
||||
success := result.GetSuccess()
|
||||
if success == nil {
|
||||
return "", false
|
||||
}
|
||||
switch output := success.GetOutput().(type) {
|
||||
case *agentv1.ReadSuccess_Content:
|
||||
return normalizeLineEndingsToLF(output.Content), true
|
||||
case *agentv1.ReadSuccess_Data:
|
||||
text, ok, _ := decodeReadDataAsEditableText(output.Data)
|
||||
return normalizeLineEndingsToLF(text), ok
|
||||
}
|
||||
if success.GetOutputBlobId() != nil {
|
||||
return "", false
|
||||
}
|
||||
return "", true
|
||||
}
|
||||
|
||||
func summarizeEditResult(result *agentv1.EditResult) string {
|
||||
if result == nil {
|
||||
return "edit result missing"
|
||||
}
|
||||
switch item := result.GetResult().(type) {
|
||||
case *agentv1.EditResult_Success:
|
||||
return item.Success.GetAfterFullFileContent()
|
||||
case *agentv1.EditResult_FileNotFound:
|
||||
return fmt.Sprintf("file not found: %s", item.FileNotFound.GetPath())
|
||||
case *agentv1.EditResult_ReadPermissionDenied:
|
||||
return fmt.Sprintf("permission denied: %s", item.ReadPermissionDenied.GetPath())
|
||||
case *agentv1.EditResult_WritePermissionDenied:
|
||||
return item.WritePermissionDenied.GetError()
|
||||
case *agentv1.EditResult_Rejected:
|
||||
return item.Rejected.GetReason()
|
||||
case *agentv1.EditResult_Error:
|
||||
return item.Error.GetError()
|
||||
default:
|
||||
return "unknown edit result"
|
||||
}
|
||||
}
|
||||
|
||||
func summarizeEditResultWithoutPath(result *agentv1.EditResult) string {
|
||||
if result == nil {
|
||||
return "edit result missing"
|
||||
}
|
||||
switch item := result.GetResult().(type) {
|
||||
case *agentv1.EditResult_Success:
|
||||
return item.Success.GetAfterFullFileContent()
|
||||
case *agentv1.EditResult_FileNotFound:
|
||||
return "file not found"
|
||||
case *agentv1.EditResult_ReadPermissionDenied:
|
||||
return "permission denied"
|
||||
case *agentv1.EditResult_WritePermissionDenied:
|
||||
return item.WritePermissionDenied.GetError()
|
||||
case *agentv1.EditResult_Rejected:
|
||||
return item.Rejected.GetReason()
|
||||
case *agentv1.EditResult_Error:
|
||||
return item.Error.GetError()
|
||||
default:
|
||||
return "unknown edit result"
|
||||
}
|
||||
}
|
||||
|
||||
func editResultWithoutPath(result *agentv1.EditResult) *agentv1.EditResult {
|
||||
if result == nil {
|
||||
return nil
|
||||
}
|
||||
switch item := result.GetResult().(type) {
|
||||
case *agentv1.EditResult_Success:
|
||||
return &agentv1.EditResult{Result: &agentv1.EditResult_Success{Success: &agentv1.EditSuccess{
|
||||
AfterFullFileContent: item.Success.GetAfterFullFileContent(),
|
||||
BeforeFullFileContent: item.Success.BeforeFullFileContent,
|
||||
DiffString: item.Success.DiffString,
|
||||
LinesAdded: item.Success.LinesAdded,
|
||||
LinesRemoved: item.Success.LinesRemoved,
|
||||
Message: item.Success.Message,
|
||||
}}}
|
||||
case *agentv1.EditResult_FileNotFound:
|
||||
return &agentv1.EditResult{Result: &agentv1.EditResult_FileNotFound{FileNotFound: &agentv1.EditFileNotFound{}}}
|
||||
case *agentv1.EditResult_ReadPermissionDenied:
|
||||
return &agentv1.EditResult{Result: &agentv1.EditResult_ReadPermissionDenied{ReadPermissionDenied: &agentv1.EditReadPermissionDenied{}}}
|
||||
case *agentv1.EditResult_WritePermissionDenied:
|
||||
return &agentv1.EditResult{Result: &agentv1.EditResult_WritePermissionDenied{WritePermissionDenied: &agentv1.EditWritePermissionDenied{Error: item.WritePermissionDenied.GetError()}}}
|
||||
case *agentv1.EditResult_Rejected:
|
||||
return &agentv1.EditResult{Result: &agentv1.EditResult_Rejected{Rejected: &agentv1.EditRejected{Reason: item.Rejected.GetReason()}}}
|
||||
case *agentv1.EditResult_Error:
|
||||
return &agentv1.EditResult{Result: &agentv1.EditResult_Error{Error: &agentv1.EditError{
|
||||
Error: item.Error.GetError(),
|
||||
ModelVisibleError: item.Error.ModelVisibleError,
|
||||
}}}
|
||||
default:
|
||||
return result
|
||||
}
|
||||
}
|
||||
|
||||
func summarizeStructuredEditResult(result *agentv1.EditResult) string {
|
||||
if result == nil {
|
||||
return "edit result missing"
|
||||
}
|
||||
if _, ok := result.GetResult().(*agentv1.EditResult_Success); ok {
|
||||
encoded, err := protojson.Marshal(result)
|
||||
if err == nil {
|
||||
return string(encoded)
|
||||
}
|
||||
}
|
||||
return summarizeEditResult(result)
|
||||
}
|
||||
|
||||
func hiddenWriteControlError(message *agentv1.ExecClientControlMessage) string {
|
||||
if message == nil {
|
||||
return "write operation failed"
|
||||
}
|
||||
switch item := message.GetMessage().(type) {
|
||||
case *agentv1.ExecClientControlMessage_Throw:
|
||||
return firstNonEmpty(strings.TrimSpace(item.Throw.GetError()), "write operation failed")
|
||||
case *agentv1.ExecClientControlMessage_StreamClose:
|
||||
return "write operation closed unexpectedly"
|
||||
default:
|
||||
return "write operation failed"
|
||||
}
|
||||
}
|
||||
|
||||
func readJSONStringAny(args map[string]any, keys ...string) (string, bool, bool) {
|
||||
for _, key := range keys {
|
||||
value, ok := args[key]
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
typed, ok := value.(string)
|
||||
if !ok {
|
||||
return "", true, false
|
||||
}
|
||||
return typed, true, true
|
||||
}
|
||||
return "", false, false
|
||||
}
|
||||
|
||||
func readJSONIntAny(args map[string]any, keys ...string) (int, bool, bool) {
|
||||
value, found, err := runtimecore.ReadIntArg(args, keys...)
|
||||
if err != nil {
|
||||
return 0, true, false
|
||||
}
|
||||
return value, found, found
|
||||
}
|
||||
|
||||
func readJSONBoolAny(args map[string]any, keys ...string) (bool, bool, bool) {
|
||||
for _, key := range keys {
|
||||
value, ok := args[key]
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
typed, ok := value.(bool)
|
||||
if !ok {
|
||||
return false, true, false
|
||||
}
|
||||
return typed, true, true
|
||||
}
|
||||
return false, false, false
|
||||
}
|
||||
|
||||
func int32Ptr(value int32) *int32 {
|
||||
return &value
|
||||
}
|
||||
@@ -0,0 +1,612 @@
|
||||
// events.go 负责构造对外兼容的 legacy RunSSE 消息。
|
||||
package forwarder
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"strings"
|
||||
|
||||
"google.golang.org/protobuf/proto"
|
||||
"google.golang.org/protobuf/types/known/structpb"
|
||||
|
||||
"cursor/gen/agentv1"
|
||||
execbridge "cursor/internal/backend/agent/bridge/exec"
|
||||
runtimecore "cursor/internal/backend/agent/core"
|
||||
)
|
||||
|
||||
// buildHeartbeatMessage 构造一个服务端心跳消息。
|
||||
func buildHeartbeatMessage() *agentv1.AgentServerMessage {
|
||||
return &agentv1.AgentServerMessage{
|
||||
Message: &agentv1.AgentServerMessage_InteractionUpdate{
|
||||
InteractionUpdate: &agentv1.InteractionUpdate{
|
||||
Message: &agentv1.InteractionUpdate_Heartbeat{
|
||||
Heartbeat: &agentv1.HeartbeatUpdate{},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// buildTextDeltaMessage 构造文本增量消息。
|
||||
func buildTextDeltaMessage(text string) *agentv1.AgentServerMessage {
|
||||
return &agentv1.AgentServerMessage{
|
||||
Message: &agentv1.AgentServerMessage_InteractionUpdate{
|
||||
InteractionUpdate: &agentv1.InteractionUpdate{
|
||||
Message: &agentv1.InteractionUpdate_TextDelta{
|
||||
TextDelta: &agentv1.TextDeltaUpdate{Text: text},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// buildThinkingDeltaMessage 构造思考文本增量消息。
|
||||
func buildThinkingDeltaMessage(text string, style agentv1.ThinkingStyle) *agentv1.AgentServerMessage {
|
||||
styleCopy := style
|
||||
return &agentv1.AgentServerMessage{
|
||||
Message: &agentv1.AgentServerMessage_InteractionUpdate{
|
||||
InteractionUpdate: &agentv1.InteractionUpdate{
|
||||
Message: &agentv1.InteractionUpdate_ThinkingDelta{
|
||||
ThinkingDelta: &agentv1.ThinkingDeltaUpdate{
|
||||
Text: text,
|
||||
ThinkingStyle: &styleCopy,
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// buildThinkingCompletedMessage 构造思考阶段结束消息。
|
||||
func buildThinkingCompletedMessage(durationMS int32) *agentv1.AgentServerMessage {
|
||||
return &agentv1.AgentServerMessage{
|
||||
Message: &agentv1.AgentServerMessage_InteractionUpdate{
|
||||
InteractionUpdate: &agentv1.InteractionUpdate{
|
||||
Message: &agentv1.InteractionUpdate_ThinkingCompleted{
|
||||
ThinkingCompleted: &agentv1.ThinkingCompletedUpdate{
|
||||
ThinkingDurationMs: durationMS,
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// buildSummaryStartedMessage 构造摘要开始消息。
|
||||
func buildSummaryStartedMessage() *agentv1.AgentServerMessage {
|
||||
return &agentv1.AgentServerMessage{
|
||||
Message: &agentv1.AgentServerMessage_InteractionUpdate{
|
||||
InteractionUpdate: &agentv1.InteractionUpdate{
|
||||
Message: &agentv1.InteractionUpdate_SummaryStarted{
|
||||
SummaryStarted: &agentv1.SummaryStartedUpdate{},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// buildSummaryMessage 构造摘要增量/完成内容消息。
|
||||
func buildSummaryMessage(summary string) *agentv1.AgentServerMessage {
|
||||
return &agentv1.AgentServerMessage{
|
||||
Message: &agentv1.AgentServerMessage_InteractionUpdate{
|
||||
InteractionUpdate: &agentv1.InteractionUpdate{
|
||||
Message: &agentv1.InteractionUpdate_Summary{
|
||||
Summary: &agentv1.SummaryUpdate{Summary: strings.TrimSpace(summary)},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// buildSummaryCompletedMessage 构造摘要完成消息。
|
||||
// field 11 在 legacy aiserver 协议里会被客户端按 GetThoughtAnnotationRequest 解码,
|
||||
// 所以这里需要把 request_id 写入 field 1,驱动客户端继续查询 "Chat context summarized" 注解。
|
||||
func buildSummaryCompletedMessage(requestID string) *agentv1.AgentServerMessage {
|
||||
var message *string
|
||||
if strings.TrimSpace(requestID) != "" {
|
||||
value := strings.TrimSpace(requestID)
|
||||
message = &value
|
||||
}
|
||||
return &agentv1.AgentServerMessage{
|
||||
Message: &agentv1.AgentServerMessage_InteractionUpdate{
|
||||
InteractionUpdate: &agentv1.InteractionUpdate{
|
||||
Message: &agentv1.InteractionUpdate_SummaryCompleted{
|
||||
SummaryCompleted: &agentv1.SummaryCompletedUpdate{
|
||||
HookMessage: message,
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// buildToolCallStartedMessage 构造工具调用开始消息。
|
||||
func buildToolCallStartedMessage(callID string, modelCallID string, toolCall *agentv1.ToolCall) *agentv1.AgentServerMessage {
|
||||
return &agentv1.AgentServerMessage{
|
||||
Message: &agentv1.AgentServerMessage_InteractionUpdate{
|
||||
InteractionUpdate: &agentv1.InteractionUpdate{
|
||||
Message: &agentv1.InteractionUpdate_ToolCallStarted{
|
||||
ToolCallStarted: &agentv1.ToolCallStartedUpdate{
|
||||
CallId: callID,
|
||||
ToolCall: toolCall,
|
||||
ModelCallId: modelCallID,
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// buildPartialToolCallMessage 构造工具调用参数流式生成中的兼容消息。
|
||||
func buildPartialToolCallMessage(callID string, modelCallID string, toolCall *agentv1.ToolCall, argsTextDelta string) *agentv1.AgentServerMessage {
|
||||
return &agentv1.AgentServerMessage{
|
||||
Message: &agentv1.AgentServerMessage_InteractionUpdate{
|
||||
InteractionUpdate: &agentv1.InteractionUpdate{
|
||||
Message: &agentv1.InteractionUpdate_PartialToolCall{
|
||||
PartialToolCall: &agentv1.PartialToolCallUpdate{
|
||||
CallId: callID,
|
||||
ToolCall: toolCall,
|
||||
ArgsTextDelta: argsTextDelta,
|
||||
ModelCallId: modelCallID,
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// buildToolCallDeltaMessage 构造工具调用流式增量消息。
|
||||
func buildToolCallDeltaMessage(callID string, modelCallID string, delta *agentv1.ToolCallDelta) *agentv1.AgentServerMessage {
|
||||
return &agentv1.AgentServerMessage{
|
||||
Message: &agentv1.AgentServerMessage_InteractionUpdate{
|
||||
InteractionUpdate: &agentv1.InteractionUpdate{
|
||||
Message: &agentv1.InteractionUpdate_ToolCallDelta{
|
||||
ToolCallDelta: &agentv1.ToolCallDeltaUpdate{
|
||||
CallId: callID,
|
||||
ToolCallDelta: delta,
|
||||
ModelCallId: modelCallID,
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// buildToolCallCompletedMessage 构造工具调用完成消息。
|
||||
func buildToolCallCompletedMessage(callID string, modelCallID string, toolCall *agentv1.ToolCall) *agentv1.AgentServerMessage {
|
||||
return &agentv1.AgentServerMessage{
|
||||
Message: &agentv1.AgentServerMessage_InteractionUpdate{
|
||||
InteractionUpdate: &agentv1.InteractionUpdate{
|
||||
Message: &agentv1.InteractionUpdate_ToolCallCompleted{
|
||||
ToolCallCompleted: &agentv1.ToolCallCompletedUpdate{
|
||||
CallId: callID,
|
||||
ToolCall: toolCall,
|
||||
ModelCallId: modelCallID,
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// buildShellOutputDeltaMessage 把 shell 流输出包装成兼容消息。
|
||||
func buildShellOutputDeltaMessage(delta *agentv1.ShellOutputDeltaUpdate) *agentv1.AgentServerMessage {
|
||||
return &agentv1.AgentServerMessage{
|
||||
Message: &agentv1.AgentServerMessage_InteractionUpdate{
|
||||
InteractionUpdate: &agentv1.InteractionUpdate{
|
||||
Message: &agentv1.InteractionUpdate_ShellOutputDelta{
|
||||
ShellOutputDelta: delta,
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// buildTurnEndedMessage 构造 turn 结束消息,并携带标准化后的 token 统计。
|
||||
func buildTurnEndedMessage(inputTokens int64, outputTokens int64, cacheReadTokens int64, cacheWriteTokens int64) *agentv1.AgentServerMessage {
|
||||
inputTokensValue := inputTokens
|
||||
outputTokensValue := outputTokens
|
||||
cacheReadTokensValue := cacheReadTokens
|
||||
cacheWriteTokensValue := cacheWriteTokens
|
||||
return &agentv1.AgentServerMessage{
|
||||
Message: &agentv1.AgentServerMessage_InteractionUpdate{
|
||||
InteractionUpdate: &agentv1.InteractionUpdate{
|
||||
Message: &agentv1.InteractionUpdate_TurnEnded{
|
||||
TurnEnded: &agentv1.TurnEndedUpdate{
|
||||
InputTokens: &inputTokensValue,
|
||||
OutputTokens: &outputTokensValue,
|
||||
CacheReadTokens: &cacheReadTokensValue,
|
||||
CacheWriteTokens: &cacheWriteTokensValue,
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// buildCheckpointMessage 根据投影出的状态生成 legacy checkpoint 消息。
|
||||
func buildCheckpointMessage(state *agentv1.ConversationStateStructure) *agentv1.AgentServerMessage {
|
||||
cloned := &agentv1.ConversationStateStructure{}
|
||||
if state != nil {
|
||||
if next, ok := proto.Clone(state).(*agentv1.ConversationStateStructure); ok && next != nil {
|
||||
cloned = next
|
||||
}
|
||||
}
|
||||
if cloned.TokenDetails == nil {
|
||||
cloned.TokenDetails = &agentv1.ConversationTokenDetails{}
|
||||
}
|
||||
if cloned.TokenDetails.MaxTokens == 0 {
|
||||
cloned.TokenDetails.MaxTokens = projectedConversationMaxTokens
|
||||
}
|
||||
cloned.TurnTimings = nil
|
||||
return &agentv1.AgentServerMessage{
|
||||
Message: &agentv1.AgentServerMessage_ConversationCheckpointUpdate{
|
||||
ConversationCheckpointUpdate: cloned,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// buildExecAbortMessage 构造对客户端执行桥的 abort 控制消息。
|
||||
func buildExecAbortMessage(pending runtimecore.PendingExec) *agentv1.AgentServerMessage {
|
||||
return &agentv1.AgentServerMessage{
|
||||
Message: &agentv1.AgentServerMessage_ExecServerControlMessage{
|
||||
ExecServerControlMessage: &agentv1.ExecServerControlMessage{
|
||||
Message: &agentv1.ExecServerControlMessage_Abort{
|
||||
Abort: &agentv1.ExecServerAbort{
|
||||
Id: pending.MessageID,
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// buildStartedToolCall 把工具意图映射为可发送给客户端的 started ToolCall 结构。
|
||||
func buildStartedToolCall(invocation runtimecore.ToolInvocation) *agentv1.ToolCall {
|
||||
switch strings.TrimSpace(invocation.ToolName) {
|
||||
case "Glob":
|
||||
args, err := execbridge.DecodeGlobToolArgs(invocation.ArgsJSON)
|
||||
if err != nil {
|
||||
args = &agentv1.GlobToolArgs{}
|
||||
}
|
||||
return &agentv1.ToolCall{
|
||||
Tool: &agentv1.ToolCall_GlobToolCall{
|
||||
GlobToolCall: &agentv1.GlobToolCall{Args: args},
|
||||
},
|
||||
}
|
||||
case "Grep":
|
||||
args, err := execbridge.DecodeGrepToolArgs(invocation.ArgsJSON, invocation.CallID)
|
||||
if err != nil && args == nil {
|
||||
args = &agentv1.GrepArgs{ToolCallId: invocation.CallID}
|
||||
}
|
||||
return &agentv1.ToolCall{
|
||||
Tool: &agentv1.ToolCall_GrepToolCall{
|
||||
GrepToolCall: &agentv1.GrepToolCall{Args: args},
|
||||
},
|
||||
}
|
||||
case "Read":
|
||||
args, err := execbridge.DecodeReadToolArgs(invocation.ArgsJSON)
|
||||
if err != nil && args == nil {
|
||||
args = &agentv1.ReadToolArgs{}
|
||||
}
|
||||
return &agentv1.ToolCall{
|
||||
Tool: &agentv1.ToolCall_ReadToolCall{
|
||||
ReadToolCall: &agentv1.ReadToolCall{
|
||||
Args: args,
|
||||
},
|
||||
},
|
||||
}
|
||||
case "TodoWrite":
|
||||
args, _ := decodeUpdateTodosArgsJSON(invocation.ArgsJSON)
|
||||
return &agentv1.ToolCall{
|
||||
Tool: &agentv1.ToolCall_UpdateTodosToolCall{
|
||||
UpdateTodosToolCall: &agentv1.UpdateTodosToolCall{Args: args},
|
||||
},
|
||||
}
|
||||
case "ReadTodos":
|
||||
args, _ := decodeReadTodosArgsJSON(invocation.ArgsJSON)
|
||||
return &agentv1.ToolCall{
|
||||
Tool: &agentv1.ToolCall_ReadTodosToolCall{
|
||||
ReadTodosToolCall: &agentv1.ReadTodosToolCall{Args: args},
|
||||
},
|
||||
}
|
||||
case "Delete":
|
||||
var input struct {
|
||||
Path string `json:"path"`
|
||||
}
|
||||
_ = json.Unmarshal(invocation.ArgsJSON, &input)
|
||||
return &agentv1.ToolCall{
|
||||
Tool: &agentv1.ToolCall_DeleteToolCall{
|
||||
DeleteToolCall: &agentv1.DeleteToolCall{
|
||||
Args: &agentv1.DeleteArgs{Path: strings.TrimSpace(input.Path), ToolCallId: invocation.CallID},
|
||||
},
|
||||
},
|
||||
}
|
||||
case "Shell":
|
||||
var input struct {
|
||||
Command string `json:"command"`
|
||||
Description string `json:"description,omitempty"`
|
||||
WorkingDirectory string `json:"working_directory,omitempty"`
|
||||
NotifyOnOutput *struct {
|
||||
Pattern string `json:"pattern"`
|
||||
Reason string `json:"reason"`
|
||||
DebounceMS *float64 `json:"debounce_ms,omitempty"`
|
||||
NotificationLimit *int32 `json:"notification_limit,omitempty"`
|
||||
} `json:"notify_on_output,omitempty"`
|
||||
}
|
||||
_ = json.Unmarshal(invocation.ArgsJSON, &input)
|
||||
return &agentv1.ToolCall{
|
||||
Tool: &agentv1.ToolCall_ShellToolCall{
|
||||
ShellToolCall: &agentv1.ShellToolCall{
|
||||
Args: &agentv1.ShellArgs{
|
||||
Command: strings.TrimSpace(input.Command),
|
||||
WorkingDirectory: strings.TrimSpace(input.WorkingDirectory),
|
||||
ToolCallId: invocation.CallID,
|
||||
Description: stringPtr(strings.TrimSpace(input.Description)),
|
||||
OutputNotification: buildShellOutputNotificationConfig(input.NotifyOnOutput),
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
case "AwaitShell":
|
||||
args, err := decodeAwaitShellArgs(invocation.ArgsJSON)
|
||||
if err != nil {
|
||||
args = awaitShellArgs{}
|
||||
}
|
||||
return buildAwaitShellToolCall(buildAwaitArgsFromAwaitShellArgs(args), nil)
|
||||
case "WriteShellStdin":
|
||||
var input struct {
|
||||
ShellID uint32 `json:"shell_id"`
|
||||
Chars string `json:"chars"`
|
||||
}
|
||||
_ = json.Unmarshal(invocation.ArgsJSON, &input)
|
||||
return &agentv1.ToolCall{
|
||||
Tool: &agentv1.ToolCall_WriteShellStdinToolCall{
|
||||
WriteShellStdinToolCall: &agentv1.WriteShellStdinToolCall{
|
||||
Args: &agentv1.WriteShellStdinArgs{
|
||||
ShellId: input.ShellID,
|
||||
Chars: input.Chars,
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
case "Task":
|
||||
return &agentv1.ToolCall{
|
||||
Tool: &agentv1.ToolCall_TaskToolCall{
|
||||
TaskToolCall: &agentv1.TaskToolCall{
|
||||
Args: buildTaskArgsFromJSON(invocation.ArgsJSON),
|
||||
},
|
||||
},
|
||||
}
|
||||
case "Ls":
|
||||
var input struct {
|
||||
Path string `json:"path"`
|
||||
Ignore []string `json:"ignore,omitempty"`
|
||||
}
|
||||
_ = json.Unmarshal(invocation.ArgsJSON, &input)
|
||||
return &agentv1.ToolCall{
|
||||
Tool: &agentv1.ToolCall_LsToolCall{
|
||||
LsToolCall: &agentv1.LsToolCall{
|
||||
Args: &agentv1.LsArgs{
|
||||
Path: strings.TrimSpace(input.Path),
|
||||
Ignore: append([]string(nil), input.Ignore...),
|
||||
ToolCallId: invocation.CallID,
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
case "Write":
|
||||
var input struct {
|
||||
Path string `json:"path"`
|
||||
Contents string `json:"contents"`
|
||||
}
|
||||
_ = json.Unmarshal(invocation.ArgsJSON, &input)
|
||||
return &agentv1.ToolCall{
|
||||
Tool: &agentv1.ToolCall_EditToolCall{
|
||||
EditToolCall: &agentv1.EditToolCall{
|
||||
Args: &agentv1.EditArgs{
|
||||
Path: strings.TrimSpace(input.Path),
|
||||
StreamContent: literalStringPtr(input.Contents),
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
case "PatchEdit":
|
||||
var input struct {
|
||||
FilePath string `json:"file_path"`
|
||||
Path string `json:"path,omitempty"`
|
||||
}
|
||||
_ = json.Unmarshal(invocation.ArgsJSON, &input)
|
||||
path := strings.TrimSpace(input.FilePath)
|
||||
if path == "" {
|
||||
path = strings.TrimSpace(input.Path)
|
||||
}
|
||||
return &agentv1.ToolCall{
|
||||
Tool: &agentv1.ToolCall_EditToolCall{
|
||||
EditToolCall: &agentv1.EditToolCall{
|
||||
Args: &agentv1.EditArgs{
|
||||
Path: path,
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
case "CallMcpTool":
|
||||
var input struct {
|
||||
Server string `json:"server"`
|
||||
ToolName string `json:"toolName"`
|
||||
Arguments map[string]any `json:"arguments,omitempty"`
|
||||
}
|
||||
_ = json.Unmarshal(invocation.ArgsJSON, &input)
|
||||
return &agentv1.ToolCall{
|
||||
Tool: &agentv1.ToolCall_McpToolCall{
|
||||
McpToolCall: &agentv1.McpToolCall{
|
||||
Args: &agentv1.McpArgs{
|
||||
Name: forwarderCanonicalMCPToolLookupName(input.Server, input.ToolName),
|
||||
Args: forwarderStructValueMap(input.Arguments),
|
||||
ToolCallId: invocation.CallID,
|
||||
ProviderIdentifier: strings.TrimSpace(input.Server),
|
||||
ToolName: strings.TrimSpace(input.ToolName),
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
case "ListMcpResources":
|
||||
var input struct {
|
||||
Server string `json:"server,omitempty"`
|
||||
}
|
||||
_ = json.Unmarshal(invocation.ArgsJSON, &input)
|
||||
return &agentv1.ToolCall{
|
||||
Tool: &agentv1.ToolCall_ListMcpResourcesToolCall{
|
||||
ListMcpResourcesToolCall: &agentv1.ListMcpResourcesToolCall{
|
||||
Args: &agentv1.ListMcpResourcesExecArgs{
|
||||
Server: stringPtr(strings.TrimSpace(input.Server)),
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
case "AskQuestion":
|
||||
var args agentv1.AskQuestionArgs
|
||||
_ = json.Unmarshal(invocation.ArgsJSON, &args)
|
||||
return &agentv1.ToolCall{
|
||||
Tool: &agentv1.ToolCall_AskQuestionToolCall{
|
||||
AskQuestionToolCall: &agentv1.AskQuestionToolCall{Args: &args},
|
||||
},
|
||||
}
|
||||
case "WebSearch":
|
||||
var args agentv1.WebSearchArgs
|
||||
_ = json.Unmarshal(invocation.ArgsJSON, &args)
|
||||
return &agentv1.ToolCall{
|
||||
Tool: &agentv1.ToolCall_WebSearchToolCall{
|
||||
WebSearchToolCall: &agentv1.WebSearchToolCall{Args: &args},
|
||||
},
|
||||
}
|
||||
case "WebFetch":
|
||||
var args agentv1.WebFetchArgs
|
||||
_ = json.Unmarshal(invocation.ArgsJSON, &args)
|
||||
return &agentv1.ToolCall{
|
||||
Tool: &agentv1.ToolCall_WebFetchToolCall{
|
||||
WebFetchToolCall: &agentv1.WebFetchToolCall{Args: &args},
|
||||
},
|
||||
}
|
||||
case "CreatePlan":
|
||||
args, err := runtimecore.DecodeCreatePlanArgsJSON(invocation.ArgsJSON)
|
||||
if err != nil {
|
||||
args = &agentv1.CreatePlanArgs{}
|
||||
}
|
||||
return &agentv1.ToolCall{
|
||||
Tool: &agentv1.ToolCall_CreatePlanToolCall{
|
||||
CreatePlanToolCall: &agentv1.CreatePlanToolCall{Args: args},
|
||||
},
|
||||
}
|
||||
case "GenerateImage":
|
||||
carrier, _ := decodeGenerateImageToolCarrier(invocation.ArgsJSON)
|
||||
args := buildGenerateImageArgsFromCarrier(carrier)
|
||||
return &agentv1.ToolCall{
|
||||
Tool: &agentv1.ToolCall_GenerateImageToolCall{
|
||||
GenerateImageToolCall: &agentv1.GenerateImageToolCall{Args: args},
|
||||
},
|
||||
}
|
||||
case "SwitchMode":
|
||||
var args agentv1.SwitchModeArgs
|
||||
_ = json.Unmarshal(invocation.ArgsJSON, &args)
|
||||
args.ToolCallId = invocation.CallID
|
||||
return &agentv1.ToolCall{
|
||||
Tool: &agentv1.ToolCall_SwitchModeToolCall{
|
||||
SwitchModeToolCall: &agentv1.SwitchModeToolCall{Args: &args},
|
||||
},
|
||||
}
|
||||
case "FetchMcpResource":
|
||||
var input struct {
|
||||
Server string `json:"server"`
|
||||
URI string `json:"uri"`
|
||||
DownloadPath string `json:"downloadPath,omitempty"`
|
||||
}
|
||||
_ = json.Unmarshal(invocation.ArgsJSON, &input)
|
||||
return &agentv1.ToolCall{
|
||||
Tool: &agentv1.ToolCall_ReadMcpResourceToolCall{
|
||||
ReadMcpResourceToolCall: &agentv1.ReadMcpResourceToolCall{
|
||||
Args: &agentv1.ReadMcpResourceExecArgs{
|
||||
Server: strings.TrimSpace(input.Server),
|
||||
Uri: strings.TrimSpace(input.URI),
|
||||
DownloadPath: stringPtr(strings.TrimSpace(input.DownloadPath)),
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
default:
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
func forwarderCanonicalMCPToolLookupName(server string, toolName string) string {
|
||||
trimmedServer := strings.TrimSpace(server)
|
||||
trimmedToolName := strings.TrimSpace(toolName)
|
||||
if trimmedToolName == "" {
|
||||
return ""
|
||||
}
|
||||
if trimmedServer == "" {
|
||||
return trimmedToolName
|
||||
}
|
||||
return trimmedServer + "-" + trimmedToolName
|
||||
}
|
||||
|
||||
func forwarderStructValueMap(items map[string]any) map[string]*structpb.Value {
|
||||
if len(items) == 0 {
|
||||
return map[string]*structpb.Value{}
|
||||
}
|
||||
result := make(map[string]*structpb.Value, len(items))
|
||||
for key, value := range items {
|
||||
item, err := structpb.NewValue(value)
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
result[key] = item
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
// stringPtr 在字符串非空时返回指针,避免把空串写入 optional 字段。
|
||||
func buildShellOutputNotificationConfig(input *struct {
|
||||
Pattern string `json:"pattern"`
|
||||
Reason string `json:"reason"`
|
||||
DebounceMS *float64 `json:"debounce_ms,omitempty"`
|
||||
NotificationLimit *int32 `json:"notification_limit,omitempty"`
|
||||
}) *agentv1.ShellOutputNotificationConfig {
|
||||
if input == nil {
|
||||
return nil
|
||||
}
|
||||
pattern := strings.TrimSpace(input.Pattern)
|
||||
reason := strings.TrimSpace(input.Reason)
|
||||
if pattern == "" || reason == "" {
|
||||
return nil
|
||||
}
|
||||
var debounce *float64
|
||||
if input.DebounceMS != nil {
|
||||
value := *input.DebounceMS / 1000
|
||||
if value < 5 {
|
||||
value = 5
|
||||
}
|
||||
debounce = &value
|
||||
}
|
||||
return &agentv1.ShellOutputNotificationConfig{
|
||||
Pattern: pattern,
|
||||
Reason: reason,
|
||||
Debounce: debounce,
|
||||
NotificationLimit: input.NotificationLimit,
|
||||
}
|
||||
}
|
||||
|
||||
func stringPtr(value string) *string {
|
||||
if strings.TrimSpace(value) == "" {
|
||||
return nil
|
||||
}
|
||||
next := strings.TrimSpace(value)
|
||||
return &next
|
||||
}
|
||||
|
||||
func literalStringPtr(value string) *string {
|
||||
if value == "" {
|
||||
return nil
|
||||
}
|
||||
next := value
|
||||
return &next
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,172 @@
|
||||
package forwarder
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"log"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
const historyMaintenanceLockStaleAfter = 30 * time.Minute
|
||||
const historyMaintenanceTempStaleAfter = 5 * time.Minute
|
||||
|
||||
func (service *Service) startHistoryMaintenance() {
|
||||
if service == nil || service.store == nil {
|
||||
return
|
||||
}
|
||||
go func() {
|
||||
if err := service.runHistoryMaintenance(); err != nil {
|
||||
log.Printf("forwarder history maintenance failed: %v", err)
|
||||
}
|
||||
}()
|
||||
}
|
||||
|
||||
func (service *Service) runHistoryMaintenance() error {
|
||||
historyRoot := strings.TrimSpace(service.store.HistoryDir())
|
||||
if historyRoot == "" {
|
||||
return nil
|
||||
}
|
||||
release, ok, err := acquireHistoryMaintenanceLock(historyRoot)
|
||||
if err != nil || !ok {
|
||||
return err
|
||||
}
|
||||
defer release()
|
||||
|
||||
entries, err := os.ReadDir(historyRoot)
|
||||
if err != nil {
|
||||
if errors.Is(err, os.ErrNotExist) {
|
||||
return nil
|
||||
}
|
||||
return err
|
||||
}
|
||||
for _, entry := range entries {
|
||||
if !entry.IsDir() {
|
||||
cleanupRootLegacyHistoryArtifact(historyRoot, entry.Name())
|
||||
continue
|
||||
}
|
||||
service.cleanupConversationLegacyArtifacts(filepath.Join(historyRoot, entry.Name()))
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func cleanupRootLegacyHistoryArtifact(historyRoot string, name string) {
|
||||
trimmedName := strings.TrimSpace(name)
|
||||
if trimmedName == "" || trimmedName == ".history-maintenance.lock" {
|
||||
return
|
||||
}
|
||||
if !isRootLegacyHistoryArtifact(trimmedName) {
|
||||
return
|
||||
}
|
||||
path := filepath.Join(historyRoot, trimmedName)
|
||||
if strings.Contains(trimmedName, ".tmp-") {
|
||||
cleanupStaleHistoryTempArtifact(path)
|
||||
return
|
||||
}
|
||||
if err := os.RemoveAll(path); err != nil && !errors.Is(err, os.ErrNotExist) {
|
||||
log.Printf("forwarder root legacy cleanup failed path=%s err=%v", path, err)
|
||||
}
|
||||
}
|
||||
|
||||
func isRootLegacyHistoryArtifact(name string) bool {
|
||||
if strings.TrimSpace(name) == usageFileName || strings.TrimSpace(name) == usageFileName+".lock" {
|
||||
return false
|
||||
}
|
||||
if strings.HasSuffix(name, ".json") || strings.HasSuffix(name, ".jsonl") {
|
||||
return true
|
||||
}
|
||||
if strings.Contains(name, ".tmp-") {
|
||||
return true
|
||||
}
|
||||
if strings.HasSuffix(name, ".lock") {
|
||||
return true
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func (service *Service) cleanupConversationLegacyArtifacts(conversationDir string) {
|
||||
for _, name := range []string{
|
||||
"turns",
|
||||
"active",
|
||||
"latest.json",
|
||||
"summary.json",
|
||||
"replay.json",
|
||||
"runtime.json",
|
||||
"request.json",
|
||||
"recovery.json",
|
||||
"conversation.json",
|
||||
"entries.jsonl",
|
||||
} {
|
||||
path := filepath.Join(conversationDir, name)
|
||||
if err := os.RemoveAll(path); err != nil && !errors.Is(err, os.ErrNotExist) {
|
||||
log.Printf("forwarder legacy cleanup failed path=%s err=%v", path, err)
|
||||
}
|
||||
}
|
||||
entries, err := os.ReadDir(conversationDir)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
for _, entry := range entries {
|
||||
if !entry.IsDir() {
|
||||
if strings.Contains(entry.Name(), ".tmp-") {
|
||||
cleanupStaleHistoryTempArtifact(filepath.Join(conversationDir, entry.Name()))
|
||||
}
|
||||
continue
|
||||
}
|
||||
if _, err := strconv.Atoi(strings.TrimSpace(entry.Name())); err != nil {
|
||||
continue
|
||||
}
|
||||
path := filepath.Join(conversationDir, entry.Name())
|
||||
if err := os.RemoveAll(path); err != nil && !errors.Is(err, os.ErrNotExist) {
|
||||
log.Printf("forwarder numeric legacy cleanup failed path=%s err=%v", path, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func cleanupStaleHistoryTempArtifact(path string) {
|
||||
info, err := os.Stat(path)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
if time.Since(info.ModTime()) < historyMaintenanceTempStaleAfter {
|
||||
return
|
||||
}
|
||||
if err := os.Remove(path); err != nil && !errors.Is(err, os.ErrNotExist) {
|
||||
log.Printf("forwarder temp cleanup failed path=%s err=%v", path, err)
|
||||
}
|
||||
}
|
||||
|
||||
func acquireHistoryMaintenanceLock(historyRoot string) (func(), bool, error) {
|
||||
if err := os.MkdirAll(historyRoot, 0o755); err != nil {
|
||||
return nil, false, err
|
||||
}
|
||||
lockPath := filepath.Join(historyRoot, ".history-maintenance.lock")
|
||||
for attempt := 0; attempt < 2; attempt++ {
|
||||
file, err := os.OpenFile(lockPath, os.O_CREATE|os.O_EXCL|os.O_WRONLY, 0o600)
|
||||
if err == nil {
|
||||
_, _ = file.WriteString(time.Now().UTC().Format(time.RFC3339Nano))
|
||||
_ = file.Close()
|
||||
return func() {
|
||||
_ = os.Remove(lockPath)
|
||||
}, true, nil
|
||||
}
|
||||
if !errors.Is(err, os.ErrExist) {
|
||||
return nil, false, err
|
||||
}
|
||||
info, statErr := os.Stat(lockPath)
|
||||
if statErr != nil {
|
||||
if errors.Is(statErr, os.ErrNotExist) {
|
||||
continue
|
||||
}
|
||||
return nil, false, statErr
|
||||
}
|
||||
if time.Since(info.ModTime()) > historyMaintenanceLockStaleAfter {
|
||||
_ = os.Remove(lockPath)
|
||||
continue
|
||||
}
|
||||
return nil, false, nil
|
||||
}
|
||||
return nil, false, nil
|
||||
}
|
||||
@@ -0,0 +1,533 @@
|
||||
package forwarder
|
||||
|
||||
import (
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"cursor/gen/agentv1"
|
||||
runtimecore "cursor/internal/backend/agent/core"
|
||||
)
|
||||
|
||||
func (service *Service) handleInteractionToolInvocation(stream *ActiveStream, invocation runtimecore.ToolInvocation) error {
|
||||
if service == nil || stream == nil {
|
||||
return fmt.Errorf("active stream is required")
|
||||
}
|
||||
serverMessage, pendingInteraction, err := service.interactionBridge.OpenQuery(invocation)
|
||||
if err != nil {
|
||||
return newRecoverableToolInvocationError(err)
|
||||
}
|
||||
pendingInteraction.ModelCallID = invocation.ModelCallID
|
||||
pendingInteraction.ReasoningContent = invocation.ReasoningContent
|
||||
pendingInteraction.ReasoningSignature = invocation.ReasoningSignature
|
||||
pendingInteraction.ReasoningSignatureSource = invocation.ReasoningSignatureSource
|
||||
if pendingInteraction.OpenedAt.IsZero() {
|
||||
pendingInteraction.OpenedAt = time.Now().UTC()
|
||||
}
|
||||
|
||||
stream.mu.Lock()
|
||||
if stream.PendingInteractions == nil {
|
||||
stream.PendingInteractions = make(map[string]runtimecore.PendingInteraction)
|
||||
}
|
||||
pendingInteraction.ProviderPass = stream.ProviderPassCount
|
||||
stream.PendingInteractions[pendingInteraction.InteractionID] = pendingInteraction
|
||||
stream.UpdatedAt = time.Now().UTC()
|
||||
stream.mu.Unlock()
|
||||
|
||||
removePending := func() {
|
||||
stream.mu.Lock()
|
||||
delete(stream.PendingInteractions, pendingInteraction.InteractionID)
|
||||
stream.UpdatedAt = time.Now().UTC()
|
||||
stream.mu.Unlock()
|
||||
}
|
||||
if err := service.publishCheckpoint(stream.RequestID, stream.ConversationID); err != nil {
|
||||
removePending()
|
||||
return err
|
||||
}
|
||||
if err := service.broker.Publish(stream.RequestID, StreamEvent{Message: serverMessage}); err != nil {
|
||||
removePending()
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (service *Service) handleInteractionResult(intent InboundIntent) error {
|
||||
stream, ok := service.broker.Get(intent.RequestID)
|
||||
if !ok || stream == nil {
|
||||
return fmt.Errorf("request is not active: %s", intent.RequestID)
|
||||
}
|
||||
if intent.InteractionResponse == nil {
|
||||
return fmt.Errorf("interaction response is required")
|
||||
}
|
||||
pending, found := selectPendingInteraction(intent.InteractionResponse, stream)
|
||||
if !found {
|
||||
return fmt.Errorf("pending interaction not found")
|
||||
}
|
||||
result, err := service.interactionBridge.ApplyInteractionResponse(intent.InteractionResponse, pending)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
markInteractionCompleted(stream, pending)
|
||||
toolName := strings.TrimSpace(deriveToolNameFromPendingInteraction(pending))
|
||||
if result.ToolCall != nil {
|
||||
applySwitchModeMetadata(stream, result.ToolCall)
|
||||
if err := service.appendToolResult(stream, result.ToolCallID, toolName, pending.ArgsJSON, result.ToolResultPayload, pending.ReasoningContent, result.ToolCall); err != nil {
|
||||
return err
|
||||
}
|
||||
} else if strings.TrimSpace(result.ToolResultPayload) != "" {
|
||||
if err := service.appendToolResult(stream, pending.ToolCallID, toolName, pending.ArgsJSON, result.ToolResultPayload, pending.ReasoningContent, nil); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
if err := service.publishToolCallCompleted(intent.RequestID, result.ToolCallID, pending.ModelCallID, result.ToolCall); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := service.applyApprovedSwitchMode(stream, stream.ConversationID, result.ToolCall); err != nil {
|
||||
return err
|
||||
}
|
||||
if switchToolCall := result.ToolCall.GetSwitchModeToolCall(); switchToolCall != nil && switchToolCall.GetResult().GetSuccess() != nil {
|
||||
targetMode, err := switchModeTarget(switchToolCall.GetArgs())
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
modeEntry, err := newModeMetadataEntry(stream.TurnSeq, intent.RequestID, targetMode, true, ModeSourceSwitchModeTool)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if _, err := service.appendConversationEntries(stream, stream.ConversationID, []HistoryEntry{
|
||||
modeEntry,
|
||||
newModeChangePromptContextEntry(stream.TurnSeq, intent.RequestID, targetMode),
|
||||
}); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
if err := service.syncSummaryCarryForward(stream.ConversationID, intent.RequestID, pending.ModelCallID); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := service.publishCheckpoint(intent.RequestID, stream.ConversationID); err != nil {
|
||||
return err
|
||||
}
|
||||
if !shouldAutoResumeAfterInteraction(pending) {
|
||||
rememberPendingProviderCompletion(stream, pendingTurnCompletion{
|
||||
ConversationID: stream.ConversationID,
|
||||
RequestID: stream.RequestID,
|
||||
TurnSeq: stream.TurnSeq,
|
||||
ModelCallID: pending.ModelCallID,
|
||||
ProviderPass: pending.ProviderPass,
|
||||
Disposition: completionDispositionCompleteAfterExternal,
|
||||
})
|
||||
}
|
||||
return service.reconcileStream(stream)
|
||||
}
|
||||
|
||||
func (service *Service) applyApprovedSwitchMode(stream *ActiveStream, conversationID string, toolCall *agentv1.ToolCall) error {
|
||||
switchToolCall := toolCall.GetSwitchModeToolCall()
|
||||
if switchToolCall == nil || switchToolCall.GetResult().GetSuccess() == nil {
|
||||
return nil
|
||||
}
|
||||
targetMode, err := switchModeTarget(switchToolCall.GetArgs())
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
targetAlias, err := modeAlias(targetMode)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
stream.mu.Lock()
|
||||
stream.Mode = targetMode
|
||||
stream.UpdatedAt = time.Now().UTC()
|
||||
stream.mu.Unlock()
|
||||
_, err = service.updateConversationMetaAndCheckpoint(stream, conversationID, func(item *ConversationFile) error {
|
||||
if item == nil {
|
||||
return nil
|
||||
}
|
||||
item.Mode = targetAlias
|
||||
return nil
|
||||
})
|
||||
return err
|
||||
}
|
||||
|
||||
func applySwitchModeMetadata(stream *ActiveStream, toolCall *agentv1.ToolCall) {
|
||||
switchToolCall := toolCall.GetSwitchModeToolCall()
|
||||
if switchToolCall == nil {
|
||||
return
|
||||
}
|
||||
success := switchToolCall.GetResult().GetSuccess()
|
||||
if success == nil {
|
||||
return
|
||||
}
|
||||
if fromAlias, err := modeAlias(currentStreamMode(stream)); err == nil {
|
||||
success.FromModeId = fromAlias
|
||||
}
|
||||
}
|
||||
|
||||
func switchModeTarget(args *agentv1.SwitchModeArgs) (agentv1.AgentMode, error) {
|
||||
if args == nil {
|
||||
return agentv1.AgentMode_AGENT_MODE_UNSPECIFIED, fmt.Errorf("switch mode args are required")
|
||||
}
|
||||
return parseTargetModeID(args.GetTargetModeId())
|
||||
}
|
||||
|
||||
func markInteractionCompleted(stream *ActiveStream, pending runtimecore.PendingInteraction) {
|
||||
if stream == nil {
|
||||
return
|
||||
}
|
||||
stream.mu.Lock()
|
||||
delete(stream.PendingInteractions, pending.InteractionID)
|
||||
stream.UpdatedAt = time.Now().UTC()
|
||||
stream.mu.Unlock()
|
||||
}
|
||||
|
||||
func deriveToolNameFromPendingInteraction(pending runtimecore.PendingInteraction) string {
|
||||
switch strings.TrimSpace(pending.InteractionKind) {
|
||||
case "ask_question":
|
||||
return "AskQuestion"
|
||||
case "create_plan":
|
||||
return "CreatePlan"
|
||||
case "web_search":
|
||||
return "WebSearch"
|
||||
case "web_fetch":
|
||||
return "WebFetch"
|
||||
case "switch_mode":
|
||||
return "SwitchMode"
|
||||
default:
|
||||
return ""
|
||||
}
|
||||
}
|
||||
|
||||
func isInteractionTool(name string) bool {
|
||||
switch strings.TrimSpace(name) {
|
||||
case "AskQuestion", "CreatePlan", "WebSearch", "WebFetch", "SwitchMode":
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func shouldAutoResumeAfterInteraction(pending runtimecore.PendingInteraction) bool {
|
||||
switch strings.TrimSpace(pending.InteractionKind) {
|
||||
case "create_plan":
|
||||
return false
|
||||
default:
|
||||
return true
|
||||
}
|
||||
}
|
||||
|
||||
type generateImageToolCarrier struct {
|
||||
Description string `json:"description,omitempty"`
|
||||
FilePath string `json:"file_path,omitempty"`
|
||||
ReferenceImagePaths []string `json:"reference_image_paths,omitempty"`
|
||||
ImageData string `json:"image_data,omitempty"`
|
||||
}
|
||||
|
||||
func isImmediateNativeTool(name string) bool {
|
||||
switch strings.TrimSpace(name) {
|
||||
case "GenerateImage", "AwaitShell":
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func (service *Service) handleImmediateNativeToolInvocation(stream *ActiveStream, invocation runtimecore.ToolInvocation) error {
|
||||
switch strings.TrimSpace(invocation.ToolName) {
|
||||
case "GenerateImage":
|
||||
return service.handleGenerateImageToolInvocation(stream, invocation)
|
||||
case "AwaitShell":
|
||||
return service.handleAwaitShellToolInvocation(stream, invocation)
|
||||
default:
|
||||
return fmt.Errorf("unsupported immediate native tool: %s", invocation.ToolName)
|
||||
}
|
||||
}
|
||||
|
||||
func (service *Service) handleGenerateImageToolInvocation(stream *ActiveStream, invocation runtimecore.ToolInvocation) error {
|
||||
carrier, decodeErr := decodeGenerateImageToolCarrier(invocation.ArgsJSON)
|
||||
args := buildGenerateImageArgsFromCarrier(carrier)
|
||||
sanitizedInvocation := invocation
|
||||
sanitizedInvocation.ArgsJSON = encodeGenerateImageArgsForHistory(args)
|
||||
if decodeErr != nil {
|
||||
result, payload := buildGenerateImageErrorResult(decodeErr.Error())
|
||||
return service.completeImmediateToolResult(stream, sanitizedInvocation, payload, buildGenerateImageToolCall(args, result))
|
||||
}
|
||||
imageData, err := normalizeGenerateImageBase64(carrier.ImageData)
|
||||
if err != nil {
|
||||
result, payload := buildGenerateImageErrorResult(err.Error())
|
||||
return service.completeImmediateToolResult(stream, sanitizedInvocation, payload, buildGenerateImageToolCall(args, result))
|
||||
}
|
||||
result, payload := buildGenerateImageSuccessResult(args.GetFilePath(), imageData)
|
||||
if err := service.completeImmediateToolResult(stream, sanitizedInvocation, payload, buildGenerateImageToolCall(args, result)); err != nil {
|
||||
return err
|
||||
}
|
||||
markProviderTerminalToolInvocation(stream)
|
||||
return nil
|
||||
}
|
||||
|
||||
func markProviderTerminalToolInvocation(stream *ActiveStream) {
|
||||
if stream == nil {
|
||||
return
|
||||
}
|
||||
stream.mu.Lock()
|
||||
stream.ProviderTerminalToolInvocation = true
|
||||
stream.UpdatedAt = time.Now().UTC()
|
||||
stream.mu.Unlock()
|
||||
}
|
||||
|
||||
func decodeGenerateImageToolCarrier(raw []byte) (generateImageToolCarrier, error) {
|
||||
var carrier generateImageToolCarrier
|
||||
if len(raw) == 0 {
|
||||
return carrier, nil
|
||||
}
|
||||
if err := json.Unmarshal(raw, &carrier); err != nil {
|
||||
return carrier, fmt.Errorf("decode GenerateImage args failed: %w", err)
|
||||
}
|
||||
carrier.Description = strings.TrimSpace(carrier.Description)
|
||||
carrier.FilePath = strings.TrimSpace(carrier.FilePath)
|
||||
carrier.ImageData = strings.TrimSpace(carrier.ImageData)
|
||||
carrier.ReferenceImagePaths = compactTrimmedStrings(carrier.ReferenceImagePaths)
|
||||
return carrier, nil
|
||||
}
|
||||
|
||||
func compactTrimmedStrings(values []string) []string {
|
||||
if len(values) == 0 {
|
||||
return nil
|
||||
}
|
||||
items := make([]string, 0, len(values))
|
||||
for _, value := range values {
|
||||
trimmed := strings.TrimSpace(value)
|
||||
if trimmed != "" {
|
||||
items = append(items, trimmed)
|
||||
}
|
||||
}
|
||||
return items
|
||||
}
|
||||
|
||||
func buildGenerateImageArgsFromCarrier(carrier generateImageToolCarrier) *agentv1.GenerateImageArgs {
|
||||
args := &agentv1.GenerateImageArgs{
|
||||
Description: strings.TrimSpace(carrier.Description),
|
||||
ReferenceImagePaths: append([]string(nil), carrier.ReferenceImagePaths...),
|
||||
}
|
||||
if filePath := strings.TrimSpace(carrier.FilePath); filePath != "" {
|
||||
args.FilePath = &filePath
|
||||
}
|
||||
return args
|
||||
}
|
||||
|
||||
func encodeGenerateImageArgsForHistory(args *agentv1.GenerateImageArgs) []byte {
|
||||
payload := map[string]any{}
|
||||
if args != nil {
|
||||
if description := strings.TrimSpace(args.GetDescription()); description != "" {
|
||||
payload["description"] = description
|
||||
}
|
||||
if filePath := strings.TrimSpace(args.GetFilePath()); filePath != "" {
|
||||
payload["file_path"] = filePath
|
||||
}
|
||||
if referenceImagePaths := compactTrimmedStrings(args.GetReferenceImagePaths()); len(referenceImagePaths) > 0 {
|
||||
payload["reference_image_paths"] = referenceImagePaths
|
||||
}
|
||||
}
|
||||
encoded, err := json.Marshal(payload)
|
||||
if err != nil {
|
||||
return []byte("{}")
|
||||
}
|
||||
return encoded
|
||||
}
|
||||
|
||||
func normalizeGenerateImageBase64(raw string) (string, error) {
|
||||
trimmed := strings.TrimSpace(raw)
|
||||
if trimmed == "" {
|
||||
return "", fmt.Errorf("GenerateImage image_data is required")
|
||||
}
|
||||
if strings.HasPrefix(strings.ToLower(trimmed), "data:image/") {
|
||||
return "", fmt.Errorf("GenerateImage image_data must be raw base64 without a data URL prefix")
|
||||
}
|
||||
compact := strings.Join(strings.Fields(trimmed), "")
|
||||
if compact == "" {
|
||||
return "", fmt.Errorf("GenerateImage image_data is required")
|
||||
}
|
||||
if _, err := base64.StdEncoding.DecodeString(compact); err != nil {
|
||||
return "", fmt.Errorf("GenerateImage image_data must be valid base64: %w", err)
|
||||
}
|
||||
return compact, nil
|
||||
}
|
||||
|
||||
func buildGenerateImageToolCall(args *agentv1.GenerateImageArgs, result *agentv1.GenerateImageResult) *agentv1.ToolCall {
|
||||
return &agentv1.ToolCall{
|
||||
Tool: &agentv1.ToolCall_GenerateImageToolCall{
|
||||
GenerateImageToolCall: &agentv1.GenerateImageToolCall{
|
||||
Args: args,
|
||||
Result: result,
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func buildGenerateImageSuccessResult(filePath string, imageData string) (*agentv1.GenerateImageResult, string) {
|
||||
trimmedFilePath := strings.TrimSpace(filePath)
|
||||
payload := fmt.Sprintf("generate image success image_data_bytes=%d", len(imageData))
|
||||
if trimmedFilePath != "" {
|
||||
payload = fmt.Sprintf("generate image success file_path=%s image_data_bytes=%d", trimmedFilePath, len(imageData))
|
||||
}
|
||||
return &agentv1.GenerateImageResult{
|
||||
Result: &agentv1.GenerateImageResult_Success{
|
||||
Success: &agentv1.GenerateImageSuccess{
|
||||
FilePath: trimmedFilePath,
|
||||
ImageData: imageData,
|
||||
},
|
||||
},
|
||||
}, payload
|
||||
}
|
||||
|
||||
func buildGenerateImageErrorResult(message string) (*agentv1.GenerateImageResult, string) {
|
||||
errorText := strings.TrimSpace(message)
|
||||
if errorText == "" {
|
||||
errorText = "image generation failed"
|
||||
}
|
||||
return &agentv1.GenerateImageResult{
|
||||
Result: &agentv1.GenerateImageResult_Error{
|
||||
Error: &agentv1.GenerateImageError{Error: errorText},
|
||||
},
|
||||
}, errorText
|
||||
}
|
||||
|
||||
func isLocalStateTool(name string) bool {
|
||||
switch strings.TrimSpace(name) {
|
||||
case "TodoWrite":
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func (service *Service) handleLocalStateToolInvocation(stream *ActiveStream, invocation runtimecore.ToolInvocation) error {
|
||||
switch strings.TrimSpace(invocation.ToolName) {
|
||||
case "TodoWrite":
|
||||
return service.handleTodoWriteToolInvocation(stream, invocation)
|
||||
default:
|
||||
return fmt.Errorf("unsupported local state tool: %s", invocation.ToolName)
|
||||
}
|
||||
}
|
||||
|
||||
func (service *Service) handleTodoWriteToolInvocation(stream *ActiveStream, invocation runtimecore.ToolInvocation) error {
|
||||
decodedArgs, err := decodeUpdateTodosArgsJSONWithPresence(invocation.ArgsJSON)
|
||||
if err != nil {
|
||||
return service.completeImmediateToolResult(stream, invocation, summarizeUpdateTodosResult(&agentv1.UpdateTodosResult{
|
||||
Result: &agentv1.UpdateTodosResult_Error{
|
||||
Error: &agentv1.UpdateTodosError{Error: err.Error()},
|
||||
},
|
||||
}), buildUpdateTodosToolCall(&agentv1.UpdateTodosArgs{}, &agentv1.UpdateTodosResult{
|
||||
Result: &agentv1.UpdateTodosResult_Error{
|
||||
Error: &agentv1.UpdateTodosError{Error: err.Error()},
|
||||
},
|
||||
}))
|
||||
}
|
||||
args := decodedArgs.Args
|
||||
conversation, _, _, err := service.snapshotCheckpointConversation(stream)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
structuredState, err := projectConversationStructuredState(conversation)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
now := time.Now().UTC()
|
||||
applyMerge := shouldMergeTodoUpdate(args, decodedArgs.MergeSet, structuredState.Todos)
|
||||
args.Merge = applyMerge
|
||||
var nextTodos []*agentv1.TodoItem
|
||||
if applyMerge {
|
||||
nextTodos, err = mergeTodoItems(structuredState.Todos, args.GetTodos(), now)
|
||||
} else {
|
||||
nextTodos, err = normalizeTodoItems(args.GetTodos(), now, true)
|
||||
}
|
||||
if err != nil {
|
||||
result := buildUpdateTodosErrorResult(args, err.Error())
|
||||
return service.completeImmediateToolResult(stream, invocation, summarizeUpdateTodosResult(result), buildUpdateTodosToolCall(args, result))
|
||||
}
|
||||
if !applyMerge {
|
||||
if missingIDs := missingActiveTodoReplacementIDs(structuredState.Todos, nextTodos); len(missingIDs) > 0 {
|
||||
result := buildUpdateTodosErrorResult(args, unsafeTodoReplaceError(missingIDs))
|
||||
return service.completeImmediateToolResult(stream, invocation, summarizeUpdateTodosResult(result), buildUpdateTodosToolCall(args, result))
|
||||
}
|
||||
}
|
||||
result := &agentv1.UpdateTodosResult{
|
||||
Result: &agentv1.UpdateTodosResult_Success{
|
||||
Success: &agentv1.UpdateTodosSuccess{
|
||||
Todos: cloneTodoItems(nextTodos),
|
||||
TotalCount: int32(len(nextTodos)),
|
||||
WasMerge: applyMerge,
|
||||
},
|
||||
},
|
||||
}
|
||||
return service.completeImmediateToolResult(stream, invocation, summarizeUpdateTodosResult(result), buildUpdateTodosToolCall(args, result))
|
||||
}
|
||||
|
||||
func shouldMergeTodoUpdate(args *agentv1.UpdateTodosArgs, mergeSet bool, existing []*agentv1.TodoItem) bool {
|
||||
if args.GetMerge() {
|
||||
return true
|
||||
}
|
||||
if mergeSet || len(existing) == 0 {
|
||||
return false
|
||||
}
|
||||
if len(missingActiveTodoReplacementIDs(existing, args.GetTodos())) > 0 {
|
||||
return true
|
||||
}
|
||||
for _, item := range args.GetTodos() {
|
||||
if item == nil {
|
||||
continue
|
||||
}
|
||||
if strings.TrimSpace(item.GetContent()) == "" {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func (service *Service) handleReadTodosToolInvocation(stream *ActiveStream, invocation runtimecore.ToolInvocation) error {
|
||||
args, err := decodeReadTodosArgsJSON(invocation.ArgsJSON)
|
||||
if err != nil {
|
||||
return service.completeImmediateToolResult(stream, invocation, summarizeReadTodosResult(&agentv1.ReadTodosResult{
|
||||
Result: &agentv1.ReadTodosResult_Error{
|
||||
Error: &agentv1.ReadTodosError{Error: err.Error()},
|
||||
},
|
||||
}), buildReadTodosToolCall(&agentv1.ReadTodosArgs{}, &agentv1.ReadTodosResult{
|
||||
Result: &agentv1.ReadTodosResult_Error{
|
||||
Error: &agentv1.ReadTodosError{Error: err.Error()},
|
||||
},
|
||||
}))
|
||||
}
|
||||
conversation, _, _, err := service.snapshotCheckpointConversation(stream)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
structuredState, err := projectConversationStructuredState(conversation)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
filtered := filterTodoItems(structuredState.Todos, args.GetStatusFilter(), args.GetIdFilter())
|
||||
result := &agentv1.ReadTodosResult{
|
||||
Result: &agentv1.ReadTodosResult_Success{
|
||||
Success: &agentv1.ReadTodosSuccess{
|
||||
Todos: filtered,
|
||||
TotalCount: int32(len(filtered)),
|
||||
},
|
||||
},
|
||||
}
|
||||
return service.completeImmediateToolResult(stream, invocation, summarizeReadTodosResult(result), buildReadTodosToolCall(args, result))
|
||||
}
|
||||
|
||||
func (service *Service) completeImmediateToolResult(stream *ActiveStream, invocation runtimecore.ToolInvocation, resultText string, toolCall *agentv1.ToolCall) error {
|
||||
if err := service.appendToolResult(stream, invocation.CallID, strings.TrimSpace(invocation.ToolName), invocation.ArgsJSON, resultText, invocation.ReasoningContent, toolCall); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := service.publishToolCallCompleted(stream.RequestID, invocation.CallID, invocation.ModelCallID, toolCall); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := service.syncSummaryCarryForward(stream.ConversationID, stream.RequestID, invocation.ModelCallID); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := service.publishCheckpoint(stream.RequestID, stream.ConversationID); err != nil {
|
||||
return err
|
||||
}
|
||||
return service.reconcileStream(stream)
|
||||
}
|
||||
@@ -0,0 +1,42 @@
|
||||
package forwarder
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"strings"
|
||||
)
|
||||
|
||||
var errProviderLoopInterrupted = errors.New("provider loop interrupted")
|
||||
|
||||
func isTerminalStreamStatus(status StreamStatus) bool {
|
||||
switch status {
|
||||
case StreamStatusCanceled, StreamStatusCompleted, StreamStatusFailed:
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func providerLoopInterruptErr(ctx context.Context, stream *ActiveStream, modelCallID string) error {
|
||||
if ctx != nil && ctx.Err() != nil {
|
||||
return errProviderLoopInterrupted
|
||||
}
|
||||
if stream == nil {
|
||||
return nil
|
||||
}
|
||||
stream.mu.Lock()
|
||||
defer stream.mu.Unlock()
|
||||
if isTerminalStreamStatus(stream.Status) {
|
||||
return errProviderLoopInterrupted
|
||||
}
|
||||
switch stream.Phase {
|
||||
case TurnPhaseCanceled, TurnPhaseCompleted, TurnPhaseFailed:
|
||||
return errProviderLoopInterrupted
|
||||
}
|
||||
expectedModelCallID := strings.TrimSpace(modelCallID)
|
||||
currentModelCallID := strings.TrimSpace(stream.CurrentModelCallID)
|
||||
if expectedModelCallID != "" && currentModelCallID != "" && currentModelCallID != expectedModelCallID {
|
||||
return errProviderLoopInterrupted
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,72 @@
|
||||
package forwarder
|
||||
|
||||
import (
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
modeladapter "cursor/internal/backend/agent/model"
|
||||
)
|
||||
|
||||
func cloneStringAnyMap(value map[string]any) map[string]any {
|
||||
if len(value) == 0 {
|
||||
return nil
|
||||
}
|
||||
payload, err := json.Marshal(value)
|
||||
if err != nil {
|
||||
return nil
|
||||
}
|
||||
var decoded map[string]any
|
||||
if err := json.Unmarshal(payload, &decoded); err != nil {
|
||||
return nil
|
||||
}
|
||||
return decoded
|
||||
}
|
||||
|
||||
func parseRFC3339Time(value string) time.Time {
|
||||
trimmed := strings.TrimSpace(value)
|
||||
if trimmed == "" {
|
||||
return time.Time{}
|
||||
}
|
||||
parsed, err := time.Parse(time.RFC3339Nano, trimmed)
|
||||
if err != nil {
|
||||
return time.Time{}
|
||||
}
|
||||
return parsed.UTC()
|
||||
}
|
||||
|
||||
func marshalCarryForwardMessage(message modeladapter.Message) ([]byte, error) {
|
||||
payload := map[string]any{
|
||||
"role": message.Role,
|
||||
"content": message.Content,
|
||||
}
|
||||
if len(message.ContentParts) > 0 {
|
||||
payload["content_parts"] = message.ContentParts
|
||||
}
|
||||
if strings.TrimSpace(message.ReasoningContent) != "" || (strings.TrimSpace(message.Role) == "assistant" && len(message.ToolCalls) > 0) {
|
||||
payload["reasoning_content"] = message.ReasoningContent
|
||||
}
|
||||
if strings.TrimSpace(message.ReasoningSignature) != "" {
|
||||
payload["reasoning_signature"] = message.ReasoningSignature
|
||||
}
|
||||
if strings.TrimSpace(message.ReasoningSignatureSource) != "" {
|
||||
payload["reasoning_signature_source"] = strings.TrimSpace(message.ReasoningSignatureSource)
|
||||
}
|
||||
if len(message.ToolCalls) > 0 {
|
||||
payload["tool_calls"] = message.ToolCalls
|
||||
}
|
||||
if strings.TrimSpace(message.ToolCallID) != "" {
|
||||
payload["tool_call_id"] = message.ToolCallID
|
||||
}
|
||||
if strings.TrimSpace(message.Name) != "" {
|
||||
payload["name"] = message.Name
|
||||
}
|
||||
return json.Marshal(payload)
|
||||
}
|
||||
|
||||
func shortSHA256(value string) string {
|
||||
sum := sha256.Sum256([]byte(value))
|
||||
return hex.EncodeToString(sum[:])[:12]
|
||||
}
|
||||
@@ -0,0 +1,73 @@
|
||||
// legacy_stream.go 提供兼容 Cursor legacy RunSSE 响应头的 HTTP 包装器。
|
||||
package forwarder
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
|
||||
"cursor/gen/agentv1"
|
||||
"cursor/gen/aiserverv1"
|
||||
|
||||
"connectrpc.com/connect"
|
||||
)
|
||||
|
||||
const legacyRunSSEContentType = "text/event-stream"
|
||||
|
||||
// NewLegacyRunSSEHandler 构造一个对外表现为 text/event-stream 的 Connect ServerStream 处理器。
|
||||
func NewLegacyRunSSEHandler(
|
||||
procedure string,
|
||||
implementation func(context.Context, *connect.Request[aiserverv1.BidiRequestId], *connect.ServerStream[agentv1.AgentServerMessage]) error,
|
||||
options ...connect.HandlerOption,
|
||||
) http.Handler {
|
||||
inner := connect.NewServerStreamHandler(procedure, implementation, options...)
|
||||
return http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {
|
||||
inner.ServeHTTP(newLegacyRunSSEHeaderWriter(writer), request)
|
||||
})
|
||||
}
|
||||
|
||||
// legacyRunSSEHeaderWriter 在真正写出 header 时覆盖 Content-Type。
|
||||
type legacyRunSSEHeaderWriter struct {
|
||||
http.ResponseWriter
|
||||
wroteHeader bool
|
||||
}
|
||||
|
||||
// newLegacyRunSSEHeaderWriter 创建一个响应头包装器。
|
||||
func newLegacyRunSSEHeaderWriter(writer http.ResponseWriter) *legacyRunSSEHeaderWriter {
|
||||
return &legacyRunSSEHeaderWriter{ResponseWriter: writer}
|
||||
}
|
||||
|
||||
// WriteHeader 在输出状态码前补齐 legacy 所需响应头。
|
||||
func (writer *legacyRunSSEHeaderWriter) WriteHeader(statusCode int) {
|
||||
writer.applyLegacyHeaders()
|
||||
writer.wroteHeader = true
|
||||
writer.ResponseWriter.WriteHeader(statusCode)
|
||||
}
|
||||
|
||||
// Write 在首次写 body 时懒加载 header。
|
||||
func (writer *legacyRunSSEHeaderWriter) Write(payload []byte) (int, error) {
|
||||
if !writer.wroteHeader {
|
||||
writer.WriteHeader(http.StatusOK)
|
||||
}
|
||||
return writer.ResponseWriter.Write(payload)
|
||||
}
|
||||
|
||||
// Flush 尝试把底层缓冲区立即刷新给客户端。
|
||||
func (writer *legacyRunSSEHeaderWriter) Flush() {
|
||||
if flusher, ok := writer.ResponseWriter.(http.Flusher); ok {
|
||||
flusher.Flush()
|
||||
}
|
||||
}
|
||||
|
||||
// Unwrap 返回底层 ResponseWriter。
|
||||
func (writer *legacyRunSSEHeaderWriter) Unwrap() http.ResponseWriter {
|
||||
return writer.ResponseWriter
|
||||
}
|
||||
|
||||
// applyLegacyHeaders 设置 Cursor legacy RunSSE 需要的响应头。
|
||||
func (writer *legacyRunSSEHeaderWriter) applyLegacyHeaders() {
|
||||
header := writer.ResponseWriter.Header()
|
||||
header.Set("Content-Type", legacyRunSSEContentType)
|
||||
if header.Get("Cache-Control") == "" {
|
||||
header.Set("Cache-Control", "no-cache")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,34 @@
|
||||
// module.go 负责把 forwarder service 装配成 legacy HTTP/Connect handler。
|
||||
package forwarder
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"connectrpc.com/connect"
|
||||
|
||||
modeladapter "cursor/internal/backend/agent/model"
|
||||
)
|
||||
|
||||
type Module struct {
|
||||
Service *Service
|
||||
LocalBidiHandler http.Handler
|
||||
LocalRunSSE http.Handler
|
||||
AiHandler http.Handler
|
||||
RepositoryServiceHandler http.Handler
|
||||
UploadServiceHandler http.Handler
|
||||
}
|
||||
|
||||
// NewModule 创建 forwarder 模块,并导出本地 Bidi / RunSSE 处理器。
|
||||
func NewModule(historyRoot string, channelService modeladapter.ChannelResolver) *Module {
|
||||
service := NewService(historyRoot, channelService)
|
||||
legacyBidiAppendProcedure := "/aiserver.v1.BidiService/BidiAppend"
|
||||
legacyRunSSEProcedure := "/agent.v1.AgentService/RunSSE"
|
||||
return &Module{
|
||||
Service: service,
|
||||
LocalBidiHandler: connect.NewUnaryHandler(legacyBidiAppendProcedure, service.BidiAppend),
|
||||
LocalRunSSE: NewLegacyRunSSEHandler(legacyRunSSEProcedure, service.RunSSE),
|
||||
AiHandler: newAIHandler(service),
|
||||
RepositoryServiceHandler: newRepositoryServiceHandler(service),
|
||||
UploadServiceHandler: newUploadServiceHandler(service),
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,745 @@
|
||||
package forwarder
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"log"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"google.golang.org/protobuf/encoding/protojson"
|
||||
|
||||
"cursor/gen/agentv1"
|
||||
execbridge "cursor/internal/backend/agent/bridge/exec"
|
||||
runtimecore "cursor/internal/backend/agent/core"
|
||||
)
|
||||
|
||||
const (
|
||||
patchEditToolName = "PatchEdit"
|
||||
|
||||
patchEditReadExecKindName = "patch_edit_read"
|
||||
patchEditWriteExecKindName = "patch_edit_write"
|
||||
patchEditPostReadExecKindName = "patch_edit_post_read"
|
||||
)
|
||||
|
||||
type patchEditArgs struct {
|
||||
Path string
|
||||
OldString string
|
||||
NewString string
|
||||
NewStringSet bool
|
||||
ReplaceAll bool
|
||||
ReplaceAllSet bool
|
||||
}
|
||||
|
||||
type pendingPatchEditPayload struct {
|
||||
ToolName string `json:"tool_name"`
|
||||
RawArgsJSON json.RawMessage `json:"raw_args_json,omitempty"`
|
||||
ArgsDecodeFailed bool `json:"args_decode_failed,omitempty"`
|
||||
Args patchEditArgs `json:"args,omitempty"`
|
||||
ResolvedPath string `json:"resolved_path"`
|
||||
BeforeContent string `json:"before_content,omitempty"`
|
||||
AfterContent string `json:"after_content,omitempty"`
|
||||
DiffString string `json:"diff_string,omitempty"`
|
||||
LinesAdded int32 `json:"lines_added,omitempty"`
|
||||
LinesRemoved int32 `json:"lines_removed,omitempty"`
|
||||
Message string `json:"message,omitempty"`
|
||||
}
|
||||
|
||||
type queuedPatchEditOperation struct {
|
||||
ToolCallID string
|
||||
ModelCallID string
|
||||
ProviderPass int
|
||||
ReasoningContent string
|
||||
ReasoningSignature string
|
||||
ReasoningSignatureSource string
|
||||
Payload pendingPatchEditPayload
|
||||
}
|
||||
|
||||
func isPatchEditToolName(name string) bool {
|
||||
return strings.TrimSpace(name) == patchEditToolName
|
||||
}
|
||||
|
||||
func isHiddenPatchEditExecKind(kind string) bool {
|
||||
switch strings.TrimSpace(kind) {
|
||||
case patchEditReadExecKindName, patchEditWriteExecKindName, patchEditPostReadExecKindName:
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func (service *Service) handlePatchEditToolInvocation(stream *ActiveStream, invocation runtimecore.ToolInvocation) error {
|
||||
if stream == nil {
|
||||
return fmt.Errorf("patch edit stream is required")
|
||||
}
|
||||
toolName := strings.TrimSpace(invocation.ToolName)
|
||||
if !isPatchEditToolName(toolName) {
|
||||
return fmt.Errorf("unsupported patch edit tool: %s", invocation.ToolName)
|
||||
}
|
||||
|
||||
payload := pendingPatchEditPayload{
|
||||
ToolName: patchEditToolName,
|
||||
RawArgsJSON: append(json.RawMessage(nil), invocation.ArgsJSON...),
|
||||
}
|
||||
args, decodeErr := decodePatchEditArgs(invocation.ArgsJSON)
|
||||
payload.Args = args
|
||||
payload.ArgsDecodeFailed = decodeErr != nil
|
||||
|
||||
resolvedPath := strings.TrimSpace(args.Path)
|
||||
payload.ResolvedPath = resolvedPath
|
||||
if decodeErr == nil {
|
||||
payload.Args.Path = resolvedPath
|
||||
visibleArgsJSON, err := payload.Args.MarshalJSON()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
invocation.ArgsJSON = visibleArgsJSON
|
||||
}
|
||||
|
||||
startedToolCall := buildStartedToolCall(invocation)
|
||||
if startedToolCall != nil {
|
||||
toolCallPayload, err := protojson.Marshal(startedToolCall)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
_, err = service.appendConversationEntries(stream, stream.ConversationID, []HistoryEntry{
|
||||
newToolCallEntryWithProviderMetadata(stream.TurnSeq, stream.RequestID, invocation.CallID, patchEditToolName, invocation.ReasoningContent, invocation.ReasoningSignature, invocation.ReasoningSignatureSource, invocation.ReasoningProviderItemID, invocation.ReasoningProviderStatus, invocation.ReasoningProviderSummary, invocation.ProviderItemID, invocation.ProviderCallID, invocation.ProviderStatus, toolCallPayload),
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
if err := service.broker.Publish(stream.RequestID, StreamEvent{
|
||||
Message: buildToolCallStartedMessage(invocation.CallID, invocation.ModelCallID, startedToolCall),
|
||||
}); err != nil {
|
||||
return err
|
||||
}
|
||||
providerPass := currentProviderPass(stream)
|
||||
if decodeErr != nil {
|
||||
return service.finishPatchEditOperation(stream, invocation.CallID, invocation.ModelCallID, providerPass, invocation.ReasoningContent, payload, buildEditErrorResult(resolvedPath, decodeErr.Error()))
|
||||
}
|
||||
return service.startOrQueuePatchEditOperation(stream, queuedPatchEditOperation{
|
||||
ToolCallID: invocation.CallID,
|
||||
ModelCallID: invocation.ModelCallID,
|
||||
ProviderPass: providerPass,
|
||||
ReasoningContent: invocation.ReasoningContent,
|
||||
ReasoningSignature: invocation.ReasoningSignature,
|
||||
ReasoningSignatureSource: invocation.ReasoningSignatureSource,
|
||||
Payload: payload,
|
||||
})
|
||||
}
|
||||
|
||||
func (service *Service) startOrQueuePatchEditOperation(stream *ActiveStream, operation queuedPatchEditOperation) error {
|
||||
key := patchEditQueueKey(operation.Payload.ResolvedPath)
|
||||
if key != "" && patchEditPathHasInFlightOperation(stream, key) {
|
||||
enqueuePatchEditOperation(stream, key, operation)
|
||||
log.Printf(
|
||||
"forwarder patch edit queued behind in-flight edit conversation_id=%s request_id=%s tool_call_id=%s path=%s",
|
||||
strings.TrimSpace(stream.ConversationID),
|
||||
strings.TrimSpace(stream.RequestID),
|
||||
strings.TrimSpace(operation.ToolCallID),
|
||||
key,
|
||||
)
|
||||
return service.publishCheckpoint(stream.RequestID, stream.ConversationID)
|
||||
}
|
||||
return service.startQueuedPatchEditOperation(stream, operation)
|
||||
}
|
||||
|
||||
func (service *Service) startQueuedPatchEditOperation(stream *ActiveStream, operation queuedPatchEditOperation) error {
|
||||
return service.startHiddenPatchEditRead(
|
||||
stream,
|
||||
operation.ToolCallID,
|
||||
operation.ModelCallID,
|
||||
operation.ProviderPass,
|
||||
operation.ReasoningContent,
|
||||
operation.ReasoningSignature,
|
||||
operation.ReasoningSignatureSource,
|
||||
operation.Payload,
|
||||
)
|
||||
}
|
||||
|
||||
func (service *Service) startHiddenPatchEditRead(stream *ActiveStream, toolCallID string, modelCallID string, providerPass int, reasoningContent string, reasoningSignature string, reasoningSignatureSource string, payload pendingPatchEditPayload) error {
|
||||
readArgsJSON, err := json.Marshal(map[string]any{"path": strings.TrimSpace(payload.ResolvedPath)})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
serverMessage, pendingExec, err := service.execBridge.OpenExec(execbridge.OpenExecContext{
|
||||
ConversationID: stream.ConversationID,
|
||||
}, runtimecore.ToolInvocation{
|
||||
CallID: strings.TrimSpace(toolCallID),
|
||||
ToolName: "Read",
|
||||
ArgsJSON: readArgsJSON,
|
||||
ModelCallID: strings.TrimSpace(modelCallID),
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
pendingArgsJSON, err := payload.MarshalJSON()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
pendingExec.ModelCallID = strings.TrimSpace(modelCallID)
|
||||
pendingExec.ProviderPass = providerPass
|
||||
pendingExec.ToolCallID = strings.TrimSpace(toolCallID)
|
||||
pendingExec.ReasoningContent = reasoningContent
|
||||
pendingExec.ReasoningSignature = strings.TrimSpace(reasoningSignature)
|
||||
pendingExec.ReasoningSignatureSource = strings.TrimSpace(reasoningSignatureSource)
|
||||
pendingExec.ExecKind = patchEditReadExecKindName
|
||||
pendingExec.ArgsJSON = pendingArgsJSON
|
||||
stream.mu.Lock()
|
||||
stream.PendingExecs[pendingExec.ExecID] = pendingExec
|
||||
stream.mu.Unlock()
|
||||
if err := service.publishCheckpoint(stream.RequestID, stream.ConversationID); err != nil {
|
||||
return err
|
||||
}
|
||||
return service.broker.Publish(stream.RequestID, StreamEvent{Message: serverMessage})
|
||||
}
|
||||
|
||||
func (service *Service) startHiddenPatchEditWrite(stream *ActiveStream, toolCallID string, modelCallID string, providerPass int, reasoningContent string, reasoningSignature string, reasoningSignatureSource string, payload pendingPatchEditPayload, beforeContent string) error {
|
||||
computation, err := computePatchEdit(payload, beforeContent)
|
||||
if err != nil {
|
||||
return service.finishPatchEditOperation(stream, strings.TrimSpace(toolCallID), strings.TrimSpace(modelCallID), providerPass, reasoningContent, payload, buildEditErrorResult(payload.ResolvedPath, err.Error()))
|
||||
}
|
||||
|
||||
transportContent := preparePatchEditWriteContentsForClient(computation.AfterContent)
|
||||
writeArgsJSON, err := json.Marshal(map[string]any{
|
||||
"path": strings.TrimSpace(payload.ResolvedPath),
|
||||
"contents": transportContent,
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
serverMessage, pendingExec, err := service.execBridge.OpenExec(execbridge.OpenExecContext{
|
||||
ConversationID: stream.ConversationID,
|
||||
}, runtimecore.ToolInvocation{
|
||||
CallID: strings.TrimSpace(toolCallID),
|
||||
ToolName: "Write",
|
||||
ArgsJSON: writeArgsJSON,
|
||||
ModelCallID: strings.TrimSpace(modelCallID),
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
payload.BeforeContent = computation.BeforeContent
|
||||
payload.AfterContent = computation.AfterContent
|
||||
payload.DiffString = computation.DiffString
|
||||
payload.LinesAdded = computation.LinesAdded
|
||||
payload.LinesRemoved = computation.LinesRemoved
|
||||
payload.Message = computation.Message
|
||||
pendingArgsJSON, err := payload.MarshalJSON()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
pendingExec.ModelCallID = strings.TrimSpace(modelCallID)
|
||||
pendingExec.ProviderPass = providerPass
|
||||
pendingExec.ToolCallID = strings.TrimSpace(toolCallID)
|
||||
pendingExec.ReasoningContent = reasoningContent
|
||||
pendingExec.ReasoningSignature = strings.TrimSpace(reasoningSignature)
|
||||
pendingExec.ReasoningSignatureSource = strings.TrimSpace(reasoningSignatureSource)
|
||||
pendingExec.ExecKind = patchEditWriteExecKindName
|
||||
pendingExec.ArgsJSON = pendingArgsJSON
|
||||
stream.mu.Lock()
|
||||
stream.PendingExecs[pendingExec.ExecID] = pendingExec
|
||||
stream.mu.Unlock()
|
||||
if err := service.publishCheckpoint(stream.RequestID, stream.ConversationID); err != nil {
|
||||
return err
|
||||
}
|
||||
return service.broker.Publish(stream.RequestID, StreamEvent{Message: serverMessage})
|
||||
}
|
||||
|
||||
func (service *Service) startHiddenPatchEditPostRead(stream *ActiveStream, toolCallID string, modelCallID string, providerPass int, reasoningContent string, reasoningSignature string, reasoningSignatureSource string, payload pendingPatchEditPayload) error {
|
||||
readArgsJSON, err := json.Marshal(map[string]any{"path": strings.TrimSpace(payload.ResolvedPath)})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
serverMessage, pendingExec, err := service.execBridge.OpenExec(execbridge.OpenExecContext{
|
||||
ConversationID: stream.ConversationID,
|
||||
}, runtimecore.ToolInvocation{
|
||||
CallID: strings.TrimSpace(toolCallID),
|
||||
ToolName: "Read",
|
||||
ArgsJSON: readArgsJSON,
|
||||
ModelCallID: strings.TrimSpace(modelCallID),
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
pendingArgsJSON, err := payload.MarshalJSON()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
pendingExec.ModelCallID = strings.TrimSpace(modelCallID)
|
||||
pendingExec.ProviderPass = providerPass
|
||||
pendingExec.ToolCallID = strings.TrimSpace(toolCallID)
|
||||
pendingExec.ReasoningContent = reasoningContent
|
||||
pendingExec.ReasoningSignature = strings.TrimSpace(reasoningSignature)
|
||||
pendingExec.ReasoningSignatureSource = strings.TrimSpace(reasoningSignatureSource)
|
||||
pendingExec.ExecKind = patchEditPostReadExecKindName
|
||||
pendingExec.ArgsJSON = pendingArgsJSON
|
||||
stream.mu.Lock()
|
||||
stream.PendingExecs[pendingExec.ExecID] = pendingExec
|
||||
stream.mu.Unlock()
|
||||
if err := service.publishCheckpoint(stream.RequestID, stream.ConversationID); err != nil {
|
||||
return err
|
||||
}
|
||||
return service.broker.Publish(stream.RequestID, StreamEvent{Message: serverMessage})
|
||||
}
|
||||
|
||||
func (service *Service) handleHiddenPatchEditExecResult(stream *ActiveStream, pending runtimecore.PendingExec, message *agentv1.ExecClientMessage) error {
|
||||
if stream == nil || message == nil {
|
||||
return fmt.Errorf("patch edit exec result requires stream and message")
|
||||
}
|
||||
payload, err := decodePendingPatchEditPayload(pending.ArgsJSON)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
switch strings.TrimSpace(pending.ExecKind) {
|
||||
case patchEditReadExecKindName:
|
||||
markExecCompleted(stream, pending)
|
||||
readResult := message.GetReadResult()
|
||||
if readResult == nil {
|
||||
return service.finishPatchEditOperation(stream, pending.ToolCallID, pending.ModelCallID, pending.ProviderPass, pending.ReasoningContent, payload, buildEditErrorResult(payload.ResolvedPath, "patch edit read result missing"))
|
||||
}
|
||||
content, ok, readErr := extractCompleteReadContentForPatchEdit(readResult)
|
||||
if !ok {
|
||||
if readErr != "" {
|
||||
return service.finishPatchEditOperation(stream, pending.ToolCallID, pending.ModelCallID, pending.ProviderPass, pending.ReasoningContent, payload, buildEditErrorResult(payload.ResolvedPath, readErr))
|
||||
}
|
||||
return service.finishPatchEditOperation(stream, pending.ToolCallID, pending.ModelCallID, pending.ProviderPass, pending.ReasoningContent, payload, buildEditResultFromReadResult(payload.ResolvedPath, readResult))
|
||||
}
|
||||
return service.startHiddenPatchEditWrite(stream, pending.ToolCallID, pending.ModelCallID, pending.ProviderPass, pending.ReasoningContent, pending.ReasoningSignature, pending.ReasoningSignatureSource, payload, content)
|
||||
case patchEditWriteExecKindName:
|
||||
markExecCompleted(stream, pending)
|
||||
writeResult := message.GetWriteResult()
|
||||
if writeResult == nil {
|
||||
return service.finishPatchEditOperation(stream, pending.ToolCallID, pending.ModelCallID, pending.ProviderPass, pending.ReasoningContent, payload, buildEditErrorResult(payload.ResolvedPath, "patch edit write result missing"))
|
||||
}
|
||||
switch item := writeResult.GetResult().(type) {
|
||||
case *agentv1.WriteResult_Success:
|
||||
payload.ResolvedPath = firstNonEmpty(strings.TrimSpace(item.Success.GetPath()), strings.TrimSpace(payload.ResolvedPath), strings.TrimSpace(patchEditPayloadPath(payload)))
|
||||
if item.Success.GetFileContentAfterWrite() != "" {
|
||||
finalAfterContent, reconciled, ok := service.reconcilePatchEditObservedContent(stream, pending.ToolCallID, payload.ResolvedPath, payload.AfterContent, item.Success.GetFileContentAfterWrite())
|
||||
if !ok {
|
||||
return service.finishPatchEditOperation(stream, pending.ToolCallID, pending.ModelCallID, pending.ProviderPass, pending.ReasoningContent, payload, buildEditErrorResult(payload.ResolvedPath, "patch edit write verification failed: observed file content differs from expected content"))
|
||||
}
|
||||
if finalAfterContent != payload.AfterContent {
|
||||
payload.DiffString, payload.LinesAdded, payload.LinesRemoved = computeEditDiff(payload.BeforeContent, finalAfterContent)
|
||||
}
|
||||
payload.AfterContent = finalAfterContent
|
||||
if reconciled {
|
||||
payload.Message = appendPatchEditMessage(payload.Message, "write result matched after client line-ending normalization")
|
||||
}
|
||||
}
|
||||
setPatchEditPayloadPath(&payload, payload.ResolvedPath)
|
||||
return service.startHiddenPatchEditPostRead(stream, pending.ToolCallID, pending.ModelCallID, pending.ProviderPass, pending.ReasoningContent, pending.ReasoningSignature, pending.ReasoningSignatureSource, payload)
|
||||
default:
|
||||
return service.finishPatchEditOperation(stream, pending.ToolCallID, pending.ModelCallID, pending.ProviderPass, pending.ReasoningContent, payload, buildEditResultFromWriteResult(payload.ResolvedPath, writeResult))
|
||||
}
|
||||
case patchEditPostReadExecKindName:
|
||||
markExecCompleted(stream, pending)
|
||||
readResult := message.GetReadResult()
|
||||
if readResult == nil {
|
||||
return service.finishPatchEditOperation(stream, pending.ToolCallID, pending.ModelCallID, pending.ProviderPass, pending.ReasoningContent, payload, buildFinalEditSuccessResult(payload.ResolvedPath, payload.AfterContent, patchEditPayloadAsEditPayload(payload)))
|
||||
}
|
||||
if content, ok := extractReadContentForEdit(readResult); ok {
|
||||
if success := readResult.GetSuccess(); success != nil {
|
||||
payload.ResolvedPath = firstNonEmpty(strings.TrimSpace(success.GetPath()), payload.ResolvedPath)
|
||||
setPatchEditPayloadPath(&payload, payload.ResolvedPath)
|
||||
}
|
||||
finalAfterContent, reconciled, ok := service.reconcilePatchEditObservedContent(stream, pending.ToolCallID, payload.ResolvedPath, payload.AfterContent, content)
|
||||
if !ok {
|
||||
return service.finishPatchEditOperation(stream, pending.ToolCallID, pending.ModelCallID, pending.ProviderPass, pending.ReasoningContent, payload, buildEditErrorResult(payload.ResolvedPath, "patch edit post-read verification failed: observed file content differs from expected content"))
|
||||
}
|
||||
if reconciled {
|
||||
payload.Message = appendPatchEditMessage(payload.Message, "post-read matched after client line-ending normalization")
|
||||
}
|
||||
return service.finishPatchEditOperation(stream, pending.ToolCallID, pending.ModelCallID, pending.ProviderPass, pending.ReasoningContent, payload, buildFinalEditSuccessResult(payload.ResolvedPath, finalAfterContent, patchEditPayloadAsEditPayload(payload)))
|
||||
}
|
||||
return service.finishPatchEditOperation(stream, pending.ToolCallID, pending.ModelCallID, pending.ProviderPass, pending.ReasoningContent, payload, buildFinalEditSuccessResult(payload.ResolvedPath, payload.AfterContent, patchEditPayloadAsEditPayload(payload)))
|
||||
default:
|
||||
return fmt.Errorf("unsupported hidden patch edit exec kind: %s", pending.ExecKind)
|
||||
}
|
||||
}
|
||||
|
||||
func (service *Service) handleHiddenPatchEditExecControl(stream *ActiveStream, pending runtimecore.PendingExec, message *agentv1.ExecClientControlMessage) error {
|
||||
if stream == nil || message == nil {
|
||||
return fmt.Errorf("patch edit exec control requires stream and message")
|
||||
}
|
||||
if _, ok := message.GetMessage().(*agentv1.ExecClientControlMessage_Heartbeat); ok {
|
||||
return nil
|
||||
}
|
||||
if _, ok := message.GetMessage().(*agentv1.ExecClientControlMessage_StreamClose); ok {
|
||||
return nil
|
||||
}
|
||||
payload, err := decodePendingPatchEditPayload(pending.ArgsJSON)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
markExecCompleted(stream, pending)
|
||||
if strings.TrimSpace(pending.ExecKind) == patchEditPostReadExecKindName {
|
||||
return service.finishPatchEditOperation(stream, pending.ToolCallID, pending.ModelCallID, pending.ProviderPass, pending.ReasoningContent, payload, buildFinalEditSuccessResult(payload.ResolvedPath, payload.AfterContent, patchEditPayloadAsEditPayload(payload)))
|
||||
}
|
||||
return service.finishPatchEditOperation(stream, pending.ToolCallID, pending.ModelCallID, pending.ProviderPass, pending.ReasoningContent, payload, buildEditErrorResult(payload.ResolvedPath, hiddenPatchEditControlError(message)))
|
||||
}
|
||||
|
||||
func (service *Service) finishPatchEditOperation(stream *ActiveStream, toolCallID string, modelCallID string, providerPass int, reasoningContent string, payload pendingPatchEditPayload, result *agentv1.EditResult) error {
|
||||
if stream == nil {
|
||||
return nil
|
||||
}
|
||||
if result == nil {
|
||||
result = buildEditErrorResult(payload.ResolvedPath, "patch edit result missing")
|
||||
}
|
||||
path := firstNonEmpty(resultPath(result), payload.ResolvedPath, patchEditPayloadPath(payload))
|
||||
setPatchEditPayloadPath(&payload, path)
|
||||
argsJSON, err := patchEditPayloadArgsJSON(payload)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
toolCall := buildCompletedEditToolCall(path, result)
|
||||
historyToolCall := buildCompletedEditToolCall(path, compactPatchEditHistoryEditResult(path, result))
|
||||
if err := service.appendToolResult(stream, strings.TrimSpace(toolCallID), patchEditToolName, argsJSON, summarizePatchEditResult(path, result), reasoningContent, historyToolCall); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := service.publishToolCallCompleted(stream.RequestID, strings.TrimSpace(toolCallID), strings.TrimSpace(modelCallID), toolCall); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := service.syncSummaryCarryForward(stream.ConversationID, stream.RequestID, strings.TrimSpace(modelCallID)); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := service.publishCheckpoint(stream.RequestID, stream.ConversationID); err != nil {
|
||||
return err
|
||||
}
|
||||
if next, ok := takeNextQueuedPatchEditOperation(stream, firstNonEmpty(payload.ResolvedPath, path)); ok {
|
||||
if err := service.startQueuedPatchEditOperation(stream, next); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return service.reconcileStream(stream)
|
||||
}
|
||||
|
||||
func patchEditQueueKey(path string) string {
|
||||
return strings.TrimSpace(path)
|
||||
}
|
||||
|
||||
func patchEditPathHasInFlightOperation(stream *ActiveStream, key string) bool {
|
||||
if stream == nil || strings.TrimSpace(key) == "" {
|
||||
return false
|
||||
}
|
||||
stream.mu.Lock()
|
||||
defer stream.mu.Unlock()
|
||||
for _, pending := range stream.PendingExecs {
|
||||
if !isHiddenPatchEditExecKind(pending.ExecKind) {
|
||||
continue
|
||||
}
|
||||
payload, err := decodePendingPatchEditPayload(pending.ArgsJSON)
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
if patchEditQueueKey(firstNonEmpty(payload.ResolvedPath, patchEditPayloadPath(payload))) == key {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func enqueuePatchEditOperation(stream *ActiveStream, key string, operation queuedPatchEditOperation) {
|
||||
if stream == nil || strings.TrimSpace(key) == "" {
|
||||
return
|
||||
}
|
||||
if len(operation.Payload.RawArgsJSON) > 0 {
|
||||
operation.Payload.RawArgsJSON = append(json.RawMessage(nil), operation.Payload.RawArgsJSON...)
|
||||
}
|
||||
stream.mu.Lock()
|
||||
if stream.PatchEditQueues == nil {
|
||||
stream.PatchEditQueues = make(map[string][]queuedPatchEditOperation)
|
||||
}
|
||||
stream.PatchEditQueues[key] = append(stream.PatchEditQueues[key], operation)
|
||||
stream.UpdatedAt = time.Now().UTC()
|
||||
stream.mu.Unlock()
|
||||
}
|
||||
|
||||
func takeNextQueuedPatchEditOperation(stream *ActiveStream, path string) (queuedPatchEditOperation, bool) {
|
||||
key := patchEditQueueKey(path)
|
||||
if stream == nil || key == "" {
|
||||
return queuedPatchEditOperation{}, false
|
||||
}
|
||||
stream.mu.Lock()
|
||||
defer stream.mu.Unlock()
|
||||
if stream.PatchEditQueues == nil {
|
||||
return queuedPatchEditOperation{}, false
|
||||
}
|
||||
queue := stream.PatchEditQueues[key]
|
||||
if len(queue) == 0 {
|
||||
delete(stream.PatchEditQueues, key)
|
||||
return queuedPatchEditOperation{}, false
|
||||
}
|
||||
next := queue[0]
|
||||
if len(queue) == 1 {
|
||||
delete(stream.PatchEditQueues, key)
|
||||
} else {
|
||||
remaining := make([]queuedPatchEditOperation, len(queue)-1)
|
||||
copy(remaining, queue[1:])
|
||||
stream.PatchEditQueues[key] = remaining
|
||||
}
|
||||
stream.UpdatedAt = time.Now().UTC()
|
||||
return next, true
|
||||
}
|
||||
|
||||
func decodePatchEditArgs(raw []byte) (patchEditArgs, error) {
|
||||
var rawArgs map[string]any
|
||||
if err := json.Unmarshal(raw, &rawArgs); err != nil {
|
||||
return patchEditArgs{}, fmt.Errorf("decode PatchEdit args failed: %w", err)
|
||||
}
|
||||
path, pathFound, pathValid := readJSONStringAny(rawArgs, "path", "file_path", "Path", "FilePath")
|
||||
oldString, oldFound, oldValid := readJSONStringAny(rawArgs, "old_string", "oldString", "OldString")
|
||||
newString, newFound, newValid := readJSONStringAny(rawArgs, "new_string", "newString", "NewString")
|
||||
replaceAll, replaceAllFound, replaceAllValid := readJSONBoolAny(rawArgs, "replace_all", "replaceAll", "ReplaceAll")
|
||||
result := patchEditArgs{
|
||||
Path: path,
|
||||
OldString: oldString,
|
||||
NewString: newString,
|
||||
NewStringSet: newFound,
|
||||
ReplaceAll: replaceAll,
|
||||
ReplaceAllSet: replaceAllFound,
|
||||
}
|
||||
switch {
|
||||
case pathFound && !pathValid:
|
||||
return result, fmt.Errorf("PatchEdit path must be a string")
|
||||
case strings.TrimSpace(result.Path) == "":
|
||||
return result, fmt.Errorf("PatchEdit path is required")
|
||||
case !isAbsoluteToolPath(result.Path):
|
||||
return result, fmt.Errorf("PatchEdit path must be absolute")
|
||||
case oldFound && !oldValid:
|
||||
return result, fmt.Errorf("PatchEdit old_string must be a string")
|
||||
case !oldFound || oldString == "":
|
||||
return result, fmt.Errorf("PatchEdit old_string is required")
|
||||
case newFound && !newValid:
|
||||
return result, fmt.Errorf("PatchEdit new_string must be a string")
|
||||
case !newFound:
|
||||
return result, fmt.Errorf("PatchEdit new_string is required")
|
||||
case replaceAllFound && !replaceAllValid:
|
||||
return result, fmt.Errorf("PatchEdit replace_all must be a boolean")
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func computePatchEdit(payload pendingPatchEditPayload, beforeContent string) (editComputation, error) {
|
||||
args := payload.Args
|
||||
occurrences := strings.Count(beforeContent, args.OldString)
|
||||
switch {
|
||||
case occurrences == 0:
|
||||
return editComputation{}, fmt.Errorf("PatchEdit old_string not found")
|
||||
case occurrences > 1 && !args.ReplaceAll:
|
||||
return editComputation{}, fmt.Errorf("PatchEdit old_string is not unique; found %d occurrences", occurrences)
|
||||
}
|
||||
replaceCount := 1
|
||||
if args.ReplaceAll {
|
||||
replaceCount = -1
|
||||
}
|
||||
afterContent := strings.Replace(beforeContent, args.OldString, args.NewString, replaceCount)
|
||||
diffString, linesAdded, linesRemoved := computeEditDiff(beforeContent, afterContent)
|
||||
return editComputation{
|
||||
BeforeContent: beforeContent,
|
||||
AfterContent: afterContent,
|
||||
DiffString: diffString,
|
||||
LinesAdded: linesAdded,
|
||||
LinesRemoved: linesRemoved,
|
||||
Message: "PatchEdit applied",
|
||||
}, nil
|
||||
}
|
||||
|
||||
func extractCompleteReadContentForPatchEdit(result *agentv1.ReadResult) (string, bool, string) {
|
||||
success := result.GetSuccess()
|
||||
if success == nil {
|
||||
return "", false, ""
|
||||
}
|
||||
if success.GetTruncated() {
|
||||
return "", false, "patch edit requires a complete Read result, but the file content was truncated"
|
||||
}
|
||||
switch output := success.GetOutput().(type) {
|
||||
case *agentv1.ReadSuccess_Content:
|
||||
return normalizeLineEndingsToLF(output.Content), true, ""
|
||||
case *agentv1.ReadSuccess_Data:
|
||||
text, ok, readErr := decodeReadDataAsEditableText(output.Data)
|
||||
if !ok {
|
||||
return "", false, "patch edit requires editable text content, but " + readErr
|
||||
}
|
||||
return normalizeLineEndingsToLF(text), true, ""
|
||||
}
|
||||
if success.GetOutputBlobId() != nil {
|
||||
return "", false, "patch edit requires inline file content, but Read returned only an output blob"
|
||||
}
|
||||
return "", false, "patch edit requires inline file content, but Read did not include content"
|
||||
}
|
||||
|
||||
func (service *Service) reconcilePatchEditObservedContent(stream *ActiveStream, toolCallID string, path string, expected string, observed string) (string, bool, bool) {
|
||||
if observed == expected {
|
||||
return observed, false, true
|
||||
}
|
||||
expectedNormalized := normalizeLineEndingsToLF(expected)
|
||||
observedNormalized := normalizeLineEndingsToLF(observed)
|
||||
lineEndingEquivalent := expectedNormalized == observedNormalized
|
||||
log.Printf(
|
||||
"forwarder patch edit observed content mismatch conversation_id=%s request_id=%s tool_call_id=%s path=%s expected_bytes=%d observed_bytes=%d expected_sha256=%s observed_sha256=%s line_ending_equivalent=%t",
|
||||
strings.TrimSpace(stream.ConversationID),
|
||||
strings.TrimSpace(stream.RequestID),
|
||||
strings.TrimSpace(toolCallID),
|
||||
strings.TrimSpace(path),
|
||||
len(expected),
|
||||
len(observed),
|
||||
shortSHA256(expected),
|
||||
shortSHA256(observed),
|
||||
lineEndingEquivalent,
|
||||
)
|
||||
if lineEndingEquivalent {
|
||||
return observed, true, true
|
||||
}
|
||||
return observed, false, false
|
||||
}
|
||||
|
||||
func appendPatchEditMessage(existing string, addition string) string {
|
||||
existing = strings.TrimSpace(existing)
|
||||
addition = strings.TrimSpace(addition)
|
||||
if existing == "" {
|
||||
return addition
|
||||
}
|
||||
if addition == "" {
|
||||
return existing
|
||||
}
|
||||
return existing + "; " + addition
|
||||
}
|
||||
|
||||
func patchEditPayloadPath(payload pendingPatchEditPayload) string {
|
||||
return strings.TrimSpace(payload.Args.Path)
|
||||
}
|
||||
|
||||
func setPatchEditPayloadPath(payload *pendingPatchEditPayload, path string) {
|
||||
if payload == nil {
|
||||
return
|
||||
}
|
||||
payload.Args.Path = strings.TrimSpace(path)
|
||||
}
|
||||
|
||||
func patchEditPayloadArgsJSON(payload pendingPatchEditPayload) ([]byte, error) {
|
||||
if payload.ArgsDecodeFailed && len(payload.RawArgsJSON) > 0 {
|
||||
return append([]byte(nil), payload.RawArgsJSON...), nil
|
||||
}
|
||||
return payload.Args.MarshalJSON()
|
||||
}
|
||||
|
||||
func patchEditPayloadAsEditPayload(payload pendingPatchEditPayload) editResultPayload {
|
||||
return editResultPayload{
|
||||
BeforeContent: payload.BeforeContent,
|
||||
AfterContent: payload.AfterContent,
|
||||
DiffString: payload.DiffString,
|
||||
LinesAdded: payload.LinesAdded,
|
||||
LinesRemoved: payload.LinesRemoved,
|
||||
Message: payload.Message,
|
||||
}
|
||||
}
|
||||
|
||||
func summarizePatchEditResult(path string, result *agentv1.EditResult) string {
|
||||
success := result.GetSuccess()
|
||||
if success == nil {
|
||||
return summarizeEditResultWithoutPath(result)
|
||||
}
|
||||
encoded, err := json.Marshal(map[string]any{
|
||||
"success": map[string]any{
|
||||
"path": firstNonEmpty(success.GetPath(), path),
|
||||
},
|
||||
})
|
||||
if err != nil {
|
||||
return fmt.Sprintf(`{"success":{"path":%q}}`, firstNonEmpty(success.GetPath(), path))
|
||||
}
|
||||
return string(encoded)
|
||||
}
|
||||
|
||||
func compactPatchEditHistoryEditResult(path string, result *agentv1.EditResult) *agentv1.EditResult {
|
||||
success := result.GetSuccess()
|
||||
if success == nil {
|
||||
return editResultWithoutPath(result)
|
||||
}
|
||||
return &agentv1.EditResult{
|
||||
Result: &agentv1.EditResult_Success{
|
||||
Success: &agentv1.EditSuccess{
|
||||
Path: firstNonEmpty(success.GetPath(), path),
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func boundedPatchEditDiffString(diffString string) string {
|
||||
return truncateProjectedReplayText(patchEditToolName, diffString, projectedPatchEditReplayLimit)
|
||||
}
|
||||
|
||||
func (args patchEditArgs) MarshalJSON() ([]byte, error) {
|
||||
payload := map[string]any{
|
||||
"path": strings.TrimSpace(args.Path),
|
||||
"old_string": args.OldString,
|
||||
"new_string": args.NewString,
|
||||
}
|
||||
if args.ReplaceAllSet || args.ReplaceAll {
|
||||
payload["replace_all"] = args.ReplaceAll
|
||||
}
|
||||
return json.Marshal(payload)
|
||||
}
|
||||
|
||||
func (args *patchEditArgs) UnmarshalJSON(raw []byte) error {
|
||||
var rawArgs map[string]any
|
||||
if err := json.Unmarshal(raw, &rawArgs); err != nil {
|
||||
return err
|
||||
}
|
||||
path, pathFound, pathValid := readJSONStringAny(rawArgs, "path", "file_path", "Path", "FilePath")
|
||||
if pathFound && !pathValid {
|
||||
return fmt.Errorf("PatchEdit path must be a string")
|
||||
}
|
||||
oldString, oldFound, oldValid := readJSONStringAny(rawArgs, "old_string", "oldString", "OldString")
|
||||
if oldFound && !oldValid {
|
||||
return fmt.Errorf("PatchEdit old_string must be a string")
|
||||
}
|
||||
newString, newFound, newValid := readJSONStringAny(rawArgs, "new_string", "newString", "NewString")
|
||||
if newFound && !newValid {
|
||||
return fmt.Errorf("PatchEdit new_string must be a string")
|
||||
}
|
||||
replaceAll, replaceAllFound, replaceAllValid := readJSONBoolAny(rawArgs, "replace_all", "replaceAll", "ReplaceAll")
|
||||
if replaceAllFound && !replaceAllValid {
|
||||
return fmt.Errorf("PatchEdit replace_all must be a boolean")
|
||||
}
|
||||
*args = patchEditArgs{
|
||||
Path: path,
|
||||
OldString: oldString,
|
||||
NewString: newString,
|
||||
NewStringSet: newFound,
|
||||
ReplaceAll: replaceAll,
|
||||
ReplaceAllSet: replaceAllFound,
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (payload pendingPatchEditPayload) MarshalJSON() ([]byte, error) {
|
||||
type alias pendingPatchEditPayload
|
||||
return json.Marshal(alias(payload))
|
||||
}
|
||||
|
||||
func decodePendingPatchEditPayload(raw []byte) (pendingPatchEditPayload, error) {
|
||||
var payload pendingPatchEditPayload
|
||||
if err := json.Unmarshal(raw, &payload); err != nil {
|
||||
return pendingPatchEditPayload{}, fmt.Errorf("decode patch edit pending payload failed: %w", err)
|
||||
}
|
||||
return payload, nil
|
||||
}
|
||||
|
||||
func hiddenPatchEditControlError(message *agentv1.ExecClientControlMessage) string {
|
||||
if message == nil {
|
||||
return "patch edit operation failed"
|
||||
}
|
||||
switch item := message.GetMessage().(type) {
|
||||
case *agentv1.ExecClientControlMessage_Throw:
|
||||
return firstNonEmpty(strings.TrimSpace(item.Throw.GetError()), "patch edit operation failed")
|
||||
case *agentv1.ExecClientControlMessage_StreamClose:
|
||||
return "patch edit operation closed unexpectedly"
|
||||
default:
|
||||
return "patch edit operation failed"
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,296 @@
|
||||
package forwarder
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"cursor/gen/agentv1"
|
||||
)
|
||||
|
||||
type streamPathContext struct {
|
||||
workspacePaths []string
|
||||
terminalsFolder string
|
||||
requestFileContents map[string]string
|
||||
}
|
||||
|
||||
func updateStreamRequestContextData(stream *ActiveStream, requestContext *agentv1.RequestContext) {
|
||||
if stream == nil {
|
||||
return
|
||||
}
|
||||
|
||||
var workspacePaths []string
|
||||
var terminalsFolder string
|
||||
var fileContents map[string]string
|
||||
if requestContext != nil {
|
||||
env := requestContext.GetEnv()
|
||||
workspacePaths = compactWorkspacePaths(env.GetWorkspacePaths(), env.GetProjectFolder())
|
||||
terminalsFolder = strings.TrimSpace(env.GetTerminalsFolder())
|
||||
fileContents = cloneStringMap(requestContext.GetFileContents())
|
||||
}
|
||||
|
||||
stream.mu.Lock()
|
||||
stream.WorkspacePaths = workspacePaths
|
||||
stream.TerminalsFolder = terminalsFolder
|
||||
stream.RequestFileContents = fileContents
|
||||
stream.UpdatedAt = time.Now().UTC()
|
||||
stream.mu.Unlock()
|
||||
}
|
||||
|
||||
func snapshotStreamPathContext(stream *ActiveStream) streamPathContext {
|
||||
if stream == nil {
|
||||
return streamPathContext{}
|
||||
}
|
||||
|
||||
stream.mu.Lock()
|
||||
defer stream.mu.Unlock()
|
||||
|
||||
return streamPathContext{
|
||||
workspacePaths: append([]string(nil), stream.WorkspacePaths...),
|
||||
terminalsFolder: strings.TrimSpace(stream.TerminalsFolder),
|
||||
requestFileContents: cloneStringMap(stream.RequestFileContents),
|
||||
}
|
||||
}
|
||||
|
||||
func compactWorkspacePaths(workspacePaths []string, projectFolder string) []string {
|
||||
items := append([]string(nil), workspacePaths...)
|
||||
if trimmedProjectFolder := strings.TrimSpace(projectFolder); trimmedProjectFolder != "" {
|
||||
items = append(items, trimmedProjectFolder)
|
||||
}
|
||||
seen := make(map[string]struct{}, len(items))
|
||||
result := make([]string, 0, len(items))
|
||||
for _, item := range items {
|
||||
trimmed := strings.TrimSpace(item)
|
||||
if trimmed == "" {
|
||||
continue
|
||||
}
|
||||
cleaned := filepath.Clean(trimmed)
|
||||
if _, ok := seen[cleaned]; ok {
|
||||
continue
|
||||
}
|
||||
seen[cleaned] = struct{}{}
|
||||
result = append(result, cleaned)
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
func cloneStringMap(input map[string]string) map[string]string {
|
||||
if len(input) == 0 {
|
||||
return nil
|
||||
}
|
||||
cloned := make(map[string]string, len(input))
|
||||
for key, value := range input {
|
||||
trimmedKey := strings.TrimSpace(key)
|
||||
if trimmedKey == "" {
|
||||
continue
|
||||
}
|
||||
cloned[trimmedKey] = value
|
||||
}
|
||||
if len(cloned) == 0 {
|
||||
return nil
|
||||
}
|
||||
return cloned
|
||||
}
|
||||
|
||||
func resolveWorkspacePath(path string, workspaceRoots []string, requireExisting bool) (string, bool) {
|
||||
trimmedPath := strings.TrimSpace(path)
|
||||
if trimmedPath == "" {
|
||||
return "", false
|
||||
}
|
||||
cleanedPath := filepath.Clean(trimmedPath)
|
||||
if pathCandidateUsable(cleanedPath, requireExisting) {
|
||||
return cleanedPath, true
|
||||
}
|
||||
|
||||
if len(workspaceRoots) == 0 {
|
||||
return "", false
|
||||
}
|
||||
|
||||
if !filepath.IsAbs(cleanedPath) {
|
||||
for _, workspaceRoot := range workspaceRoots {
|
||||
candidate := filepath.Join(workspaceRoot, cleanedPath)
|
||||
if pathCandidateUsable(candidate, requireExisting) {
|
||||
return filepath.Clean(candidate), true
|
||||
}
|
||||
}
|
||||
return "", false
|
||||
}
|
||||
|
||||
pathParts := splitPathParts(cleanedPath)
|
||||
if len(pathParts) == 0 {
|
||||
return "", false
|
||||
}
|
||||
|
||||
for _, workspaceRoot := range workspaceRoots {
|
||||
rootBase := strings.TrimSpace(filepath.Base(workspaceRoot))
|
||||
if rootBase == "" || rootBase == "." || rootBase == string(filepath.Separator) {
|
||||
continue
|
||||
}
|
||||
if index := lastIndexFold(pathParts, rootBase); index >= 0 {
|
||||
candidate := workspaceRoot
|
||||
if index+1 < len(pathParts) {
|
||||
candidate = filepath.Join(append([]string{workspaceRoot}, pathParts[index+1:]...)...)
|
||||
}
|
||||
if pathCandidateUsable(candidate, requireExisting) {
|
||||
return filepath.Clean(candidate), true
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
for suffixLen := len(pathParts); suffixLen >= 1; suffixLen-- {
|
||||
suffixParts := append([]string(nil), pathParts[len(pathParts)-suffixLen:]...)
|
||||
for _, workspaceRoot := range workspaceRoots {
|
||||
candidate := joinWorkspaceRootWithSuffix(workspaceRoot, suffixParts)
|
||||
if pathCandidateUsable(candidate, requireExisting) {
|
||||
return filepath.Clean(candidate), true
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return "", false
|
||||
}
|
||||
|
||||
func isAbsoluteToolPath(path string) bool {
|
||||
trimmed := strings.TrimSpace(path)
|
||||
if trimmed == "" {
|
||||
return false
|
||||
}
|
||||
if strings.HasPrefix(trimmed, "/") {
|
||||
return true
|
||||
}
|
||||
if strings.HasPrefix(trimmed, `\\`) {
|
||||
return true
|
||||
}
|
||||
return len(trimmed) >= 3 && isASCIIAlpha(trimmed[0]) && trimmed[1] == ':' && isPathSeparator(trimmed[2])
|
||||
}
|
||||
|
||||
func isASCIIAlpha(value byte) bool {
|
||||
return (value >= 'A' && value <= 'Z') || (value >= 'a' && value <= 'z')
|
||||
}
|
||||
|
||||
func isPathSeparator(value byte) bool {
|
||||
return value == '\\' || value == '/'
|
||||
}
|
||||
|
||||
func lookupRequestFileContents(context streamPathContext, originalPath string, resolvedPath string) (string, bool) {
|
||||
if len(context.requestFileContents) == 0 {
|
||||
return "", false
|
||||
}
|
||||
aliases := buildPathAliases(originalPath, resolvedPath, context.workspacePaths)
|
||||
for _, alias := range aliases {
|
||||
if value, ok := context.requestFileContents[alias]; ok {
|
||||
return value, true
|
||||
}
|
||||
}
|
||||
return "", false
|
||||
}
|
||||
|
||||
func buildPathAliases(originalPath string, resolvedPath string, workspaceRoots []string) []string {
|
||||
seen := make(map[string]struct{}, 16)
|
||||
aliases := make([]string, 0, 16)
|
||||
addAlias := func(value string) {
|
||||
trimmed := strings.TrimSpace(value)
|
||||
if trimmed == "" {
|
||||
return
|
||||
}
|
||||
if _, ok := seen[trimmed]; ok {
|
||||
return
|
||||
}
|
||||
seen[trimmed] = struct{}{}
|
||||
aliases = append(aliases, trimmed)
|
||||
}
|
||||
|
||||
for _, candidate := range []string{originalPath, resolvedPath} {
|
||||
trimmed := strings.TrimSpace(candidate)
|
||||
if trimmed == "" {
|
||||
continue
|
||||
}
|
||||
addAlias(trimmed)
|
||||
addAlias(filepath.Clean(trimmed))
|
||||
addAlias(filepath.ToSlash(filepath.Clean(trimmed)))
|
||||
}
|
||||
|
||||
for _, candidate := range []string{originalPath, resolvedPath} {
|
||||
cleaned := filepath.Clean(strings.TrimSpace(candidate))
|
||||
if cleaned == "" {
|
||||
continue
|
||||
}
|
||||
for _, workspaceRoot := range workspaceRoots {
|
||||
if rel, err := filepath.Rel(workspaceRoot, cleaned); err == nil && rel != "." && !strings.HasPrefix(rel, ".."+string(filepath.Separator)) && rel != ".." {
|
||||
addAlias(rel)
|
||||
addAlias(filepath.ToSlash(rel))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return aliases
|
||||
}
|
||||
|
||||
func splitPathParts(path string) []string {
|
||||
cleaned := filepath.Clean(strings.TrimSpace(path))
|
||||
if cleaned == "" {
|
||||
return nil
|
||||
}
|
||||
volumeName := filepath.VolumeName(cleaned)
|
||||
if volumeName != "" {
|
||||
cleaned = strings.TrimPrefix(cleaned, volumeName)
|
||||
}
|
||||
cleaned = strings.TrimPrefix(cleaned, string(filepath.Separator))
|
||||
if cleaned == "" {
|
||||
return nil
|
||||
}
|
||||
parts := strings.Split(cleaned, string(filepath.Separator))
|
||||
result := make([]string, 0, len(parts))
|
||||
for _, part := range parts {
|
||||
trimmed := strings.TrimSpace(part)
|
||||
if trimmed != "" {
|
||||
result = append(result, trimmed)
|
||||
}
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
func lastIndexFold(items []string, target string) int {
|
||||
for index := len(items) - 1; index >= 0; index-- {
|
||||
if strings.EqualFold(strings.TrimSpace(items[index]), strings.TrimSpace(target)) {
|
||||
return index
|
||||
}
|
||||
}
|
||||
return -1
|
||||
}
|
||||
|
||||
func joinWorkspaceRootWithSuffix(workspaceRoot string, suffixParts []string) string {
|
||||
root := filepath.Clean(strings.TrimSpace(workspaceRoot))
|
||||
if root == "" {
|
||||
return ""
|
||||
}
|
||||
parts := append([]string(nil), suffixParts...)
|
||||
if len(parts) > 0 && strings.EqualFold(parts[0], filepath.Base(root)) {
|
||||
parts = parts[1:]
|
||||
}
|
||||
if len(parts) == 0 {
|
||||
return root
|
||||
}
|
||||
return filepath.Join(append([]string{root}, parts...)...)
|
||||
}
|
||||
|
||||
func pathCandidateUsable(path string, requireExisting bool) bool {
|
||||
cleaned := filepath.Clean(strings.TrimSpace(path))
|
||||
if cleaned == "" {
|
||||
return false
|
||||
}
|
||||
info, err := os.Stat(cleaned)
|
||||
if err == nil {
|
||||
return !info.IsDir() || requireExisting
|
||||
}
|
||||
if requireExisting {
|
||||
return false
|
||||
}
|
||||
parent := filepath.Dir(cleaned)
|
||||
if parent == "" || parent == "." || parent == cleaned {
|
||||
return false
|
||||
}
|
||||
parentInfo, parentErr := os.Stat(parent)
|
||||
return parentErr == nil && parentInfo.IsDir()
|
||||
}
|
||||
@@ -0,0 +1,37 @@
|
||||
package forwarder
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestIsAbsoluteToolPath(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
path string
|
||||
want bool
|
||||
}{
|
||||
{name: "empty", path: "", want: false},
|
||||
{name: "spaces", path: " ", want: false},
|
||||
{name: "relative", path: "internal/backend/forwarder/path_resolution.go", want: false},
|
||||
{name: "dot relative", path: "./path_resolution.go", want: false},
|
||||
{name: "parent relative", path: "../path_resolution.go", want: false},
|
||||
{name: "windows drive relative", path: `C:path_resolution.go`, want: false},
|
||||
{name: "windows drive only", path: `C:`, want: false},
|
||||
{name: "invalid drive prefix", path: `1:/path_resolution.go`, want: false},
|
||||
{name: "posix absolute", path: "/Users/example/path_resolution.go", want: true},
|
||||
{name: "posix absolute with spaces", path: " /tmp/path_resolution.go ", want: true},
|
||||
{name: "windows backslash absolute", path: `C:\Users\example\path_resolution.go`, want: true},
|
||||
{name: "windows slash absolute", path: `C:/Users/example/path_resolution.go`, want: true},
|
||||
{name: "windows unc backslash", path: `\\server\share\path_resolution.go`, want: true},
|
||||
{name: "windows unc slash", path: `//server/share/path_resolution.go`, want: true},
|
||||
{name: "windows extended drive", path: `\\?\C:\Users\example\path_resolution.go`, want: true},
|
||||
{name: "windows extended unc", path: `\\?\UNC\server\share\path_resolution.go`, want: true},
|
||||
{name: "windows device path", path: `\\.\C:\Users\example\path_resolution.go`, want: true},
|
||||
}
|
||||
|
||||
for _, test := range tests {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
if got := isAbsoluteToolPath(test.path); got != test.want {
|
||||
t.Fatalf("isAbsoluteToolPath(%q) = %v, want %v", test.path, got, test.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,98 @@
|
||||
package forwarder
|
||||
|
||||
import "strings"
|
||||
|
||||
// prepareWriteContentsForClient applies a narrow Windows compatibility shim for
|
||||
// whole-file Write. Some Windows clients write CRLF payloads as CRCRLF on disk,
|
||||
// so we send LF-only content when the target path is clearly a Windows path and
|
||||
// the logical contents already use CRLF. The client's text-mode write then
|
||||
// expands LF back to the intended CRLF on disk.
|
||||
func prepareWriteContentsForClient(path string, contents string) string {
|
||||
if !looksLikeWindowsPath(path) || !strings.Contains(contents, "\r\n") {
|
||||
return contents
|
||||
}
|
||||
return normalizeCRLFToLF(contents)
|
||||
}
|
||||
|
||||
// preparePatchEditWriteContentsForClient sends LF-only transport for PatchEdit.
|
||||
// The client write path may expand LF to the platform newline, so sending CRLF
|
||||
// can become CRCRLF on disk.
|
||||
func preparePatchEditWriteContentsForClient(contents string) string {
|
||||
if !strings.Contains(contents, "\r") {
|
||||
return contents
|
||||
}
|
||||
return normalizeObservedLineEndingsToLF(contents)
|
||||
}
|
||||
|
||||
// reconcilePostWriteObservedContent only corrects one known Windows anomaly:
|
||||
// every line ending in the expected content was duplicated during write-back
|
||||
// verification, which makes the file appear to contain a blank line between
|
||||
// each real line. Some read paths normalize line endings to LF, so the
|
||||
// comparison also checks the normalized representation.
|
||||
func reconcilePostWriteObservedContent(expected string, observed string) (string, bool) {
|
||||
if observed == expected {
|
||||
return observed, false
|
||||
}
|
||||
if observed == systematicallyDoubledLineEndings(expected) {
|
||||
return expected, true
|
||||
}
|
||||
expectedNormalized := normalizeObservedLineEndingsToLF(expected)
|
||||
observedNormalized := normalizeObservedLineEndingsToLF(observed)
|
||||
if observedNormalized == expectedNormalized {
|
||||
return expected, true
|
||||
}
|
||||
if observedNormalized == systematicallyDoubledLineEndings(expectedNormalized) {
|
||||
return expected, true
|
||||
}
|
||||
return observed, false
|
||||
}
|
||||
|
||||
func systematicallyDoubledLineEndings(content string) string {
|
||||
if content == "" {
|
||||
return ""
|
||||
}
|
||||
|
||||
var builder strings.Builder
|
||||
builder.Grow(len(content) * 2)
|
||||
for index := 0; index < len(content); {
|
||||
switch {
|
||||
case content[index] == '\r' && index+1 < len(content) && content[index+1] == '\n':
|
||||
builder.WriteString("\r\n\r\n")
|
||||
index += 2
|
||||
case content[index] == '\n':
|
||||
builder.WriteString("\n\n")
|
||||
index++
|
||||
default:
|
||||
builder.WriteByte(content[index])
|
||||
index++
|
||||
}
|
||||
}
|
||||
return builder.String()
|
||||
}
|
||||
|
||||
func normalizeObservedLineEndingsToLF(content string) string {
|
||||
if content == "" {
|
||||
return ""
|
||||
}
|
||||
normalized := normalizeCRLFToLF(content)
|
||||
return strings.ReplaceAll(normalized, "\r", "\n")
|
||||
}
|
||||
|
||||
func normalizeCRLFToLF(content string) string {
|
||||
if content == "" || !strings.Contains(content, "\r\n") {
|
||||
return content
|
||||
}
|
||||
return strings.ReplaceAll(content, "\r\n", "\n")
|
||||
}
|
||||
|
||||
func looksLikeWindowsPath(path string) bool {
|
||||
trimmed := strings.TrimSpace(path)
|
||||
if len(trimmed) >= 3 && isASCIILetter(trimmed[0]) && trimmed[1] == ':' {
|
||||
return trimmed[2] == '\\' || trimmed[2] == '/'
|
||||
}
|
||||
return strings.HasPrefix(trimmed, `\\`)
|
||||
}
|
||||
|
||||
func isASCIILetter(value byte) bool {
|
||||
return (value >= 'a' && value <= 'z') || (value >= 'A' && value <= 'Z')
|
||||
}
|
||||
@@ -0,0 +1,18 @@
|
||||
//go:build !windows
|
||||
|
||||
package forwarder
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"os"
|
||||
"syscall"
|
||||
)
|
||||
|
||||
func processExists(pid int) bool {
|
||||
process, err := os.FindProcess(pid)
|
||||
if err != nil || process == nil {
|
||||
return false
|
||||
}
|
||||
err = process.Signal(syscall.Signal(0))
|
||||
return err == nil || errors.Is(err, syscall.EPERM)
|
||||
}
|
||||
@@ -0,0 +1,18 @@
|
||||
//go:build windows
|
||||
|
||||
package forwarder
|
||||
|
||||
import (
|
||||
"errors"
|
||||
|
||||
"golang.org/x/sys/windows"
|
||||
)
|
||||
|
||||
func processExists(pid int) bool {
|
||||
handle, err := windows.OpenProcess(windows.PROCESS_QUERY_LIMITED_INFORMATION, false, uint32(pid))
|
||||
if err == nil {
|
||||
_ = windows.CloseHandle(handle)
|
||||
return true
|
||||
}
|
||||
return errors.Is(err, windows.ERROR_ACCESS_DENIED)
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,130 @@
|
||||
package forwarder
|
||||
|
||||
import (
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"strings"
|
||||
|
||||
modeladapter "cursor/internal/backend/agent/model"
|
||||
)
|
||||
|
||||
func newPromptContextMessage(source string, message modeladapter.Message, persist bool) PromptContextMessage {
|
||||
context := PromptContextMessage{
|
||||
Source: strings.TrimSpace(source),
|
||||
Message: message,
|
||||
Persist: persist,
|
||||
}
|
||||
context.ContentHash = promptContextContentHash(context.Message)
|
||||
return context
|
||||
}
|
||||
|
||||
func normalizePromptContextMessage(context PromptContextMessage) PromptContextMessage {
|
||||
context.Source = strings.TrimSpace(context.Source)
|
||||
context.Message.Role = strings.TrimSpace(context.Message.Role)
|
||||
context.Message.Content = strings.TrimSpace(context.Message.Content)
|
||||
context.ContentHash = strings.TrimSpace(context.ContentHash)
|
||||
if context.ContentHash == "" {
|
||||
context.ContentHash = promptContextContentHash(context.Message)
|
||||
}
|
||||
return context
|
||||
}
|
||||
|
||||
func promptContextContentHash(message modeladapter.Message) string {
|
||||
sum := sha256.Sum256([]byte(strings.TrimSpace(message.Role) + "\x00" + strings.TrimSpace(message.Content)))
|
||||
return hex.EncodeToString(sum[:])
|
||||
}
|
||||
|
||||
func filterCurrentTurnPromptContexts(conversation *ConversationFile, contexts []PromptContextMessage) []PromptContextMessage {
|
||||
if len(contexts) == 0 {
|
||||
return nil
|
||||
}
|
||||
seen := collectCurrentTurnPromptContextKeys(conversation)
|
||||
filtered := make([]PromptContextMessage, 0, len(contexts))
|
||||
for _, context := range contexts {
|
||||
context = normalizePromptContextMessage(context)
|
||||
if !isReplayablePromptContext(context) {
|
||||
continue
|
||||
}
|
||||
if context.Persist {
|
||||
if _, ok := seen[promptContextKey(context)]; ok {
|
||||
continue
|
||||
}
|
||||
}
|
||||
filtered = append(filtered, context)
|
||||
}
|
||||
return filtered
|
||||
}
|
||||
|
||||
func collectCurrentTurnPromptContextKeys(conversation *ConversationFile) map[string]struct{} {
|
||||
keys := make(map[string]struct{})
|
||||
if conversation == nil {
|
||||
return keys
|
||||
}
|
||||
currentTurnSeq := conversation.NextTurnSeq - 1
|
||||
if currentTurnSeq <= 0 {
|
||||
return keys
|
||||
}
|
||||
for _, entry := range conversation.Entries {
|
||||
if entry.TurnSeq != currentTurnSeq || strings.TrimSpace(entry.Kind) != "prompt_context" {
|
||||
continue
|
||||
}
|
||||
var payload promptContextEntryPayload
|
||||
if err := json.Unmarshal(entry.Payload, &payload); err != nil {
|
||||
continue
|
||||
}
|
||||
context := normalizePromptContextMessage(PromptContextMessage{
|
||||
Source: payload.Source,
|
||||
ContentHash: payload.ContentHash,
|
||||
Message: modeladapter.Message{
|
||||
Role: payload.Role,
|
||||
Content: payload.Content,
|
||||
},
|
||||
Persist: true,
|
||||
})
|
||||
if !isReplayablePromptContext(context) {
|
||||
continue
|
||||
}
|
||||
keys[promptContextKey(context)] = struct{}{}
|
||||
}
|
||||
return keys
|
||||
}
|
||||
|
||||
func promptContextKey(context PromptContextMessage) string {
|
||||
context = normalizePromptContextMessage(context)
|
||||
return context.Source + "\x00" + context.ContentHash
|
||||
}
|
||||
|
||||
func isReplayablePromptContext(context PromptContextMessage) bool {
|
||||
context = normalizePromptContextMessage(context)
|
||||
if context.Source == "" || strings.TrimSpace(context.Message.Content) == "" {
|
||||
return false
|
||||
}
|
||||
if context.Message.Role == "" {
|
||||
context.Message.Role = "user"
|
||||
}
|
||||
if context.Message.Role != "user" && context.Message.Role != "system" {
|
||||
return false
|
||||
}
|
||||
if strings.TrimSpace(context.Message.Name) != "" || strings.TrimSpace(context.Message.ToolCallID) != "" {
|
||||
return false
|
||||
}
|
||||
return len(context.Message.ToolCalls) == 0 && len(context.Message.ContentParts) == 0
|
||||
}
|
||||
|
||||
func newPromptContextEntry(turnSeq int64, requestID string, context PromptContextMessage) HistoryEntry {
|
||||
context = normalizePromptContextMessage(context)
|
||||
payload, _ := json.Marshal(promptContextEntryPayload{
|
||||
Source: context.Source,
|
||||
Role: firstNonEmpty(strings.TrimSpace(context.Message.Role), "user"),
|
||||
Content: strings.TrimSpace(context.Message.Content),
|
||||
ContentHash: context.ContentHash,
|
||||
})
|
||||
return HistoryEntry{
|
||||
TurnSeq: turnSeq,
|
||||
RequestID: strings.TrimSpace(requestID),
|
||||
Role: "system",
|
||||
Kind: "prompt_context",
|
||||
Payload: payload,
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,323 @@
|
||||
package forwarder
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"sort"
|
||||
"strings"
|
||||
"unicode/utf8"
|
||||
|
||||
"google.golang.org/protobuf/proto"
|
||||
|
||||
"cursor/gen/agentv1"
|
||||
)
|
||||
|
||||
const (
|
||||
promptGuardUserTextChars = 32000
|
||||
promptGuardUserRichTextChars = 32000
|
||||
promptGuardSubagentReminderChars = 12000
|
||||
promptGuardSelectedFileChars = 16000
|
||||
promptGuardSelectedFilesTotalChars = 64000
|
||||
promptGuardSelectedFilesMaxCount = 12
|
||||
promptGuardRequestFileChars = 16000
|
||||
promptGuardRequestFilesTotalChars = 64000
|
||||
promptGuardRequestFilesMaxCount = 12
|
||||
promptGuardRuleChars = 6000
|
||||
promptGuardRulesTotalChars = 24000
|
||||
promptGuardRulesMaxCount = 40
|
||||
promptGuardSkillDescriptionChars = 800
|
||||
promptGuardSkillDescriptionsTotal = 16000
|
||||
promptGuardSkillDescriptorsMaxCount = 32
|
||||
promptGuardAgentSkillContentChars = 6000
|
||||
promptGuardAgentSkillsMaxCount = 16
|
||||
promptGuardRealtimeTextChars = 12000
|
||||
promptGuardCompiledMessageChars = 120000
|
||||
)
|
||||
|
||||
func normalizeUserMessageForStorage(userMessage *agentv1.UserMessage) *agentv1.UserMessage {
|
||||
if userMessage == nil {
|
||||
return nil
|
||||
}
|
||||
cloned, ok := proto.Clone(userMessage).(*agentv1.UserMessage)
|
||||
if !ok || cloned == nil {
|
||||
return userMessage
|
||||
}
|
||||
cloned.Text = truncatePromptGuardText("user_message.text", cloned.GetText(), promptGuardUserTextChars)
|
||||
if richText := strings.TrimSpace(cloned.GetRichText()); richText != "" {
|
||||
cloned.RichText = stringPtr(truncatePromptGuardText("user_message.rich_text", richText, promptGuardUserRichTextChars))
|
||||
}
|
||||
if reminder := strings.TrimSpace(cloned.GetSubagentSystemReminder()); reminder != "" {
|
||||
cloned.SubagentSystemReminder = stringPtr(truncatePromptGuardText("user_message.subagent_system_reminder", reminder, promptGuardSubagentReminderChars))
|
||||
}
|
||||
cloned.SelectedContext = guardSelectedContext(cloned.GetSelectedContext())
|
||||
return cloned
|
||||
}
|
||||
|
||||
func guardRequestContextForStorage(requestContext *agentv1.RequestContext) {
|
||||
if requestContext == nil {
|
||||
return
|
||||
}
|
||||
requestContext.Rules = guardCursorRules(requestContext.GetRules())
|
||||
requestContext.FileContents = guardStringMap(
|
||||
requestContext.GetFileContents(),
|
||||
"request_context.file_contents",
|
||||
promptGuardRequestFileChars,
|
||||
promptGuardRequestFilesTotalChars,
|
||||
promptGuardRequestFilesMaxCount,
|
||||
)
|
||||
if summary := strings.TrimSpace(requestContext.GetUserIntentSummary()); summary != "" {
|
||||
requestContext.UserIntentSummary = stringPtr(truncatePromptGuardText("request_context.user_intent_summary", summary, promptGuardRealtimeTextChars))
|
||||
}
|
||||
if hooks := strings.TrimSpace(requestContext.GetHooksAdditionalContext()); hooks != "" {
|
||||
requestContext.HooksAdditionalContext = stringPtr(truncatePromptGuardText("request_context.hooks_additional_context", hooks, promptGuardRealtimeTextChars))
|
||||
}
|
||||
if commit := strings.TrimSpace(requestContext.GetCommitAttributionMessage()); commit != "" {
|
||||
requestContext.CommitAttributionMessage = stringPtr(truncatePromptGuardText("request_context.commit_attribution_message", commit, promptGuardRealtimeTextChars))
|
||||
}
|
||||
if pr := strings.TrimSpace(requestContext.GetPrAttributionMessage()); pr != "" {
|
||||
requestContext.PrAttributionMessage = stringPtr(truncatePromptGuardText("request_context.pr_attribution_message", pr, promptGuardRealtimeTextChars))
|
||||
}
|
||||
}
|
||||
|
||||
func guardCompiledConversationForProvider(compiled CompiledConversation) CompiledConversation {
|
||||
for index := range compiled.Messages {
|
||||
message := &compiled.Messages[index]
|
||||
if strings.TrimSpace(message.Role) == "system" {
|
||||
continue
|
||||
}
|
||||
if strings.TrimSpace(message.Content) != "" {
|
||||
message.Content = truncatePromptGuardText("compiled."+firstNonEmpty(strings.TrimSpace(message.Role), "message"), message.Content, promptGuardCompiledMessageChars)
|
||||
}
|
||||
for partIndex := range message.ContentParts {
|
||||
if strings.TrimSpace(message.ContentParts[partIndex].Text) == "" {
|
||||
continue
|
||||
}
|
||||
message.ContentParts[partIndex].Text = truncatePromptGuardText("compiled.content_part", message.ContentParts[partIndex].Text, promptGuardCompiledMessageChars)
|
||||
}
|
||||
}
|
||||
return compiled
|
||||
}
|
||||
|
||||
func guardSelectedContext(selectedContext *agentv1.SelectedContext) *agentv1.SelectedContext {
|
||||
if selectedContext == nil {
|
||||
return nil
|
||||
}
|
||||
cloned, ok := proto.Clone(selectedContext).(*agentv1.SelectedContext)
|
||||
if !ok || cloned == nil {
|
||||
return selectedContext
|
||||
}
|
||||
cloned.Files = guardSelectedFiles(cloned.GetFiles())
|
||||
cloned.SelectedSkills = guardAgentSkills(cloned.GetSelectedSkills())
|
||||
cloned.ExtraContext = guardStringSlice(cloned.GetExtraContext(), "selected_context.extra_context", promptGuardRealtimeTextChars, promptGuardRealtimeTextChars, promptGuardAgentSkillsMaxCount)
|
||||
return cloned
|
||||
}
|
||||
|
||||
func guardSelectedFiles(files []*agentv1.SelectedFile) []*agentv1.SelectedFile {
|
||||
if len(files) == 0 {
|
||||
return nil
|
||||
}
|
||||
result := make([]*agentv1.SelectedFile, 0, minInt(len(files), promptGuardSelectedFilesMaxCount))
|
||||
remaining := promptGuardSelectedFilesTotalChars
|
||||
for _, file := range files {
|
||||
if file == nil || len(result) >= promptGuardSelectedFilesMaxCount {
|
||||
continue
|
||||
}
|
||||
content := strings.TrimSpace(file.GetContent())
|
||||
if content == "" {
|
||||
continue
|
||||
}
|
||||
limit := minInt(promptGuardSelectedFileChars, remaining)
|
||||
if limit <= 0 {
|
||||
break
|
||||
}
|
||||
cloned, ok := proto.Clone(file).(*agentv1.SelectedFile)
|
||||
if !ok || cloned == nil {
|
||||
continue
|
||||
}
|
||||
cloned.Content = truncatePromptGuardText("selected_context.files.content", content, limit)
|
||||
remaining -= promptGuardRuneCount(cloned.GetContent())
|
||||
result = append(result, cloned)
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
func guardCursorRules(rules []*agentv1.CursorRule) []*agentv1.CursorRule {
|
||||
if len(rules) == 0 {
|
||||
return nil
|
||||
}
|
||||
result := make([]*agentv1.CursorRule, 0, minInt(len(rules), promptGuardRulesMaxCount))
|
||||
remaining := promptGuardRulesTotalChars
|
||||
for _, rule := range rules {
|
||||
if rule == nil || len(result) >= promptGuardRulesMaxCount {
|
||||
continue
|
||||
}
|
||||
content := strings.TrimSpace(rule.GetContent())
|
||||
if content == "" {
|
||||
continue
|
||||
}
|
||||
limit := minInt(promptGuardRuleChars, remaining)
|
||||
if limit <= 0 {
|
||||
break
|
||||
}
|
||||
cloned, ok := proto.Clone(rule).(*agentv1.CursorRule)
|
||||
if !ok || cloned == nil {
|
||||
continue
|
||||
}
|
||||
cloned.Content = truncatePromptGuardText("request_context.rules.content", content, limit)
|
||||
remaining -= promptGuardRuneCount(cloned.GetContent())
|
||||
result = append(result, cloned)
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
func guardSkillDescriptors(descriptors []*agentv1.SkillDescriptor) []*agentv1.SkillDescriptor {
|
||||
if len(descriptors) == 0 {
|
||||
return nil
|
||||
}
|
||||
result := make([]*agentv1.SkillDescriptor, 0, minInt(len(descriptors), promptGuardSkillDescriptorsMaxCount))
|
||||
remaining := promptGuardSkillDescriptionsTotal
|
||||
for _, descriptor := range descriptors {
|
||||
if descriptor == nil || len(result) >= promptGuardSkillDescriptorsMaxCount {
|
||||
continue
|
||||
}
|
||||
description := strings.TrimSpace(descriptor.GetDescription())
|
||||
if description == "" {
|
||||
continue
|
||||
}
|
||||
limit := minInt(promptGuardSkillDescriptionChars, remaining)
|
||||
if limit <= 0 {
|
||||
break
|
||||
}
|
||||
cloned, ok := proto.Clone(descriptor).(*agentv1.SkillDescriptor)
|
||||
if !ok || cloned == nil {
|
||||
continue
|
||||
}
|
||||
cloned.Description = truncatePromptGuardText("skill.description", description, limit)
|
||||
remaining -= promptGuardRuneCount(cloned.GetDescription())
|
||||
result = append(result, cloned)
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
func guardAgentSkills(skills []*agentv1.AgentSkill) []*agentv1.AgentSkill {
|
||||
if len(skills) == 0 {
|
||||
return nil
|
||||
}
|
||||
result := make([]*agentv1.AgentSkill, 0, minInt(len(skills), promptGuardAgentSkillsMaxCount))
|
||||
for _, skill := range skills {
|
||||
if skill == nil || len(result) >= promptGuardAgentSkillsMaxCount {
|
||||
continue
|
||||
}
|
||||
cloned, ok := proto.Clone(skill).(*agentv1.AgentSkill)
|
||||
if !ok || cloned == nil {
|
||||
continue
|
||||
}
|
||||
if content := strings.TrimSpace(cloned.GetContent()); content != "" {
|
||||
cloned.Content = truncatePromptGuardText("agent_skill.content", content, promptGuardAgentSkillContentChars)
|
||||
}
|
||||
if description := strings.TrimSpace(cloned.GetDescription()); description != "" {
|
||||
cloned.Description = truncatePromptGuardText("agent_skill.description", description, promptGuardSkillDescriptionChars)
|
||||
}
|
||||
result = append(result, cloned)
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
func guardStringMap(input map[string]string, label string, itemLimit int, totalLimit int, maxItems int) map[string]string {
|
||||
if len(input) == 0 {
|
||||
return nil
|
||||
}
|
||||
keys := make([]string, 0, len(input))
|
||||
values := make(map[string]string, len(input))
|
||||
for key, value := range input {
|
||||
trimmedKey := strings.TrimSpace(key)
|
||||
trimmedValue := strings.TrimSpace(value)
|
||||
if trimmedKey == "" || trimmedValue == "" {
|
||||
continue
|
||||
}
|
||||
if _, exists := values[trimmedKey]; exists {
|
||||
continue
|
||||
}
|
||||
keys = append(keys, trimmedKey)
|
||||
values[trimmedKey] = trimmedValue
|
||||
}
|
||||
sort.Strings(keys)
|
||||
result := make(map[string]string, minInt(len(keys), maxItems))
|
||||
remaining := totalLimit
|
||||
for _, key := range keys {
|
||||
if len(result) >= maxItems {
|
||||
break
|
||||
}
|
||||
content := values[key]
|
||||
if content == "" {
|
||||
continue
|
||||
}
|
||||
limit := minInt(itemLimit, remaining)
|
||||
if limit <= 0 {
|
||||
break
|
||||
}
|
||||
result[key] = truncatePromptGuardText(label, content, limit)
|
||||
remaining -= promptGuardRuneCount(result[key])
|
||||
}
|
||||
if len(result) == 0 {
|
||||
return nil
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
func guardStringSlice(input []string, label string, itemLimit int, totalLimit int, maxItems int) []string {
|
||||
if len(input) == 0 {
|
||||
return nil
|
||||
}
|
||||
result := make([]string, 0, minInt(len(input), maxItems))
|
||||
remaining := totalLimit
|
||||
for _, item := range input {
|
||||
if len(result) >= maxItems {
|
||||
break
|
||||
}
|
||||
content := strings.TrimSpace(item)
|
||||
if content == "" {
|
||||
continue
|
||||
}
|
||||
limit := minInt(itemLimit, remaining)
|
||||
if limit <= 0 {
|
||||
break
|
||||
}
|
||||
truncated := truncatePromptGuardText(label, content, limit)
|
||||
remaining -= promptGuardRuneCount(truncated)
|
||||
result = append(result, truncated)
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
func truncatePromptGuardText(label string, text string, limit int) string {
|
||||
if limit <= 0 {
|
||||
return ""
|
||||
}
|
||||
if promptGuardRuneCount(text) <= limit {
|
||||
return text
|
||||
}
|
||||
runes := []rune(text)
|
||||
notice := fmt.Sprintf("\n\n[truncated: %s exceeded %d chars; kept head and tail from %d chars]\n\n", strings.TrimSpace(label), limit, len(runes))
|
||||
noticeRunes := []rune(notice)
|
||||
keep := limit - len(noticeRunes)
|
||||
if keep <= 0 {
|
||||
return string(runes[:limit])
|
||||
}
|
||||
head := keep * 2 / 3
|
||||
tail := keep - head
|
||||
if tail <= 0 {
|
||||
return string(runes[:head]) + notice
|
||||
}
|
||||
return string(runes[:head]) + notice + string(runes[len(runes)-tail:])
|
||||
}
|
||||
|
||||
func promptGuardRuneCount(text string) int {
|
||||
return utf8.RuneCountInString(text)
|
||||
}
|
||||
|
||||
func minInt(a int, b int) int {
|
||||
if a < b {
|
||||
return a
|
||||
}
|
||||
return b
|
||||
}
|
||||
@@ -0,0 +1,62 @@
|
||||
// provider.go 把 forwarder 的 canonical 请求转交给现有的 provider adapter 层。
|
||||
package forwarder
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"strings"
|
||||
|
||||
modeladapter "cursor/internal/backend/agent/model"
|
||||
)
|
||||
|
||||
type DefaultProviderGateway struct {
|
||||
router modeladapter.ModelAdapterRouter
|
||||
}
|
||||
|
||||
// NewProviderGateway 创建默认 provider 网关。
|
||||
func NewProviderGateway(resolver modeladapter.ChannelResolver) *DefaultProviderGateway {
|
||||
return &DefaultProviderGateway{
|
||||
router: modeladapter.NewRouter(resolver),
|
||||
}
|
||||
}
|
||||
|
||||
// StartStream 把 forwarder 的 provider 请求翻译成 modeladapter.StreamRequest 并发起流式调用。
|
||||
func (gateway *DefaultProviderGateway) StartStream(ctx context.Context, req ProviderRequest, sink func(modeladapter.ModelEvent) error) error {
|
||||
if ctx == nil {
|
||||
ctx = context.Background()
|
||||
}
|
||||
requestKnobs := make(map[string]any, len(req.RequestKnobs)+2)
|
||||
for key, value := range req.RequestKnobs {
|
||||
requestKnobs[key] = value
|
||||
}
|
||||
requestKnobs["stream"] = true
|
||||
if req.MaxTokens > 0 {
|
||||
requestKnobs["max_tokens"] = req.MaxTokens
|
||||
}
|
||||
if strings.TrimSpace(req.ThinkingEffort) != "" {
|
||||
requestKnobs["runtime_thinking_effort"] = strings.TrimSpace(req.ThinkingEffort)
|
||||
}
|
||||
err := gateway.router.Stream(ctx, modeladapter.StreamRequest{
|
||||
RequestID: req.RequestID,
|
||||
RunID: req.RunID,
|
||||
ModelCallID: req.ModelCallID,
|
||||
ConversationID: req.ConversationID,
|
||||
Mode: req.Mode,
|
||||
ModelID: req.ModelID,
|
||||
ThinkingEffort: req.ThinkingEffort,
|
||||
Messages: req.Messages,
|
||||
StableMessageCount: req.StableMessageCount,
|
||||
Tools: append([]json.RawMessage(nil), req.Tools...),
|
||||
MaxTokens: req.MaxTokens,
|
||||
Stream: true,
|
||||
RequestKnobs: requestKnobs,
|
||||
CompileSummary: req.CompileSummary,
|
||||
Observer: req.Observer,
|
||||
ArtifactPaths: req.ArtifactPaths,
|
||||
RequestBodyOverride: req.RequestBodyOverride,
|
||||
}, sink)
|
||||
if err != nil {
|
||||
return providerTerminalError{cause: err}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,16 @@
|
||||
// reasoning_metadata.go 保存 history/replay 共用的 reasoning metadata 判定。
|
||||
package forwarder
|
||||
|
||||
import (
|
||||
"strings"
|
||||
|
||||
modeladapter "cursor/internal/backend/agent/model"
|
||||
)
|
||||
|
||||
func hasReplayableReasoningPayload(reasoningContent string, reasoningSignature string, reasoningSignatureSource string) bool {
|
||||
if strings.TrimSpace(reasoningContent) != "" {
|
||||
return true
|
||||
}
|
||||
return strings.TrimSpace(reasoningSignature) != "" &&
|
||||
strings.TrimSpace(reasoningSignatureSource) == modeladapter.ReasoningSignatureSourceOpenAIResponses
|
||||
}
|
||||
@@ -0,0 +1,457 @@
|
||||
// reminders.go 负责按 mode 和上下文生成最小的 system reminder 集合。
|
||||
package forwarder
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
|
||||
"cursor/gen/agentv1"
|
||||
modeladapter "cursor/internal/backend/agent/model"
|
||||
promptassets "cursor/prompt"
|
||||
)
|
||||
|
||||
type DefaultReminderInjector struct {
|
||||
}
|
||||
|
||||
const (
|
||||
promptContextSourcePlanTurnContract = "plan_turn_contract"
|
||||
promptContextSourceActiveModeContract = "active_mode_contract"
|
||||
promptContextSourceLatestUserIntent = "latest_user_intent"
|
||||
promptContextSourceCurrentUserRequest = "current_user_request"
|
||||
promptContextSourceSubagentContract = "subagent_contract"
|
||||
promptContextSourceSubagentEmptyStopRecovery = "subagent_empty_stop_recovery"
|
||||
promptContextSourceDebugModeReminder = "debug_mode_reminder"
|
||||
)
|
||||
|
||||
// NewReminderInjector 创建默认 reminder 注入器。
|
||||
func NewReminderInjector() *DefaultReminderInjector {
|
||||
return &DefaultReminderInjector{}
|
||||
}
|
||||
|
||||
// Inject 根据 mode、最近用户输入和工具上下文生成本轮附加提醒。
|
||||
func (injector *DefaultReminderInjector) Inject(mode agentv1.AgentMode, conversation *ConversationFile, replayMessages []modeladapter.Message, latestUserText string, toolNames []string) PromptReminders {
|
||||
_ = toolNames
|
||||
reminders := make([]string, 0, 6)
|
||||
normalizedMode, err := validateSupportedActiveMode(mode)
|
||||
if err != nil {
|
||||
normalizedMode = agentv1.AgentMode_AGENT_MODE_AGENT
|
||||
}
|
||||
if conversation != nil && isChildConversationSubagentTypeName(conversation.SubagentTypeName) {
|
||||
return appendCurrentTurnAttentionReminders(PromptReminders{
|
||||
SystemParts: reminders,
|
||||
PromptContexts: []PromptContextMessage{
|
||||
newPromptContextReminder(promptContextSourceSubagentContract, subagentContractText()),
|
||||
newPromptContextReminder(promptContextSourceActiveModeContract, currentModeContractText(normalizedMode, true)),
|
||||
},
|
||||
}, latestUserText)
|
||||
}
|
||||
if normalizedMode == agentv1.AgentMode_AGENT_MODE_DEBUG {
|
||||
return debugModePromptReminders(conversation)
|
||||
}
|
||||
reminders = append(reminders, "If multiple <current_plan> or <todo_list> blocks appear in the conversation, treat the last block of each type as the current source of truth.")
|
||||
switch normalizedMode {
|
||||
case agentv1.AgentMode_AGENT_MODE_ASK:
|
||||
reminders = append(reminders, "You are in ask mode. Prefer direct answers and only use tools when they are necessary to answer accurately.")
|
||||
reminders = append(reminders, "Lead with the conclusion, keep the response concise, and avoid unsolicited example code or long bullet lists.")
|
||||
case agentv1.AgentMode_AGENT_MODE_PLAN:
|
||||
reminders = append(reminders, "You are currently working in plan mode for the user. Prioritize investigation, decomposition, tradeoff analysis, and producing or refining a concrete plan.")
|
||||
reminders = append(reminders, "Do not directly modify files in plan mode. Avoid direct file-editing tools such as Write, Delete, and PatchEdit.")
|
||||
reminders = append(reminders, "For non-trivial plan-mode work, first do a quick reconnaissance yourself, then launch 2-4 parallel Task subagents with subagent_type=\"explore\" to investigate distinct angles before CreatePlan. Avoid using exactly one subagent for broad tasks: either handle narrow tasks directly, or split broad tasks into multiple independent investigations. Synthesize the subagent results yourself; do not delegate the final plan.")
|
||||
reminders = append(reminders, "For narrow, well-scoped tasks, you can investigate directly and then create a lean plan with only the essential stages, tradeoffs, and next steps.")
|
||||
if hasCurrentPlan(conversation) {
|
||||
reminders = append(reminders, "A current plan already exists. Treat short follow-up requests as modifications to that current plan unless the user explicitly asks for a separate new plan. When calling CreatePlan for an existing plan, send the complete revised plan, preserve relevant existing content, incorporate the user's requested changes, and omit the name field. The CreatePlan name field is only allowed on the first CreatePlan call; never use a later name to rename or create a separate plan.")
|
||||
}
|
||||
case agentv1.AgentMode_AGENT_MODE_MULTITASK:
|
||||
reminders = append(reminders, "You are in multitask mode. Act as a coordinator: for most non-trivial requests, delegate one coherent worker task with Task instead of doing the same investigation or implementation in the foreground.")
|
||||
reminders = append(reminders, "After delegating the only coherent worker task for a request, do not continue the same work in the foreground. Only do distinct coordination work, answer a new independent question, or synthesize after multiple workers return.")
|
||||
reminders = append(reminders, "Do not wait, sleep, or poll just for a running worker to complete. End the response unless there is separate useful coordination to do.")
|
||||
reminders = append(reminders, "Do not over-decompose small or medium tasks into many sibling workers. Use multiple sibling workers only for clearly independent top-level workstreams.")
|
||||
default:
|
||||
reminders = append(reminders, "You are in agent mode. Use the available tools when they materially improve correctness or efficiency.")
|
||||
reminders = append(reminders, "When reporting progress or completion, lead with the result, mention only key changes or verification, and avoid long recaps, exhaustive lists, or unsolicited example code.")
|
||||
}
|
||||
|
||||
result := PromptReminders{
|
||||
SystemParts: reminders,
|
||||
PromptContexts: []PromptContextMessage{
|
||||
newPromptContextReminder(promptContextSourceActiveModeContract, currentModeContractText(normalizedMode, false)),
|
||||
},
|
||||
}
|
||||
if normalizedMode == agentv1.AgentMode_AGENT_MODE_PLAN {
|
||||
if reminder := strings.TrimSpace(promptassets.MustReadPlanSystemReminder()); reminder != "" {
|
||||
result.PromptContexts = append([]PromptContextMessage{
|
||||
newPromptContextReminder(promptContextSourcePlanTurnContract, reminder),
|
||||
}, result.PromptContexts...)
|
||||
}
|
||||
return appendCurrentTurnAttentionReminders(result, latestUserText)
|
||||
}
|
||||
if normalizedMode != agentv1.AgentMode_AGENT_MODE_AGENT {
|
||||
return appendCurrentTurnAttentionReminders(result, latestUserText)
|
||||
}
|
||||
|
||||
candidate, ok := extractLatestSuccessfulEditReminder(replayMessages)
|
||||
if !ok {
|
||||
return appendCurrentTurnAttentionReminders(result, latestUserText)
|
||||
}
|
||||
result.PromptContexts = append(result.PromptContexts, newPromptContextMessage(
|
||||
"latest_edit_reminder",
|
||||
modeladapter.Message{
|
||||
Role: "user",
|
||||
Content: strings.TrimSpace(fmt.Sprintf(`<system_reminder>
|
||||
You recently successfully edited %q.
|
||||
|
||||
For this file, the latest source of truth is the most recent successful %s, not earlier reads or memory.
|
||||
|
||||
When modifying this file:
|
||||
- use PatchEdit with path, old_string, new_string, and optional replace_all
|
||||
- copy old_string exactly from the latest file content; line endings are not normalized or treated equivalently during matching
|
||||
- replace_all defaults to false, so old_string must match exactly one occurrence unless you intentionally set replace_all to true
|
||||
- new_string may be empty to delete old_string
|
||||
- preserve spaces, tabs, indentation, punctuation, and line endings exactly in old_string
|
||||
- only read the file again yourself if you need exact current content or extra context to choose a unique old_string
|
||||
</system_reminder>`, candidate.Path, candidate.SourceField)),
|
||||
},
|
||||
false,
|
||||
))
|
||||
return appendCurrentTurnAttentionReminders(result, latestUserText)
|
||||
}
|
||||
|
||||
func appendCurrentTurnAttentionReminders(result PromptReminders, latestUserText string) PromptReminders {
|
||||
if reminder := latestUserIntentReminderText(latestUserText); reminder != "" {
|
||||
result.PromptContexts = append(result.PromptContexts, newPromptContextReminder(promptContextSourceLatestUserIntent, reminder))
|
||||
}
|
||||
return appendCurrentUserRequestReminder(result, latestUserText)
|
||||
}
|
||||
|
||||
func appendCurrentUserRequestReminder(result PromptReminders, latestUserText string) PromptReminders {
|
||||
if strings.TrimSpace(latestUserText) == "" {
|
||||
return result
|
||||
}
|
||||
result.PromptContexts = append(result.PromptContexts, newCurrentUserRequestReminder(latestUserText))
|
||||
return result
|
||||
}
|
||||
|
||||
func latestUserIntentReminderText(latestUserText string) string {
|
||||
normalizedText := strings.ToLower(strings.TrimSpace(latestUserText))
|
||||
switch {
|
||||
case strings.Contains(normalizedText, "review"), strings.Contains(latestUserText, "评审"), strings.Contains(latestUserText, "审查"):
|
||||
return "When reviewing code, focus on bugs, regressions, behavioral risks, and missing tests."
|
||||
case strings.Contains(normalizedText, "plan"), strings.Contains(latestUserText, "计划"):
|
||||
return "Prefer clear staged plans with concrete checkpoints."
|
||||
default:
|
||||
return ""
|
||||
}
|
||||
}
|
||||
|
||||
func newCurrentUserRequestReminder(latestUserText string) PromptContextMessage {
|
||||
return newPromptContextMessage(
|
||||
promptContextSourceCurrentUserRequest,
|
||||
modeladapter.Message{
|
||||
Role: "user",
|
||||
Content: strings.TrimSpace(fmt.Sprintf(`<current_user_request>
|
||||
%s
|
||||
</current_user_request>
|
||||
|
||||
Handle this request. Use any latest tool results as evidence when present, and provide the requested answer once you have enough information. Treat surrounding system reminders as constraints, not as the user's task.`, strings.TrimSpace(latestUserText))),
|
||||
},
|
||||
false,
|
||||
)
|
||||
}
|
||||
|
||||
func newPromptContextReminder(source string, content string) PromptContextMessage {
|
||||
return newPromptContextMessage(
|
||||
source,
|
||||
modeladapter.Message{
|
||||
Role: "user",
|
||||
Content: wrapSystemReminder(content),
|
||||
},
|
||||
false,
|
||||
)
|
||||
}
|
||||
|
||||
func subagentContractText() string {
|
||||
return strings.Join([]string{
|
||||
"The turn that contains this reminder runs inside a subagent child conversation. Work as an investigator for the parent agent, not as the final user-facing assistant.",
|
||||
"Return a short textual result: lead with the conclusion, keep only the key evidence, and do not produce a long response.",
|
||||
"Use the available agent tools when they materially improve correctness or efficiency. Do not ask the user questions. If required information is missing, report the gap to the parent agent instead of asking the user directly.",
|
||||
}, "\n\n")
|
||||
}
|
||||
|
||||
func currentModeContractText(mode agentv1.AgentMode, childSubagent bool) string {
|
||||
if childSubagent {
|
||||
return "For the turn that contains this reminder, the active mode is a subagent child conversation. Use the available agent tools, but do not call AskQuestion. Return only a concise investigation result for the parent agent."
|
||||
}
|
||||
switch normalizeMode(mode) {
|
||||
case agentv1.AgentMode_AGENT_MODE_PLAN:
|
||||
return "For the turn that contains this reminder, the active mode is plan. Do not modify files or system state. Use CreatePlan when the plan is ready or needs updating."
|
||||
case agentv1.AgentMode_AGENT_MODE_ASK:
|
||||
return "For the turn that contains this reminder, the active mode is ask. Prefer a direct answer. Use tools only when they materially improve accuracy, and do not call CreatePlan."
|
||||
case agentv1.AgentMode_AGENT_MODE_DEBUG:
|
||||
return "For the turn that contains this reminder, the active mode is debug. Follow the Debug Mode workflow from the static debug prompt: inspect or reproduce before editing, keep 3-5 concrete hypotheses, use the injected debug session log path when temporary instrumentation is useful, and verify with runtime evidence. Do not call CreatePlan or SwitchMode."
|
||||
case agentv1.AgentMode_AGENT_MODE_MULTITASK:
|
||||
return "For the turn that contains this reminder, the active mode is multitask. Act as the foreground coordinator: delegate most non-trivial work to a coherent worker with Task, avoid duplicating delegated work in the foreground, and do not wait just for a worker to finish."
|
||||
default:
|
||||
return "For the turn that contains this reminder, the active mode is agent. CreatePlan is not available in this mode; do not call CreatePlan. If the user explicitly asks to create or revise a plan, call SwitchMode to return to plan mode first. If there is an accepted or current plan, execute or continue the implementation using the available agent-mode tools."
|
||||
}
|
||||
}
|
||||
|
||||
func debugModePromptReminders(conversation *ConversationFile) PromptReminders {
|
||||
initial := !hasPreviousPromptContextSource(conversation, promptContextSourceDebugModeReminder)
|
||||
content := renderDebugSystemReminder(promptassets.MustReadDebugSystemReminder(initial), conversation)
|
||||
if strings.TrimSpace(content) == "" {
|
||||
return PromptReminders{}
|
||||
}
|
||||
return PromptReminders{
|
||||
PromptContexts: []PromptContextMessage{
|
||||
newPromptContextMessage(promptContextSourceDebugModeReminder, modeladapter.Message{
|
||||
Role: "user",
|
||||
Content: strings.TrimSpace(content),
|
||||
}, false),
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func hasPreviousPromptContextSource(conversation *ConversationFile, source string) bool {
|
||||
if conversation == nil {
|
||||
return false
|
||||
}
|
||||
needle := strings.TrimSpace(source)
|
||||
if needle == "" {
|
||||
return false
|
||||
}
|
||||
for _, entry := range conversation.Entries {
|
||||
if strings.TrimSpace(entry.Kind) != "prompt_context" {
|
||||
continue
|
||||
}
|
||||
var payload promptContextEntryPayload
|
||||
if err := json.Unmarshal(entry.Payload, &payload); err != nil {
|
||||
continue
|
||||
}
|
||||
if strings.TrimSpace(payload.Source) == needle {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func renderDebugSystemReminder(template string, conversation *ConversationFile) string {
|
||||
sessionID := firstNonEmpty(debugSessionID(conversation), "debug")
|
||||
logPath := debugLogPath(sessionID)
|
||||
serverEndpoint := debugServerEndpoint(sessionID)
|
||||
result := strings.TrimSpace(template)
|
||||
for _, replacement := range []struct {
|
||||
placeholder string
|
||||
value string
|
||||
}{
|
||||
{placeholder: "{{DEBUG_SERVER_ENDPOINT}}", value: serverEndpoint},
|
||||
{placeholder: "{{DEBUG_LOG_PATH}}", value: logPath},
|
||||
{placeholder: "{{DEBUG_SESSION_ID}}", value: sessionID},
|
||||
} {
|
||||
result = strings.ReplaceAll(result, replacement.placeholder, replacement.value)
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
func debugLogPath(sessionID string) string {
|
||||
filename := fmt.Sprintf("debug-%s.log", firstNonEmpty(strings.TrimSpace(sessionID), "debug"))
|
||||
cwd, err := os.Getwd()
|
||||
if err != nil || strings.TrimSpace(cwd) == "" {
|
||||
return filepath.Join(".cursor", filename)
|
||||
}
|
||||
return filepath.Join(cwd, ".cursor", filename)
|
||||
}
|
||||
|
||||
func debugServerEndpoint(sessionID string) string {
|
||||
return "http://127.0.0.1:7337/ingest/" + firstNonEmpty(strings.TrimSpace(sessionID), "debug")
|
||||
}
|
||||
|
||||
func debugSessionID(conversation *ConversationFile) string {
|
||||
if conversation == nil {
|
||||
return ""
|
||||
}
|
||||
for _, candidate := range []string{
|
||||
conversation.CurrentRequestID,
|
||||
conversation.CurrentLoopID,
|
||||
conversation.ConversationID,
|
||||
} {
|
||||
if normalized := normalizeDebugSessionID(candidate); normalized != "" {
|
||||
return normalized
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func normalizeDebugSessionID(raw string) string {
|
||||
trimmed := strings.TrimSpace(raw)
|
||||
if trimmed == "" {
|
||||
return ""
|
||||
}
|
||||
var builder strings.Builder
|
||||
for _, value := range trimmed {
|
||||
switch {
|
||||
case value >= 'a' && value <= 'z':
|
||||
builder.WriteRune(value)
|
||||
case value >= 'A' && value <= 'Z':
|
||||
builder.WriteRune(value)
|
||||
case value >= '0' && value <= '9':
|
||||
builder.WriteRune(value)
|
||||
case value == '-', value == '_':
|
||||
builder.WriteRune(value)
|
||||
default:
|
||||
builder.WriteRune('-')
|
||||
}
|
||||
}
|
||||
return strings.Trim(builder.String(), "-")
|
||||
}
|
||||
|
||||
func wrapSystemReminder(content string) string {
|
||||
trimmed := strings.TrimSpace(content)
|
||||
if trimmed == "" {
|
||||
return ""
|
||||
}
|
||||
if strings.HasPrefix(trimmed, "<system_reminder>") && strings.HasSuffix(trimmed, "</system_reminder>") {
|
||||
return trimmed
|
||||
}
|
||||
return "<system_reminder>\n" + trimmed + "\n</system_reminder>"
|
||||
}
|
||||
|
||||
func hasCurrentPlan(conversation *ConversationFile) bool {
|
||||
if conversation == nil {
|
||||
return false
|
||||
}
|
||||
state, err := projectConversationStructuredState(conversation)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
return state.HasPlan || len(state.Plans) > 0
|
||||
}
|
||||
|
||||
type editReminderCandidate struct {
|
||||
ToolCallID string
|
||||
Path string
|
||||
SourceField string
|
||||
}
|
||||
|
||||
type replayEditSuccess struct {
|
||||
Success *struct {
|
||||
Path string `json:"path"`
|
||||
AfterFullFileContent string `json:"afterFullFileContent"`
|
||||
AfterFullFileContentSnake string `json:"after_full_file_content"`
|
||||
DiffString string `json:"diffString"`
|
||||
DiffStringSnake string `json:"diff_string"`
|
||||
} `json:"success"`
|
||||
}
|
||||
|
||||
func extractLatestSuccessfulEditReminder(messages []modeladapter.Message) (editReminderCandidate, bool) {
|
||||
if len(messages) == 0 {
|
||||
return editReminderCandidate{}, false
|
||||
}
|
||||
toolPaths := collectReplayEditToolPaths(messages)
|
||||
for index := len(messages) - 1; index >= 0; index-- {
|
||||
message := messages[index]
|
||||
if strings.TrimSpace(message.Role) != "tool" {
|
||||
continue
|
||||
}
|
||||
toolName := strings.TrimSpace(message.Name)
|
||||
if !isEditReminderToolName(toolName) {
|
||||
continue
|
||||
}
|
||||
toolCallID := strings.TrimSpace(message.ToolCallID)
|
||||
if toolCallID == "" {
|
||||
continue
|
||||
}
|
||||
var result replayEditSuccess
|
||||
if err := json.Unmarshal([]byte(message.Content), &result); err != nil || result.Success == nil {
|
||||
continue
|
||||
}
|
||||
path := strings.TrimSpace(result.Success.Path)
|
||||
if path == "" {
|
||||
path = toolPaths[toolCallID]
|
||||
}
|
||||
if path == "" {
|
||||
continue
|
||||
}
|
||||
sourceField := editReminderSourceField(result)
|
||||
if sourceField == "" {
|
||||
continue
|
||||
}
|
||||
return editReminderCandidate{
|
||||
ToolCallID: toolCallID,
|
||||
Path: path,
|
||||
SourceField: sourceField,
|
||||
}, true
|
||||
}
|
||||
return editReminderCandidate{}, false
|
||||
}
|
||||
|
||||
func editReminderSourceField(result replayEditSuccess) string {
|
||||
if result.Success == nil {
|
||||
return ""
|
||||
}
|
||||
diffString := strings.TrimSpace(firstNonEmpty(result.Success.DiffStringSnake, result.Success.DiffString))
|
||||
if diffString != "" && !looksLikeProjectedReplayTruncation(diffString) {
|
||||
return "`success.diff_string`"
|
||||
}
|
||||
afterContent := strings.TrimSpace(firstNonEmpty(result.Success.AfterFullFileContentSnake, result.Success.AfterFullFileContent))
|
||||
if afterContent != "" && !looksLikeProjectedReplayTruncation(afterContent) {
|
||||
return "`success.after_full_file_content`"
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func looksLikeProjectedReplayTruncation(value string) bool {
|
||||
trimmed := strings.TrimSpace(value)
|
||||
return strings.Contains(trimmed, "[truncated:") || strings.Contains(trimmed, "[tool result replay truncated:")
|
||||
}
|
||||
|
||||
func collectReplayEditToolPaths(messages []modeladapter.Message) map[string]string {
|
||||
paths := make(map[string]string)
|
||||
for _, message := range messages {
|
||||
if strings.TrimSpace(message.Role) != "assistant" || len(message.ToolCalls) == 0 {
|
||||
continue
|
||||
}
|
||||
for _, toolCall := range message.ToolCalls {
|
||||
toolName := strings.TrimSpace(toolCall.Function.Name)
|
||||
if !isEditReminderToolName(toolName) {
|
||||
continue
|
||||
}
|
||||
toolCallID := strings.TrimSpace(toolCall.ID)
|
||||
if toolCallID == "" {
|
||||
continue
|
||||
}
|
||||
if path := extractPathFromToolArguments(toolCall.Function.Arguments); path != "" {
|
||||
paths[toolCallID] = path
|
||||
}
|
||||
}
|
||||
}
|
||||
return paths
|
||||
}
|
||||
|
||||
func isEditReminderToolName(toolName string) bool {
|
||||
switch strings.TrimSpace(toolName) {
|
||||
case "Write", "PatchEdit", "PatchEditLines", "PatchEditSpan":
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func extractPathFromToolArguments(arguments string) string {
|
||||
trimmed := strings.TrimSpace(arguments)
|
||||
if trimmed == "" || trimmed == "{}" || trimmed == "null" {
|
||||
return ""
|
||||
}
|
||||
var payload map[string]any
|
||||
if err := json.Unmarshal([]byte(trimmed), &payload); err != nil {
|
||||
return ""
|
||||
}
|
||||
for _, key := range []string{"path", "file_path"} {
|
||||
if value, ok := payload[key].(string); ok && strings.TrimSpace(value) != "" {
|
||||
return strings.TrimSpace(value)
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
@@ -0,0 +1,358 @@
|
||||
package forwarder
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"connectrpc.com/connect"
|
||||
|
||||
"cursor/gen/aiserverv1"
|
||||
"cursor/internal/logger"
|
||||
)
|
||||
|
||||
const (
|
||||
RepositoryServiceFastRepoInitHandshakeV2Procedure = "/aiserver.v1.RepositoryService/FastRepoInitHandshakeV2"
|
||||
RepositoryServiceFastRepoInitHandshakeProcedure = "/aiserver.v1.RepositoryService/FastRepoInitHandshake"
|
||||
RepositoryServiceFastRepoSyncCompleteProcedure = "/aiserver.v1.RepositoryService/FastRepoSyncComplete"
|
||||
RepositoryServiceSyncMerkleSubtreeV2Procedure = "/aiserver.v1.RepositoryService/SyncMerkleSubtreeV2"
|
||||
RepositoryServiceSyncMerkleSubtreeProcedure = "/aiserver.v1.RepositoryService/SyncMerkleSubtree"
|
||||
RepositoryServiceFastUpdateFileV2Procedure = "/aiserver.v1.RepositoryService/FastUpdateFileV2"
|
||||
RepositoryServiceFastUpdateFileProcedure = "/aiserver.v1.RepositoryService/FastUpdateFile"
|
||||
RepositoryServiceEnsureIndexCreatedProcedure = "/aiserver.v1.RepositoryService/EnsureIndexCreated"
|
||||
RepositoryServiceGetCopyStatusProcedure = "/aiserver.v1.RepositoryService/GetCopyStatus"
|
||||
RepositoryServiceGetUploadLimitsProcedure = "/aiserver.v1.RepositoryService/GetUploadLimits"
|
||||
RepositoryServiceGetNumFilesToSendProcedure = "/aiserver.v1.RepositoryService/GetNumFilesToSend"
|
||||
RepositoryServiceGetAvailableChunkingStrategiesProcedure = "/aiserver.v1.RepositoryService/GetAvailableChunkingStrategies"
|
||||
RepositoryServiceGetHighLevelFolderDescriptionProcedure = "/aiserver.v1.RepositoryService/GetHighLevelFolderDescription"
|
||||
RepositoryServiceRepositoryStatusProcedure = "/aiserver.v1.RepositoryService/RepositoryStatus"
|
||||
RepositoryServiceBatchRepositoryStatusProcedure = "/aiserver.v1.RepositoryService/BatchRepositoryStatus"
|
||||
)
|
||||
|
||||
func newRepositoryServiceHandler(service *Service) http.Handler {
|
||||
mux := http.NewServeMux()
|
||||
mux.Handle(
|
||||
RepositoryServiceFastRepoInitHandshakeV2Procedure,
|
||||
connect.NewUnaryHandler(RepositoryServiceFastRepoInitHandshakeV2Procedure, service.FastRepoInitHandshakeV2),
|
||||
)
|
||||
mux.Handle(
|
||||
RepositoryServiceFastRepoInitHandshakeProcedure,
|
||||
connect.NewUnaryHandler(RepositoryServiceFastRepoInitHandshakeProcedure, service.FastRepoInitHandshake),
|
||||
)
|
||||
mux.Handle(
|
||||
RepositoryServiceFastRepoSyncCompleteProcedure,
|
||||
connect.NewUnaryHandler(RepositoryServiceFastRepoSyncCompleteProcedure, service.FastRepoSyncComplete),
|
||||
)
|
||||
mux.Handle(
|
||||
RepositoryServiceSyncMerkleSubtreeV2Procedure,
|
||||
connect.NewUnaryHandler(RepositoryServiceSyncMerkleSubtreeV2Procedure, service.SyncMerkleSubtreeV2),
|
||||
)
|
||||
mux.Handle(
|
||||
RepositoryServiceSyncMerkleSubtreeProcedure,
|
||||
connect.NewUnaryHandler(RepositoryServiceSyncMerkleSubtreeProcedure, service.SyncMerkleSubtree),
|
||||
)
|
||||
mux.Handle(
|
||||
RepositoryServiceFastUpdateFileV2Procedure,
|
||||
connect.NewUnaryHandler(RepositoryServiceFastUpdateFileV2Procedure, service.FastUpdateFileV2),
|
||||
)
|
||||
mux.Handle(
|
||||
RepositoryServiceFastUpdateFileProcedure,
|
||||
connect.NewUnaryHandler(RepositoryServiceFastUpdateFileProcedure, service.FastUpdateFile),
|
||||
)
|
||||
mux.Handle(
|
||||
RepositoryServiceEnsureIndexCreatedProcedure,
|
||||
connect.NewUnaryHandler(RepositoryServiceEnsureIndexCreatedProcedure, service.EnsureIndexCreated),
|
||||
)
|
||||
mux.Handle(
|
||||
RepositoryServiceGetCopyStatusProcedure,
|
||||
connect.NewUnaryHandler(RepositoryServiceGetCopyStatusProcedure, service.GetCopyStatus),
|
||||
)
|
||||
mux.Handle(
|
||||
RepositoryServiceGetUploadLimitsProcedure,
|
||||
connect.NewUnaryHandler(RepositoryServiceGetUploadLimitsProcedure, service.GetUploadLimits),
|
||||
)
|
||||
mux.Handle(
|
||||
RepositoryServiceGetNumFilesToSendProcedure,
|
||||
connect.NewUnaryHandler(RepositoryServiceGetNumFilesToSendProcedure, service.GetNumFilesToSend),
|
||||
)
|
||||
mux.Handle(
|
||||
RepositoryServiceGetAvailableChunkingStrategiesProcedure,
|
||||
connect.NewUnaryHandler(RepositoryServiceGetAvailableChunkingStrategiesProcedure, service.GetAvailableChunkingStrategies),
|
||||
)
|
||||
mux.Handle(
|
||||
RepositoryServiceGetHighLevelFolderDescriptionProcedure,
|
||||
connect.NewUnaryHandler(RepositoryServiceGetHighLevelFolderDescriptionProcedure, service.GetHighLevelFolderDescription),
|
||||
)
|
||||
mux.Handle(
|
||||
RepositoryServiceRepositoryStatusProcedure,
|
||||
connect.NewUnaryHandler(RepositoryServiceRepositoryStatusProcedure, service.RepositoryStatus),
|
||||
)
|
||||
mux.Handle(
|
||||
RepositoryServiceBatchRepositoryStatusProcedure,
|
||||
connect.NewUnaryHandler(RepositoryServiceBatchRepositoryStatusProcedure, service.BatchRepositoryStatus),
|
||||
)
|
||||
mux.Handle("/", http.NotFoundHandler())
|
||||
return mux
|
||||
}
|
||||
|
||||
func (service *Service) FastRepoInitHandshakeV2(_ context.Context, req *connect.Request[aiserverv1.FastRepoInitHandshakeV2Request]) (*connect.Response[aiserverv1.FastRepoInitHandshakeV2Response], error) {
|
||||
repo := req.Msg.GetRepository()
|
||||
codebaseID := repositoryCodebaseID(repo)
|
||||
if err := service.ensureRepositoryCodebase(codebaseID, repo); err != nil {
|
||||
return nil, connect.NewError(connect.CodeInternal, err)
|
||||
}
|
||||
logger.Infof("RepositoryService FastRepoInitHandshakeV2 ready codebase_id=%s repo=%s root_hash_present=%t", codebaseID, repositoryLogName(repo), strings.TrimSpace(req.Msg.GetRootHash()) != "")
|
||||
return connect.NewResponse(&aiserverv1.FastRepoInitHandshakeV2Response{
|
||||
Status: aiserverv1.FastRepoInitHandshakeV2Response_STATUS_SUCCESS,
|
||||
Codebases: []*aiserverv1.RepositoryCodebaseInfo{
|
||||
{
|
||||
CodebaseId: codebaseID,
|
||||
Status: aiserverv1.RepositoryCodebaseInfo_STATUS_UP_TO_DATE,
|
||||
},
|
||||
},
|
||||
}), nil
|
||||
}
|
||||
|
||||
func (service *Service) FastRepoInitHandshake(_ context.Context, req *connect.Request[aiserverv1.FastRepoInitHandshakeRequest]) (*connect.Response[aiserverv1.FastRepoInitHandshakeResponse], error) {
|
||||
repo := req.Msg.GetRepository()
|
||||
codebaseID := repositoryCodebaseID(repo)
|
||||
if err := service.ensureRepositoryCodebase(codebaseID, repo); err != nil {
|
||||
return nil, connect.NewError(connect.CodeInternal, err)
|
||||
}
|
||||
repoName := repositoryRepoName(repo)
|
||||
logger.Infof("RepositoryService FastRepoInitHandshake ready codebase_id=%s repo=%s root_hash_present=%t", codebaseID, repositoryLogName(repo), strings.TrimSpace(req.Msg.GetRootHash()) != "")
|
||||
return connect.NewResponse(&aiserverv1.FastRepoInitHandshakeResponse{
|
||||
Status: aiserverv1.FastRepoInitHandshakeResponse_STATUS_UP_TO_DATE,
|
||||
RepoName: repoName,
|
||||
}), nil
|
||||
}
|
||||
|
||||
func (service *Service) FastRepoSyncComplete(_ context.Context, req *connect.Request[aiserverv1.FastRepoSyncCompleteRequest]) (*connect.Response[aiserverv1.FastRepoSyncCompleteResponse], error) {
|
||||
store, err := service.requireCodebaseIndexStore()
|
||||
if err != nil {
|
||||
return nil, connect.NewError(connect.CodeInternal, err)
|
||||
}
|
||||
for _, codebase := range req.Msg.GetCodebases() {
|
||||
codebaseID := strings.TrimSpace(codebase.GetCodebaseId())
|
||||
if codebaseID == "" {
|
||||
continue
|
||||
}
|
||||
if err := store.MarkRepositorySyncComplete(codebaseID, codebase.GetPathKeyHash()); err != nil {
|
||||
return nil, connect.NewError(connect.CodeInternal, err)
|
||||
}
|
||||
}
|
||||
logger.Infof("RepositoryService FastRepoSyncComplete acknowledged codebase_count=%d", len(req.Msg.GetCodebases()))
|
||||
return connect.NewResponse(&aiserverv1.FastRepoSyncCompleteResponse{}), nil
|
||||
}
|
||||
|
||||
func (service *Service) SyncMerkleSubtreeV2(_ context.Context, req *connect.Request[aiserverv1.SyncMerkleSubtreeV2Request]) (*connect.Response[aiserverv1.SyncMerkleSubtreeV2Response], error) {
|
||||
logger.Infof("RepositoryService SyncMerkleSubtreeV2 matched codebase_id=%s partial_paths=%d", strings.TrimSpace(req.Msg.GetCodebaseId()), len(req.Msg.GetLocalPartialPaths()))
|
||||
return connect.NewResponse(&aiserverv1.SyncMerkleSubtreeV2Response{
|
||||
Result: &aiserverv1.SyncMerkleSubtreeV2Response_Match{Match: true},
|
||||
Results: repositoryPartialPathMatchResults(len(req.Msg.GetLocalPartialPaths())),
|
||||
}), nil
|
||||
}
|
||||
|
||||
func (service *Service) SyncMerkleSubtree(_ context.Context, req *connect.Request[aiserverv1.SyncMerkleSubtreeRequest]) (*connect.Response[aiserverv1.SyncMerkleSubtreeResponse], error) {
|
||||
logger.Infof("RepositoryService SyncMerkleSubtree matched repo=%s", repositoryLogName(req.Msg.GetRepository()))
|
||||
return connect.NewResponse(&aiserverv1.SyncMerkleSubtreeResponse{
|
||||
Result: &aiserverv1.SyncMerkleSubtreeResponse_Match{Match: true},
|
||||
}), nil
|
||||
}
|
||||
|
||||
func (service *Service) FastUpdateFileV2(_ context.Context, req *connect.Request[aiserverv1.FastUpdateFileV2Request]) (*connect.Response[aiserverv1.FastUpdateFileV2Response], error) {
|
||||
logger.Infof("RepositoryService FastUpdateFileV2 accepted codebase_id=%s update_type=%s file_updates=%d", strings.TrimSpace(req.Msg.GetCodebaseId()), req.Msg.GetUpdateType().String(), len(req.Msg.GetFileUpdates()))
|
||||
return connect.NewResponse(&aiserverv1.FastUpdateFileV2Response{
|
||||
Status: aiserverv1.FastUpdateFileV2Response_STATUS_SUCCESS,
|
||||
}), nil
|
||||
}
|
||||
|
||||
func (service *Service) FastUpdateFile(_ context.Context, req *connect.Request[aiserverv1.FastUpdateFileRequest]) (*connect.Response[aiserverv1.FastUpdateFileResponse], error) {
|
||||
logger.Infof("RepositoryService FastUpdateFile accepted repo=%s update_type=%s", repositoryLogName(req.Msg.GetRepository()), req.Msg.GetUpdateType().String())
|
||||
return connect.NewResponse(&aiserverv1.FastUpdateFileResponse{
|
||||
Status: aiserverv1.FastUpdateFileResponse_STATUS_SUCCESS,
|
||||
}), nil
|
||||
}
|
||||
|
||||
func (service *Service) EnsureIndexCreated(_ context.Context, req *connect.Request[aiserverv1.EnsureIndexCreatedRequest]) (*connect.Response[aiserverv1.EnsureIndexCreatedResponse], error) {
|
||||
repo := req.Msg.GetRepository()
|
||||
codebaseID := repositoryCodebaseID(repo)
|
||||
if err := service.ensureRepositoryCodebase(codebaseID, repo); err != nil {
|
||||
return nil, connect.NewError(connect.CodeInternal, err)
|
||||
}
|
||||
logger.Infof("RepositoryService EnsureIndexCreated ready codebase_id=%s repo=%s", codebaseID, repositoryLogName(repo))
|
||||
return connect.NewResponse(&aiserverv1.EnsureIndexCreatedResponse{}), nil
|
||||
}
|
||||
|
||||
func (service *Service) GetCopyStatus(_ context.Context, req *connect.Request[aiserverv1.GetCopyStatusRequest]) (*connect.Response[aiserverv1.GetCopyStatusResponse], error) {
|
||||
store, err := service.requireCodebaseIndexStore()
|
||||
if err != nil {
|
||||
return nil, connect.NewError(connect.CodeInternal, err)
|
||||
}
|
||||
if codebaseID := strings.TrimSpace(req.Msg.GetCodebaseId()); codebaseID != "" {
|
||||
if err := store.MarkRepositoryCopyComplete(codebaseID, req.Msg.GetCopyTaskHandle()); err != nil {
|
||||
return nil, connect.NewError(connect.CodeInternal, err)
|
||||
}
|
||||
}
|
||||
completed := aiserverv1.GetCopyStatusResponse_COMPLETED_STATUS_UP_TO_DATE
|
||||
logger.Infof("RepositoryService GetCopyStatus completed codebase_id=%s copy_task_handle=%s", strings.TrimSpace(req.Msg.GetCodebaseId()), strings.TrimSpace(req.Msg.GetCopyTaskHandle()))
|
||||
return connect.NewResponse(&aiserverv1.GetCopyStatusResponse{
|
||||
Phase: aiserverv1.GetCopyStatusResponse_PHASE_COMPLETED,
|
||||
PercentDone: 1,
|
||||
CompletedStatus: &completed,
|
||||
}), nil
|
||||
}
|
||||
|
||||
func (service *Service) GetUploadLimits(_ context.Context, req *connect.Request[aiserverv1.GetUploadLimitsRequest]) (*connect.Response[aiserverv1.GetUploadLimitsResponse], error) {
|
||||
logger.Infof("RepositoryService GetUploadLimits allowed repo=%s", repositoryLogName(req.Msg.GetRepository()))
|
||||
return connect.NewResponse(&aiserverv1.GetUploadLimitsResponse{
|
||||
SoftLimit: 1_000_000,
|
||||
HardLimit: 1_000_000,
|
||||
}), nil
|
||||
}
|
||||
|
||||
func (service *Service) GetNumFilesToSend(_ context.Context, req *connect.Request[aiserverv1.GetNumFilesToSendRequest]) (*connect.Response[aiserverv1.GetNumFilesToSendResponse], error) {
|
||||
logger.Infof("RepositoryService GetNumFilesToSend none repo=%s", repositoryLogName(req.Msg.GetRepository()))
|
||||
return connect.NewResponse(&aiserverv1.GetNumFilesToSendResponse{NumFiles: 0}), nil
|
||||
}
|
||||
|
||||
func (service *Service) GetAvailableChunkingStrategies(_ context.Context, req *connect.Request[aiserverv1.GetAvailableChunkingStrategiesRequest]) (*connect.Response[aiserverv1.GetAvailableChunkingStrategiesResponse], error) {
|
||||
logger.Infof("RepositoryService GetAvailableChunkingStrategies default repo=%s", repositoryLogName(req.Msg.GetRepository()))
|
||||
return connect.NewResponse(&aiserverv1.GetAvailableChunkingStrategiesResponse{
|
||||
ChunkingStrategies: []aiserverv1.ChunkingStrategy{aiserverv1.ChunkingStrategy_CHUNKING_STRATEGY_DEFAULT},
|
||||
}), nil
|
||||
}
|
||||
|
||||
func (service *Service) GetHighLevelFolderDescription(_ context.Context, req *connect.Request[aiserverv1.GetHighLevelFolderDescriptionRequest]) (*connect.Response[aiserverv1.GetHighLevelFolderDescriptionResponse], error) {
|
||||
logger.Infof("RepositoryService GetHighLevelFolderDescription empty workspace_root=%s top_level_count=%d", strings.TrimSpace(req.Msg.GetWorkspaceRootPath()), len(req.Msg.GetTopLevelRelativeWorkspacePaths()))
|
||||
return connect.NewResponse(&aiserverv1.GetHighLevelFolderDescriptionResponse{}), nil
|
||||
}
|
||||
|
||||
func (service *Service) RepositoryStatus(_ context.Context, req *connect.Request[aiserverv1.RepositoryStatusRequest]) (*connect.Response[aiserverv1.RepositoryStatusResponse], error) {
|
||||
logger.Infof("RepositoryService RepositoryStatus synced repo=%s", repositoryLogName(req.Msg.GetRepository()))
|
||||
return connect.NewResponse(repositoryStatusSyncedResponse()), nil
|
||||
}
|
||||
|
||||
func (service *Service) BatchRepositoryStatus(_ context.Context, req *connect.Request[aiserverv1.BatchRepositoryStatusRequest]) (*connect.Response[aiserverv1.BatchRepositoryStatusResponse], error) {
|
||||
requests := req.Msg.GetRequests()
|
||||
responses := make([]*aiserverv1.RepositoryStatusResponse, 0, len(requests))
|
||||
for _, item := range requests {
|
||||
logger.Infof("RepositoryService BatchRepositoryStatus synced repo=%s", repositoryLogName(item.GetRepository()))
|
||||
responses = append(responses, repositoryStatusSyncedResponse())
|
||||
}
|
||||
return connect.NewResponse(&aiserverv1.BatchRepositoryStatusResponse{Responses: responses}), nil
|
||||
}
|
||||
|
||||
func repositoryStatusSyncedResponse() *aiserverv1.RepositoryStatusResponse {
|
||||
owner := true
|
||||
return &aiserverv1.RepositoryStatusResponse{
|
||||
IsOwner: &owner,
|
||||
Status: &aiserverv1.RepositoryStatusResponse_Synced_{
|
||||
Synced: &aiserverv1.RepositoryStatusResponse_Synced{
|
||||
Branch: "local",
|
||||
Commit: "local",
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func (service *Service) ensureRepositoryCodebase(codebaseID string, repo *aiserverv1.RepositoryInfo) error {
|
||||
store, err := service.requireCodebaseIndexStore()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
_, err = store.EnsureRepositoryIndexed(RepositoryIndexStateRecord{
|
||||
CodebaseID: strings.TrimSpace(codebaseID),
|
||||
Repo: repositoryRepoName(repo),
|
||||
TargetDir: strings.TrimSpace(repo.GetRelativeWorkspacePath()),
|
||||
Status: repositoryIndexStatusReady,
|
||||
LastHandshakeAt: time.Now().UTC(),
|
||||
})
|
||||
return err
|
||||
}
|
||||
|
||||
func repositoryPartialPathMatchResults(count int) []*aiserverv1.SyncMerkleSubtreeV2Response_PartialPathResult {
|
||||
if count <= 0 {
|
||||
return nil
|
||||
}
|
||||
results := make([]*aiserverv1.SyncMerkleSubtreeV2Response_PartialPathResult, 0, count)
|
||||
for index := 0; index < count; index++ {
|
||||
results = append(results, &aiserverv1.SyncMerkleSubtreeV2Response_PartialPathResult{
|
||||
Result: &aiserverv1.SyncMerkleSubtreeV2Response_PartialPathResult_Match{Match: true},
|
||||
})
|
||||
}
|
||||
return results
|
||||
}
|
||||
|
||||
func repositoryCodebaseID(repo *aiserverv1.RepositoryInfo, extras ...string) string {
|
||||
parts := repositoryIdentityParts(repo)
|
||||
for _, extra := range extras {
|
||||
trimmed := strings.TrimSpace(extra)
|
||||
if trimmed != "" {
|
||||
parts = append(parts, trimmed)
|
||||
}
|
||||
}
|
||||
if len(parts) == 0 {
|
||||
return "cb_local"
|
||||
}
|
||||
sum := sha256.Sum256([]byte(strings.Join(parts, "\x00")))
|
||||
return "cb_" + hex.EncodeToString(sum[:])[:24]
|
||||
}
|
||||
|
||||
func repositoryIdentityParts(repo *aiserverv1.RepositoryInfo) []string {
|
||||
if repo == nil {
|
||||
return nil
|
||||
}
|
||||
parts := make([]string, 0, 6+len(repo.GetRemoteUrls()))
|
||||
for _, value := range []string{
|
||||
repo.GetWorkspaceUri(),
|
||||
repo.GetRelativeWorkspacePath(),
|
||||
repo.GetRepoOwner(),
|
||||
repo.GetRepoName(),
|
||||
} {
|
||||
trimmed := strings.TrimSpace(value)
|
||||
if trimmed != "" {
|
||||
parts = append(parts, trimmed)
|
||||
}
|
||||
}
|
||||
for _, remoteURL := range repo.GetRemoteUrls() {
|
||||
trimmed := strings.TrimSpace(remoteURL)
|
||||
if trimmed != "" {
|
||||
parts = append(parts, trimmed)
|
||||
}
|
||||
}
|
||||
if repo.GetOrthogonalTransformSeed() != 0 {
|
||||
parts = append(parts, fmt.Sprintf("seed:%f", repo.GetOrthogonalTransformSeed()))
|
||||
}
|
||||
return parts
|
||||
}
|
||||
|
||||
func repositoryRepoName(repo *aiserverv1.RepositoryInfo) string {
|
||||
if repo == nil {
|
||||
return "local"
|
||||
}
|
||||
if owner := strings.TrimSpace(repo.GetRepoOwner()); owner != "" {
|
||||
if name := strings.TrimSpace(repo.GetRepoName()); name != "" {
|
||||
return owner + "/" + name
|
||||
}
|
||||
}
|
||||
for _, value := range []string{repo.GetRepoName(), repo.GetRelativeWorkspacePath(), repo.GetWorkspaceUri()} {
|
||||
trimmed := strings.TrimSpace(value)
|
||||
if trimmed != "" {
|
||||
return filepath.Base(trimmed)
|
||||
}
|
||||
}
|
||||
return "local"
|
||||
}
|
||||
|
||||
func repositoryLogName(repo *aiserverv1.RepositoryInfo) string {
|
||||
if repo == nil {
|
||||
return ""
|
||||
}
|
||||
return strings.TrimSpace(strings.Join(repositoryIdentityParts(repo), "|"))
|
||||
}
|
||||
@@ -0,0 +1,395 @@
|
||||
package forwarder
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
|
||||
"google.golang.org/protobuf/proto"
|
||||
|
||||
"cursor/gen/agentv1"
|
||||
)
|
||||
|
||||
func normalizeRequestContextForStorage(requestContext *agentv1.RequestContext) *agentv1.RequestContext {
|
||||
if requestContext == nil {
|
||||
return nil
|
||||
}
|
||||
return normalizeRequestContextForStorageMode(requestContext, true)
|
||||
}
|
||||
|
||||
func normalizeRequestContextForStorageMode(requestContext *agentv1.RequestContext, includeStatic bool) *agentv1.RequestContext {
|
||||
if requestContext == nil {
|
||||
return nil
|
||||
}
|
||||
if !includeStatic {
|
||||
return normalizeRealtimeRequestContextForStorage(requestContext)
|
||||
}
|
||||
|
||||
cloned, ok := proto.Clone(requestContext).(*agentv1.RequestContext)
|
||||
if !ok || cloned == nil {
|
||||
return requestContext
|
||||
}
|
||||
|
||||
descriptors := collectSkillDescriptors(cloned)
|
||||
cloned.Rules = filterNonSkillRules(cloned.GetRules())
|
||||
cloned.AgentSkills = nil
|
||||
descriptors = guardSkillDescriptors(descriptors)
|
||||
if len(descriptors) > 0 {
|
||||
cloned.SkillOptions = &agentv1.SkillOptions{SkillDescriptors: descriptors}
|
||||
} else {
|
||||
cloned.SkillOptions = nil
|
||||
}
|
||||
guardRequestContextForStorage(cloned)
|
||||
return cloned
|
||||
}
|
||||
|
||||
func normalizeRealtimeRequestContextForStorage(requestContext *agentv1.RequestContext) *agentv1.RequestContext {
|
||||
if requestContext == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
normalized := &agentv1.RequestContext{}
|
||||
normalized.Rules = guardCursorRules(filterNonSkillRules(requestContext.GetRules()))
|
||||
if fileContents := normalizeRealtimeFileContents(requestContext.GetFileContents()); len(fileContents) > 0 {
|
||||
normalized.FileContents = fileContents
|
||||
}
|
||||
if summary := strings.TrimSpace(requestContext.GetUserIntentSummary()); summary != "" {
|
||||
normalized.UserIntentSummary = stringPtr(truncatePromptGuardText("request_context.user_intent_summary", summary, promptGuardRealtimeTextChars))
|
||||
}
|
||||
if hooks := strings.TrimSpace(requestContext.GetHooksAdditionalContext()); hooks != "" {
|
||||
normalized.HooksAdditionalContext = stringPtr(truncatePromptGuardText("request_context.hooks_additional_context", hooks, promptGuardRealtimeTextChars))
|
||||
}
|
||||
if commit := strings.TrimSpace(requestContext.GetCommitAttributionMessage()); commit != "" {
|
||||
normalized.CommitAttributionMessage = stringPtr(truncatePromptGuardText("request_context.commit_attribution_message", commit, promptGuardRealtimeTextChars))
|
||||
}
|
||||
if pr := strings.TrimSpace(requestContext.GetPrAttributionMessage()); pr != "" {
|
||||
normalized.PrAttributionMessage = stringPtr(truncatePromptGuardText("request_context.pr_attribution_message", pr, promptGuardRealtimeTextChars))
|
||||
}
|
||||
if !hasRealtimeRequestContextContent(normalized) {
|
||||
return nil
|
||||
}
|
||||
return normalized
|
||||
}
|
||||
|
||||
func normalizeRealtimeFileContents(fileContents map[string]string) map[string]string {
|
||||
if len(fileContents) == 0 {
|
||||
return nil
|
||||
}
|
||||
normalized := make(map[string]string, len(fileContents))
|
||||
for path, content := range fileContents {
|
||||
trimmedPath := strings.TrimSpace(path)
|
||||
if trimmedPath == "" || strings.TrimSpace(content) == "" {
|
||||
continue
|
||||
}
|
||||
normalized[trimmedPath] = content
|
||||
}
|
||||
if len(normalized) == 0 {
|
||||
return nil
|
||||
}
|
||||
return guardStringMap(normalized, "request_context.file_contents", promptGuardRequestFileChars, promptGuardRequestFilesTotalChars, promptGuardRequestFilesMaxCount)
|
||||
}
|
||||
|
||||
func hasRealtimeRequestContextContent(requestContext *agentv1.RequestContext) bool {
|
||||
if requestContext == nil {
|
||||
return false
|
||||
}
|
||||
return len(requestContext.GetRules()) > 0 ||
|
||||
len(requestContext.GetFileContents()) > 0 ||
|
||||
strings.TrimSpace(requestContext.GetUserIntentSummary()) != "" ||
|
||||
strings.TrimSpace(requestContext.GetHooksAdditionalContext()) != "" ||
|
||||
strings.TrimSpace(requestContext.GetCommitAttributionMessage()) != "" ||
|
||||
strings.TrimSpace(requestContext.GetPrAttributionMessage()) != ""
|
||||
}
|
||||
|
||||
func collectSkillDescriptors(requestContext *agentv1.RequestContext) []*agentv1.SkillDescriptor {
|
||||
if requestContext == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
orderedKeys := make([]string, 0, 16)
|
||||
descriptorsByKey := make(map[string]*agentv1.SkillDescriptor)
|
||||
addDescriptor := func(descriptor *agentv1.SkillDescriptor) {
|
||||
normalized := normalizeSkillDescriptor(descriptor)
|
||||
if normalized == nil {
|
||||
return
|
||||
}
|
||||
key := skillDescriptorKey(normalized)
|
||||
if key == "" {
|
||||
return
|
||||
}
|
||||
if existing, ok := descriptorsByKey[key]; ok {
|
||||
mergeSkillDescriptor(existing, normalized)
|
||||
return
|
||||
}
|
||||
orderedKeys = append(orderedKeys, key)
|
||||
descriptorsByKey[key] = normalized
|
||||
}
|
||||
|
||||
for _, descriptor := range requestContext.GetSkillOptions().GetSkillDescriptors() {
|
||||
addDescriptor(descriptor)
|
||||
}
|
||||
for _, skill := range requestContext.GetAgentSkills() {
|
||||
addDescriptor(skillDescriptorFromAgentSkill(skill))
|
||||
}
|
||||
for _, rule := range requestContext.GetRules() {
|
||||
if isSkillFilePath(rule.GetFullPath()) {
|
||||
addDescriptor(skillDescriptorFromCursorRule(rule))
|
||||
}
|
||||
}
|
||||
|
||||
descriptors := make([]*agentv1.SkillDescriptor, 0, len(orderedKeys))
|
||||
for _, key := range orderedKeys {
|
||||
if descriptor := descriptorsByKey[key]; descriptor != nil {
|
||||
descriptors = append(descriptors, descriptor)
|
||||
}
|
||||
}
|
||||
return descriptors
|
||||
}
|
||||
|
||||
func filterNonSkillRules(rules []*agentv1.CursorRule) []*agentv1.CursorRule {
|
||||
if len(rules) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
filtered := make([]*agentv1.CursorRule, 0, len(rules))
|
||||
for _, rule := range rules {
|
||||
if rule == nil {
|
||||
continue
|
||||
}
|
||||
if isSkillFilePath(rule.GetFullPath()) {
|
||||
continue
|
||||
}
|
||||
filtered = append(filtered, rule)
|
||||
}
|
||||
return filtered
|
||||
}
|
||||
|
||||
func skillDescriptorFromAgentSkill(skill *agentv1.AgentSkill) *agentv1.SkillDescriptor {
|
||||
if skill == nil {
|
||||
return nil
|
||||
}
|
||||
fullPath := strings.TrimSpace(skill.GetFullPath())
|
||||
if fullPath == "" {
|
||||
return nil
|
||||
}
|
||||
|
||||
parseError := strings.TrimSpace(skill.GetParseError())
|
||||
descriptor := &agentv1.SkillDescriptor{
|
||||
Name: inferSkillName(fullPath),
|
||||
Description: strings.TrimSpace(skill.GetDescription()),
|
||||
FolderPath: filepath.Dir(fullPath),
|
||||
Enabled: parseError == "",
|
||||
ReadmeFilePath: fullPath,
|
||||
PackageType: inferSkillPackageType(fullPath),
|
||||
}
|
||||
if parseError != "" {
|
||||
descriptor.ParseError = &parseError
|
||||
}
|
||||
return descriptor
|
||||
}
|
||||
|
||||
func skillDescriptorFromCursorRule(rule *agentv1.CursorRule) *agentv1.SkillDescriptor {
|
||||
if rule == nil {
|
||||
return nil
|
||||
}
|
||||
fullPath := strings.TrimSpace(rule.GetFullPath())
|
||||
if fullPath == "" || !isSkillFilePath(fullPath) {
|
||||
return nil
|
||||
}
|
||||
|
||||
description := ""
|
||||
if rule.GetType() != nil && rule.GetType().GetAgentFetched() != nil {
|
||||
description = strings.TrimSpace(rule.GetType().GetAgentFetched().GetDescription())
|
||||
}
|
||||
parseError := strings.TrimSpace(rule.GetParseError())
|
||||
descriptor := &agentv1.SkillDescriptor{
|
||||
Name: inferSkillName(fullPath),
|
||||
Description: description,
|
||||
FolderPath: filepath.Dir(fullPath),
|
||||
Enabled: parseError == "",
|
||||
ReadmeFilePath: fullPath,
|
||||
PackageType: inferSkillPackageType(fullPath),
|
||||
}
|
||||
if parseError != "" {
|
||||
descriptor.ParseError = &parseError
|
||||
}
|
||||
return descriptor
|
||||
}
|
||||
|
||||
func normalizeSkillDescriptor(descriptor *agentv1.SkillDescriptor) *agentv1.SkillDescriptor {
|
||||
if descriptor == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
cloned := proto.Clone(descriptor).(*agentv1.SkillDescriptor)
|
||||
cloned.Name = strings.TrimSpace(cloned.GetName())
|
||||
cloned.Description = strings.TrimSpace(cloned.GetDescription())
|
||||
cloned.FolderPath = strings.TrimSpace(cloned.GetFolderPath())
|
||||
cloned.ReadmeFilePath = strings.TrimSpace(cloned.GetReadmeFilePath())
|
||||
if cloned.ReadmeFilePath == "" && cloned.FolderPath != "" {
|
||||
cloned.ReadmeFilePath = filepath.Join(cloned.FolderPath, "SKILL.md")
|
||||
}
|
||||
if cloned.FolderPath == "" && cloned.ReadmeFilePath != "" {
|
||||
cloned.FolderPath = filepath.Dir(cloned.ReadmeFilePath)
|
||||
}
|
||||
if cloned.Name == "" {
|
||||
cloned.Name = inferSkillName(cloned.GetReadmeFilePath())
|
||||
}
|
||||
if cloned.PackageType == agentv1.PackageType_PACKAGE_TYPE_UNSPECIFIED {
|
||||
cloned.PackageType = inferSkillPackageType(firstNonEmpty(cloned.GetReadmeFilePath(), cloned.GetFolderPath()))
|
||||
}
|
||||
if cloned.GetReadmeFilePath() == "" || cloned.GetDescription() == "" {
|
||||
return nil
|
||||
}
|
||||
return cloned
|
||||
}
|
||||
|
||||
func mergeSkillDescriptor(target *agentv1.SkillDescriptor, candidate *agentv1.SkillDescriptor) {
|
||||
if target == nil || candidate == nil {
|
||||
return
|
||||
}
|
||||
if strings.TrimSpace(target.GetDescription()) == "" && strings.TrimSpace(candidate.GetDescription()) != "" {
|
||||
target.Description = strings.TrimSpace(candidate.GetDescription())
|
||||
}
|
||||
if strings.TrimSpace(target.GetReadmeFilePath()) == "" && strings.TrimSpace(candidate.GetReadmeFilePath()) != "" {
|
||||
target.ReadmeFilePath = strings.TrimSpace(candidate.GetReadmeFilePath())
|
||||
}
|
||||
if strings.TrimSpace(target.GetFolderPath()) == "" && strings.TrimSpace(candidate.GetFolderPath()) != "" {
|
||||
target.FolderPath = strings.TrimSpace(candidate.GetFolderPath())
|
||||
}
|
||||
if !target.GetEnabled() && candidate.GetEnabled() {
|
||||
target.Enabled = true
|
||||
}
|
||||
if target.GetPackageType() == agentv1.PackageType_PACKAGE_TYPE_UNSPECIFIED && candidate.GetPackageType() != agentv1.PackageType_PACKAGE_TYPE_UNSPECIFIED {
|
||||
target.PackageType = candidate.GetPackageType()
|
||||
}
|
||||
if strings.TrimSpace(target.GetParseError()) == "" && strings.TrimSpace(candidate.GetParseError()) != "" {
|
||||
parseError := strings.TrimSpace(candidate.GetParseError())
|
||||
target.ParseError = &parseError
|
||||
}
|
||||
}
|
||||
|
||||
func skillDescriptorKey(descriptor *agentv1.SkillDescriptor) string {
|
||||
if descriptor == nil {
|
||||
return ""
|
||||
}
|
||||
if readme := strings.TrimSpace(descriptor.GetReadmeFilePath()); readme != "" {
|
||||
return readme
|
||||
}
|
||||
if folder := strings.TrimSpace(descriptor.GetFolderPath()); folder != "" {
|
||||
return folder
|
||||
}
|
||||
return strings.TrimSpace(descriptor.GetName())
|
||||
}
|
||||
|
||||
func buildSkillDiscoveryMessage(requestContext *agentv1.RequestContext) string {
|
||||
normalized := normalizeRequestContextForStorageMode(requestContext, true)
|
||||
if normalized == nil || normalized.GetSkillOptions() == nil || len(normalized.GetSkillOptions().GetSkillDescriptors()) == 0 {
|
||||
return ""
|
||||
}
|
||||
|
||||
lines := []string{
|
||||
"<agent_skills>",
|
||||
"When users ask you to perform tasks, check if any of the available skills below can help complete the task more effectively. Skills provide specialized capabilities and domain knowledge. To use a skill, read the skill file at the provided absolute path using the Read tool, then follow the instructions within. When a skill is relevant, read and follow it IMMEDIATELY as your first action. NEVER just announce or mention a skill without actually reading and following it. Only use skills listed below.",
|
||||
"",
|
||||
`<available_skills description="Skills the agent can use. Use the Read tool with the provided absolute path to fetch full contents.">`,
|
||||
}
|
||||
for _, descriptor := range normalized.GetSkillOptions().GetSkillDescriptors() {
|
||||
if descriptor == nil {
|
||||
continue
|
||||
}
|
||||
fullPath := strings.TrimSpace(descriptor.GetReadmeFilePath())
|
||||
description := strings.TrimSpace(descriptor.GetDescription())
|
||||
if fullPath == "" || description == "" {
|
||||
continue
|
||||
}
|
||||
lines = append(lines, fmt.Sprintf(`<agent_skill fullPath="%s">%s</agent_skill>`, escapeSkillPromptXML(fullPath), escapeSkillPromptXML(description)))
|
||||
}
|
||||
lines = append(lines, "</available_skills>", "</agent_skills>")
|
||||
if len(lines) == 6 {
|
||||
return ""
|
||||
}
|
||||
return strings.Join(lines, "\n\n")
|
||||
}
|
||||
|
||||
func collectMCPToolServers(requestContext *agentv1.RequestContext) map[string]string {
|
||||
if requestContext == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
servers := make(map[string]string)
|
||||
addDescriptors := func(descriptors []*agentv1.McpDescriptor) {
|
||||
for _, descriptor := range descriptors {
|
||||
if descriptor == nil {
|
||||
continue
|
||||
}
|
||||
serverIdentifier := firstNonEmpty(descriptor.GetServerIdentifier(), descriptor.GetServerName())
|
||||
if strings.TrimSpace(serverIdentifier) == "" {
|
||||
continue
|
||||
}
|
||||
for _, tool := range descriptor.GetTools() {
|
||||
if tool == nil {
|
||||
continue
|
||||
}
|
||||
toolName := strings.TrimSpace(tool.GetToolName())
|
||||
if toolName == "" {
|
||||
continue
|
||||
}
|
||||
if _, exists := servers[toolName]; exists {
|
||||
continue
|
||||
}
|
||||
servers[toolName] = strings.TrimSpace(serverIdentifier)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
addDescriptors(requestContext.GetMcpFileSystemOptions().GetMcpDescriptors())
|
||||
addDescriptors(requestContext.GetMcpMetaToolOptions().GetMcpDescriptors())
|
||||
if len(servers) == 0 {
|
||||
return nil
|
||||
}
|
||||
return servers
|
||||
}
|
||||
|
||||
func escapeSkillPromptXML(value string) string {
|
||||
replacer := strings.NewReplacer(
|
||||
"&", "&",
|
||||
`"`, """,
|
||||
"<", "<",
|
||||
">", ">",
|
||||
)
|
||||
return replacer.Replace(strings.TrimSpace(value))
|
||||
}
|
||||
|
||||
func inferSkillName(path string) string {
|
||||
trimmed := strings.TrimSpace(path)
|
||||
if trimmed == "" {
|
||||
return ""
|
||||
}
|
||||
if strings.EqualFold(filepath.Base(trimmed), "SKILL.md") {
|
||||
return strings.TrimSpace(filepath.Base(filepath.Dir(trimmed)))
|
||||
}
|
||||
return strings.TrimSpace(strings.TrimSuffix(filepath.Base(trimmed), filepath.Ext(trimmed)))
|
||||
}
|
||||
|
||||
func inferSkillPackageType(path string) agentv1.PackageType {
|
||||
normalized := strings.ToLower(strings.TrimSpace(path))
|
||||
switch {
|
||||
case strings.Contains(normalized, "/.agents/skills/"):
|
||||
return agentv1.PackageType_PACKAGE_TYPE_CURSOR_PROJECT
|
||||
case strings.Contains(normalized, "/.cursor/skills/"),
|
||||
strings.Contains(normalized, "/.cursor/skills-cursor/"),
|
||||
strings.Contains(normalized, "/.codex/skills/"):
|
||||
return agentv1.PackageType_PACKAGE_TYPE_CURSOR_PERSONAL
|
||||
default:
|
||||
return agentv1.PackageType_PACKAGE_TYPE_UNSPECIFIED
|
||||
}
|
||||
}
|
||||
|
||||
func isSkillFilePath(path string) bool {
|
||||
trimmed := strings.TrimSpace(path)
|
||||
if trimmed == "" {
|
||||
return false
|
||||
}
|
||||
return strings.EqualFold(filepath.Base(trimmed), "SKILL.md")
|
||||
}
|
||||
@@ -0,0 +1,299 @@
|
||||
package forwarder
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"google.golang.org/protobuf/encoding/protojson"
|
||||
|
||||
"cursor/gen/agentv1"
|
||||
)
|
||||
|
||||
type runRewindDecision struct {
|
||||
Evaluated bool
|
||||
Apply bool
|
||||
Reason string
|
||||
SkipReason string
|
||||
IncomingMessageID string
|
||||
HasClientTurnCount bool
|
||||
ClientTurnCount int
|
||||
ServerTailTurnSeq int64
|
||||
ServerNextTurnSeq int64
|
||||
TargetTurnSeq int64
|
||||
TargetEntrySeq int64
|
||||
TargetRequestID string
|
||||
MatchCount int
|
||||
DroppedEntryCount int
|
||||
DroppedTurnCount int
|
||||
DroppedSeqStart int64
|
||||
DroppedSeqEnd int64
|
||||
PrefixEntries []HistoryEntry
|
||||
}
|
||||
|
||||
type runRewindMatch struct {
|
||||
Entry HistoryEntry
|
||||
}
|
||||
|
||||
func (service *Service) decideRunRewind(intent InboundIntent, conversation *ConversationFile) runRewindDecision {
|
||||
decision := runRewindDecision{ClientTurnCount: -1}
|
||||
if !shouldEvaluateRunRewind(intent) {
|
||||
return decision
|
||||
}
|
||||
decision.Evaluated = true
|
||||
decision.IncomingMessageID = strings.TrimSpace(intent.UserMessage.GetMessageId())
|
||||
decision.HasClientTurnCount = intent.ConversationState != nil
|
||||
if decision.HasClientTurnCount {
|
||||
decision.ClientTurnCount = len(intent.ConversationState.GetTurns())
|
||||
}
|
||||
if conversation != nil {
|
||||
decision.ServerTailTurnSeq = maxHistoryTurnSeq(conversation.Entries)
|
||||
decision.ServerNextTurnSeq = conversation.NextTurnSeq
|
||||
}
|
||||
if decision.IncomingMessageID == "" {
|
||||
decision.SkipReason = "missing_message_id"
|
||||
return decision
|
||||
}
|
||||
if conversation == nil || len(conversation.Entries) == 0 {
|
||||
decision.SkipReason = "message_id_not_found"
|
||||
return decision
|
||||
}
|
||||
|
||||
matches := findUserMessageEntriesByMessageID(conversation.Entries, decision.IncomingMessageID)
|
||||
decision.MatchCount = len(matches)
|
||||
if len(matches) == 0 {
|
||||
decision.SkipReason = "message_id_not_found"
|
||||
return decision
|
||||
}
|
||||
selected, selectReason := selectRunRewindMatch(matches, decision.ClientTurnCount, decision.HasClientTurnCount)
|
||||
decision.TargetTurnSeq = selected.Entry.TurnSeq
|
||||
decision.TargetEntrySeq = selected.Entry.Seq
|
||||
decision.TargetRequestID = strings.TrimSpace(selected.Entry.RequestID)
|
||||
if decision.TargetTurnSeq <= 0 {
|
||||
decision.SkipReason = "target_turn_seq_missing"
|
||||
return decision
|
||||
}
|
||||
|
||||
serverTailBeyondTarget := decision.ServerTailTurnSeq > decision.TargetTurnSeq
|
||||
clientBehindServerTail := decision.HasClientTurnCount && int64(decision.ClientTurnCount) < decision.ServerTailTurnSeq
|
||||
if !serverTailBeyondTarget && !clientBehindServerTail {
|
||||
decision.SkipReason = "message_id_at_active_tail"
|
||||
return decision
|
||||
}
|
||||
|
||||
decision.Apply = true
|
||||
decision.Reason = selectReason
|
||||
decision.PrefixEntries = prefixEntriesBeforeTurn(conversation.Entries, decision.TargetTurnSeq)
|
||||
decision.DroppedEntryCount, decision.DroppedTurnCount, decision.DroppedSeqStart, decision.DroppedSeqEnd = droppedEntryStats(conversation.Entries, decision.TargetTurnSeq)
|
||||
return decision
|
||||
}
|
||||
|
||||
func shouldEvaluateRunRewind(intent InboundIntent) bool {
|
||||
if strings.TrimSpace(intent.Kind) != "run" || intent.Prewarm || intent.UserMessage == nil {
|
||||
return false
|
||||
}
|
||||
return conversationActionCase(intent.ClientMessage) == "user_message_action"
|
||||
}
|
||||
|
||||
func findUserMessageEntriesByMessageID(entries []HistoryEntry, messageID string) []runRewindMatch {
|
||||
messageID = strings.TrimSpace(messageID)
|
||||
if messageID == "" || len(entries) == 0 {
|
||||
return nil
|
||||
}
|
||||
matches := make([]runRewindMatch, 0, 1)
|
||||
for _, entry := range entries {
|
||||
if strings.TrimSpace(entry.Kind) != "user_message" || len(entry.Payload) == 0 {
|
||||
continue
|
||||
}
|
||||
userMessage := &agentv1.UserMessage{}
|
||||
if err := protojson.Unmarshal(entry.Payload, userMessage); err != nil {
|
||||
continue
|
||||
}
|
||||
if strings.TrimSpace(userMessage.GetMessageId()) != messageID {
|
||||
continue
|
||||
}
|
||||
matches = append(matches, runRewindMatch{Entry: entry})
|
||||
}
|
||||
return matches
|
||||
}
|
||||
|
||||
func selectRunRewindMatch(matches []runRewindMatch, clientTurnCount int, hasClientTurnCount bool) (runRewindMatch, string) {
|
||||
if len(matches) == 0 {
|
||||
return runRewindMatch{}, "no_match"
|
||||
}
|
||||
if hasClientTurnCount && clientTurnCount >= 0 {
|
||||
targetTurnSeq := int64(clientTurnCount) + 1
|
||||
for _, match := range matches {
|
||||
if match.Entry.TurnSeq == targetTurnSeq {
|
||||
return match, "client_turn_count_aligned"
|
||||
}
|
||||
}
|
||||
var candidate *runRewindMatch
|
||||
for index := range matches {
|
||||
match := matches[index]
|
||||
if match.Entry.TurnSeq <= int64(clientTurnCount) {
|
||||
continue
|
||||
}
|
||||
if candidate == nil || earlierHistoryEntry(match.Entry, candidate.Entry) {
|
||||
candidate = &match
|
||||
}
|
||||
}
|
||||
if candidate != nil {
|
||||
return *candidate, "first_match_after_client_turn_count"
|
||||
}
|
||||
}
|
||||
selected := matches[0]
|
||||
for _, match := range matches[1:] {
|
||||
if earlierHistoryEntry(match.Entry, selected.Entry) {
|
||||
selected = match
|
||||
}
|
||||
}
|
||||
return selected, "earliest_message_id_match"
|
||||
}
|
||||
|
||||
func earlierHistoryEntry(left HistoryEntry, right HistoryEntry) bool {
|
||||
if left.TurnSeq != right.TurnSeq {
|
||||
return left.TurnSeq < right.TurnSeq
|
||||
}
|
||||
if left.Seq != right.Seq {
|
||||
return left.Seq < right.Seq
|
||||
}
|
||||
return left.CreatedAt.Before(right.CreatedAt)
|
||||
}
|
||||
|
||||
func maxHistoryTurnSeq(entries []HistoryEntry) int64 {
|
||||
var maxTurnSeq int64
|
||||
for _, entry := range entries {
|
||||
if entry.TurnSeq > maxTurnSeq {
|
||||
maxTurnSeq = entry.TurnSeq
|
||||
}
|
||||
}
|
||||
return maxTurnSeq
|
||||
}
|
||||
|
||||
func prefixEntriesBeforeTurn(entries []HistoryEntry, targetTurnSeq int64) []HistoryEntry {
|
||||
if len(entries) == 0 {
|
||||
return nil
|
||||
}
|
||||
prefix := make([]HistoryEntry, 0, len(entries))
|
||||
for _, entry := range entries {
|
||||
if entry.TurnSeq < targetTurnSeq {
|
||||
prefix = append(prefix, entry)
|
||||
}
|
||||
}
|
||||
return prefix
|
||||
}
|
||||
|
||||
func droppedEntryStats(entries []HistoryEntry, targetTurnSeq int64) (int, int, int64, int64) {
|
||||
if len(entries) == 0 {
|
||||
return 0, 0, 0, 0
|
||||
}
|
||||
droppedTurns := make(map[int64]struct{})
|
||||
var droppedEntries int
|
||||
var seqStart int64
|
||||
var seqEnd int64
|
||||
for _, entry := range entries {
|
||||
if entry.TurnSeq < targetTurnSeq {
|
||||
continue
|
||||
}
|
||||
droppedEntries++
|
||||
if entry.TurnSeq > 0 {
|
||||
droppedTurns[entry.TurnSeq] = struct{}{}
|
||||
}
|
||||
if entry.Seq > 0 && (seqStart == 0 || entry.Seq < seqStart) {
|
||||
seqStart = entry.Seq
|
||||
}
|
||||
if entry.Seq > seqEnd {
|
||||
seqEnd = entry.Seq
|
||||
}
|
||||
}
|
||||
return droppedEntries, len(droppedTurns), seqStart, seqEnd
|
||||
}
|
||||
|
||||
func appendReplacementRunEntries(prefix []HistoryEntry, entries []HistoryEntry) []HistoryEntry {
|
||||
replacement := make([]HistoryEntry, 0, len(prefix)+len(entries))
|
||||
replacement = append(replacement, prefix...)
|
||||
replacement = append(replacement, entries...)
|
||||
return replacement
|
||||
}
|
||||
|
||||
func (service *Service) applyRunRewindToConversation(conversation *ConversationFile, decision runRewindDecision, entries []HistoryEntry, intent InboundIntent, turnSeq int64) {
|
||||
if conversation == nil || !decision.Apply {
|
||||
return
|
||||
}
|
||||
conversation.Entries = nil
|
||||
conversation.NextEntrySeq = 1
|
||||
conversation.NextTurnSeq = 1
|
||||
appendEntriesInPlace(conversation, appendReplacementRunEntries(decision.PrefixEntries, entries))
|
||||
applyRunRewindConversationState(conversation, intent, turnSeq)
|
||||
deriveConversationLoopState(conversation)
|
||||
}
|
||||
|
||||
func applyRunRewindConversationState(conversation *ConversationFile, intent InboundIntent, turnSeq int64) {
|
||||
if conversation == nil {
|
||||
return
|
||||
}
|
||||
conversation.TokenDetailsUsedTokens = 0
|
||||
if conversation.TokenDetailsMaxTokens == 0 {
|
||||
conversation.TokenDetailsMaxTokens = projectedConversationMaxTokens
|
||||
}
|
||||
clearConversationAutoCompactionState(conversation)
|
||||
conversation.LatestRequestPrefix = nil
|
||||
conversation.LastProviderCall = nil
|
||||
conversation.CurrentLoopID = fmt.Sprintf("%d:%s", turnSeq, strings.TrimSpace(intent.RequestID))
|
||||
conversation.CurrentLoopStatus = "running"
|
||||
conversation.CurrentRequestID = strings.TrimSpace(intent.RequestID)
|
||||
conversation.CurrentTurnSeq = turnSeq
|
||||
}
|
||||
|
||||
func applyRunRewindMetadata(conversation *ConversationFile, source *ConversationFile, intent InboundIntent, turnSeq int64) {
|
||||
if conversation == nil {
|
||||
return
|
||||
}
|
||||
if source != nil {
|
||||
if strings.TrimSpace(source.ConversationID) != "" {
|
||||
conversation.ConversationID = strings.TrimSpace(source.ConversationID)
|
||||
}
|
||||
if strings.TrimSpace(source.RootConversationID) != "" {
|
||||
conversation.RootConversationID = strings.TrimSpace(source.RootConversationID)
|
||||
}
|
||||
conversation.ParentConversationID = strings.TrimSpace(source.ParentConversationID)
|
||||
conversation.ParentToolCallID = strings.TrimSpace(source.ParentToolCallID)
|
||||
conversation.SubagentTypeName = strings.TrimSpace(source.SubagentTypeName)
|
||||
if strings.TrimSpace(source.Mode) != "" {
|
||||
conversation.Mode = strings.TrimSpace(source.Mode)
|
||||
}
|
||||
if source.TokenDetailsMaxTokens > 0 {
|
||||
conversation.TokenDetailsMaxTokens = source.TokenDetailsMaxTokens
|
||||
}
|
||||
}
|
||||
applyRunRewindConversationState(conversation, intent, turnSeq)
|
||||
}
|
||||
|
||||
func (service *Service) logRunRewindDecision(requestID string, conversationID string, eventName string, decision runRewindDecision) {
|
||||
if service == nil || !decision.Evaluated {
|
||||
return
|
||||
}
|
||||
fields := map[string]any{
|
||||
"message_id": decision.IncomingMessageID,
|
||||
"apply": decision.Apply,
|
||||
"reason": decision.Reason,
|
||||
"skip_reason": decision.SkipReason,
|
||||
"target_turn_seq": decision.TargetTurnSeq,
|
||||
"target_entry_seq": decision.TargetEntrySeq,
|
||||
"target_request_id": decision.TargetRequestID,
|
||||
"server_tail_turn_seq": decision.ServerTailTurnSeq,
|
||||
"server_next_turn_seq": decision.ServerNextTurnSeq,
|
||||
"match_count": decision.MatchCount,
|
||||
"dropped_entry_count": decision.DroppedEntryCount,
|
||||
"dropped_turn_count": decision.DroppedTurnCount,
|
||||
"dropped_seq_start": decision.DroppedSeqStart,
|
||||
"dropped_seq_end": decision.DroppedSeqEnd,
|
||||
}
|
||||
if decision.HasClientTurnCount {
|
||||
fields["client_turn_count"] = decision.ClientTurnCount
|
||||
} else {
|
||||
fields["client_turn_count"] = nil
|
||||
}
|
||||
service.debug.LogRuntime(context.Background(), requestID, conversationID, eventName, fields)
|
||||
}
|
||||
@@ -0,0 +1,298 @@
|
||||
package forwarder
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"cursor/gen/agentv1"
|
||||
modeladapter "cursor/internal/backend/agent/model"
|
||||
)
|
||||
|
||||
func newRuntimeConversation(conversationID string, mode agentv1.AgentMode) (*ConversationFile, error) {
|
||||
now := time.Now().UTC()
|
||||
normalizedID := strings.TrimSpace(conversationID)
|
||||
alias, err := modeAlias(mode)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &ConversationFile{
|
||||
ConversationID: normalizedID,
|
||||
RootConversationID: normalizedID,
|
||||
Mode: alias,
|
||||
CreatedAt: now,
|
||||
UpdatedAt: now,
|
||||
NextTurnSeq: 1,
|
||||
NextEntrySeq: 1,
|
||||
Entries: make([]HistoryEntry, 0, 16),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (service *Service) bootstrapRuntimeConversation(intent InboundIntent) (*ConversationFile, agentv1.AgentMode, int64, []HistoryEntry, error) {
|
||||
if service == nil {
|
||||
return nil, agentv1.AgentMode_AGENT_MODE_AGENT, 0, nil, fmt.Errorf("forwarder service is nil")
|
||||
}
|
||||
contextWindowTokens := service.resolveContextWindowTokens(intent.ModelID)
|
||||
var conversation *ConversationFile
|
||||
var err error
|
||||
if service.store != nil {
|
||||
conversation, err = service.store.LoadConversation(intent.ConversationID)
|
||||
if err != nil {
|
||||
return nil, agentv1.AgentMode_AGENT_MODE_AGENT, 0, nil, err
|
||||
}
|
||||
}
|
||||
if conversation == nil {
|
||||
conversation, err = newRuntimeConversation(intent.ConversationID, intent.Mode)
|
||||
if err != nil {
|
||||
return nil, agentv1.AgentMode_AGENT_MODE_AGENT, 0, nil, err
|
||||
}
|
||||
}
|
||||
importedEntries := []HistoryEntry(nil)
|
||||
if len(conversation.Entries) == 0 && intent.ConversationState != nil {
|
||||
importedEntries, err = service.importConversationState(conversation, intent.ConversationState)
|
||||
if err != nil {
|
||||
return nil, agentv1.AgentMode_AGENT_MODE_AGENT, 0, nil, err
|
||||
}
|
||||
}
|
||||
if strings.TrimSpace(conversation.ConversationID) == "" {
|
||||
conversation.ConversationID = strings.TrimSpace(intent.ConversationID)
|
||||
}
|
||||
if strings.TrimSpace(conversation.RootConversationID) == "" {
|
||||
conversation.RootConversationID = strings.TrimSpace(conversation.ConversationID)
|
||||
}
|
||||
if intent.HasExplicitMode {
|
||||
alias, err := modeAlias(intent.Mode)
|
||||
if err != nil {
|
||||
return nil, agentv1.AgentMode_AGENT_MODE_AGENT, 0, nil, err
|
||||
}
|
||||
conversation.Mode = alias
|
||||
} else if strings.TrimSpace(conversation.Mode) == "" {
|
||||
alias, err := modeAlias(intent.Mode)
|
||||
if err != nil {
|
||||
return nil, agentv1.AgentMode_AGENT_MODE_AGENT, 0, nil, err
|
||||
}
|
||||
conversation.Mode = alias
|
||||
}
|
||||
if strings.TrimSpace(intent.SubagentTypeName) != "" {
|
||||
conversation.SubagentTypeName = strings.TrimSpace(intent.SubagentTypeName)
|
||||
}
|
||||
if contextWindowTokens > 0 {
|
||||
conversation.TokenDetailsMaxTokens = contextWindowTokens
|
||||
} else if conversation.TokenDetailsMaxTokens == 0 {
|
||||
conversation.TokenDetailsMaxTokens = projectedConversationMaxTokens
|
||||
}
|
||||
turnSeq := conversation.NextTurnSeq
|
||||
if turnSeq <= 0 {
|
||||
turnSeq = 1
|
||||
}
|
||||
effectiveMode, err := parseModeAlias(conversation.Mode)
|
||||
if err != nil {
|
||||
return nil, agentv1.AgentMode_AGENT_MODE_AGENT, 0, nil, err
|
||||
}
|
||||
entries, err := buildRunEntries(intent, effectiveMode, turnSeq)
|
||||
if err != nil {
|
||||
return nil, agentv1.AgentMode_AGENT_MODE_AGENT, 0, nil, err
|
||||
}
|
||||
conversation.CurrentRequestID = strings.TrimSpace(intent.RequestID)
|
||||
conversation.CurrentTurnSeq = turnSeq
|
||||
conversation.CurrentLoopID = fmt.Sprintf("%d:%s", turnSeq, strings.TrimSpace(intent.RequestID))
|
||||
conversation.CurrentLoopStatus = "running"
|
||||
if conversation.NextTurnSeq <= turnSeq {
|
||||
conversation.NextTurnSeq = turnSeq + 1
|
||||
}
|
||||
return conversation, effectiveMode, turnSeq, append(importedEntries, entries...), nil
|
||||
}
|
||||
|
||||
func (service *Service) syncConversationRecord(conversationID string, conversation *ConversationFile) error {
|
||||
if service == nil || service.store == nil || conversation == nil {
|
||||
return nil
|
||||
}
|
||||
mode, err := parseModeAlias(conversation.Mode)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if _, err := service.store.CreateConversation(
|
||||
conversationID,
|
||||
mode,
|
||||
conversation.ParentConversationID,
|
||||
conversation.ParentToolCallID,
|
||||
conversation.RootConversationID,
|
||||
); err != nil {
|
||||
return err
|
||||
}
|
||||
_, err = service.store.UpdateConversationMeta(conversationID, func(item *ConversationFile) error {
|
||||
if item == nil {
|
||||
return nil
|
||||
}
|
||||
item.ConversationID = conversation.ConversationID
|
||||
item.RootConversationID = conversation.RootConversationID
|
||||
item.ParentConversationID = conversation.ParentConversationID
|
||||
item.ParentToolCallID = conversation.ParentToolCallID
|
||||
item.SubagentTypeName = conversation.SubagentTypeName
|
||||
item.Mode = conversation.Mode
|
||||
item.TokenDetailsUsedTokens = conversation.TokenDetailsUsedTokens
|
||||
item.TokenDetailsMaxTokens = conversation.TokenDetailsMaxTokens
|
||||
item.AutoCompactionPending = conversation.AutoCompactionPending
|
||||
item.AutoCompactionPromptTokens = conversation.AutoCompactionPromptTokens
|
||||
item.AutoCompactionReserveTokens = conversation.AutoCompactionReserveTokens
|
||||
item.AutoCompactionTriggeredAt = conversation.AutoCompactionTriggeredAt
|
||||
item.AutoCompactionSourceModelCallID = conversation.AutoCompactionSourceModelCallID
|
||||
item.LatestRequestPrefix = cloneConversationRequestPrefix(conversation.LatestRequestPrefix)
|
||||
item.LastProviderCall = cloneConversationProviderCall(conversation.LastProviderCall)
|
||||
item.CreatedAt = conversation.CreatedAt
|
||||
item.UpdatedAt = conversation.UpdatedAt
|
||||
item.NextTurnSeq = conversation.NextTurnSeq
|
||||
item.NextEntrySeq = conversation.NextEntrySeq
|
||||
item.CurrentLoopID = conversation.CurrentLoopID
|
||||
item.CurrentLoopStatus = conversation.CurrentLoopStatus
|
||||
item.CurrentRequestID = conversation.CurrentRequestID
|
||||
item.CurrentTurnSeq = conversation.CurrentTurnSeq
|
||||
return nil
|
||||
})
|
||||
return err
|
||||
}
|
||||
|
||||
func (service *Service) syncSummaryCarryForward(_ string, requestID string, modelCallID string) error {
|
||||
if service == nil || strings.TrimSpace(requestID) == "" || strings.TrimSpace(modelCallID) == "" {
|
||||
return nil
|
||||
}
|
||||
stream, ok := service.broker.Get(requestID)
|
||||
if !ok || stream == nil {
|
||||
return nil
|
||||
}
|
||||
conversation, _, _, err := service.snapshotCheckpointConversation(stream)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return service.syncSummarySnapshot(stream, conversation, requestID, modelCallID)
|
||||
}
|
||||
|
||||
func (service *Service) syncSummarySnapshot(stream *ActiveStream, conversation *ConversationFile, requestID string, modelCallID string) error {
|
||||
if service == nil || conversation == nil || strings.TrimSpace(requestID) == "" || strings.TrimSpace(modelCallID) == "" {
|
||||
return nil
|
||||
}
|
||||
return service.syncConversationRecord(conversation.ConversationID, conversation)
|
||||
}
|
||||
|
||||
func (service *Service) projectSummaryInputHistoryMessages(conversation *ConversationFile) ([]modeladapter.Message, error) {
|
||||
if service == nil || service.projector == nil || conversation == nil {
|
||||
return nil, nil
|
||||
}
|
||||
inputConversation := cloneConversationFile(conversation)
|
||||
inputConversation.Entries = filterConversationEntriesByKind(conversation.Entries, "user_message", "request_context", "prompt_context")
|
||||
return service.projector.ProjectPromptReplay(inputConversation)
|
||||
}
|
||||
|
||||
func (service *Service) projectSummaryOutputMessages(conversation *ConversationFile) ([]modeladapter.Message, error) {
|
||||
if service == nil || service.projector == nil || conversation == nil {
|
||||
return nil, nil
|
||||
}
|
||||
outputConversation := cloneConversationFile(conversation)
|
||||
outputConversation.Entries = filterConversationEntriesByKind(conversation.Entries, "assistant_text", "tool_call", "tool_result")
|
||||
return service.projector.ProjectPromptReplay(outputConversation)
|
||||
}
|
||||
|
||||
func filterConversationEntriesByKind(entries []HistoryEntry, kinds ...string) []HistoryEntry {
|
||||
if len(entries) == 0 || len(kinds) == 0 {
|
||||
return nil
|
||||
}
|
||||
allowed := make(map[string]struct{}, len(kinds))
|
||||
for _, kind := range kinds {
|
||||
if trimmed := strings.TrimSpace(kind); trimmed != "" {
|
||||
allowed[trimmed] = struct{}{}
|
||||
}
|
||||
}
|
||||
filtered := make([]HistoryEntry, 0, len(entries))
|
||||
for _, entry := range entries {
|
||||
if _, ok := allowed[strings.TrimSpace(entry.Kind)]; !ok {
|
||||
continue
|
||||
}
|
||||
filtered = append(filtered, entry)
|
||||
}
|
||||
return filtered
|
||||
}
|
||||
|
||||
func buildSummaryTurnOutcome(stream *ActiveStream, conversation *ConversationFile, requestID string, modelCallID string) map[string]any {
|
||||
outcome := map[string]any{
|
||||
"status": "in_progress",
|
||||
"request_id": strings.TrimSpace(requestID),
|
||||
"run_id": strings.TrimSpace(requestID),
|
||||
"model_call_id": strings.TrimSpace(modelCallID),
|
||||
}
|
||||
if stream != nil {
|
||||
outcome["model_id"] = strings.TrimSpace(stream.ModelID)
|
||||
outcome["is_prewarm"] = false
|
||||
outcome["turn_seq"] = stream.TurnSeq
|
||||
if !stream.CreatedAt.IsZero() {
|
||||
outcome["started_at"] = stream.CreatedAt.UTC().Format(time.RFC3339Nano)
|
||||
}
|
||||
}
|
||||
if conversation != nil && !conversation.UpdatedAt.IsZero() {
|
||||
outcome["last_event_at"] = conversation.UpdatedAt.UTC().Format(time.RFC3339Nano)
|
||||
}
|
||||
for index := len(conversation.Entries) - 1; index >= 0; index-- {
|
||||
entry := conversation.Entries[index]
|
||||
if strings.TrimSpace(entry.RequestID) != strings.TrimSpace(requestID) || strings.TrimSpace(entry.Kind) != "metadata" {
|
||||
continue
|
||||
}
|
||||
var payload metadataPayload
|
||||
if err := json.Unmarshal(entry.Payload, &payload); err != nil {
|
||||
continue
|
||||
}
|
||||
switch strings.TrimSpace(payload.Type) {
|
||||
case "turn_completed":
|
||||
outcome["status"] = "completed"
|
||||
case "provider_error":
|
||||
outcome["status"] = "provider_error"
|
||||
if errorText := strings.TrimSpace(readStringValue(payload.Value["error"])); errorText != "" {
|
||||
outcome["error"] = errorText
|
||||
}
|
||||
case "failed":
|
||||
outcome["status"] = "failed"
|
||||
if errorText := strings.TrimSpace(readStringValue(payload.Value["error"])); errorText != "" {
|
||||
outcome["error"] = errorText
|
||||
}
|
||||
default:
|
||||
continue
|
||||
}
|
||||
if modelCall := strings.TrimSpace(readStringValue(payload.Value["model_call_id"])); modelCall != "" {
|
||||
outcome["model_call_id"] = modelCall
|
||||
}
|
||||
if !entry.CreatedAt.IsZero() {
|
||||
outcome["last_event_at"] = entry.CreatedAt.UTC().Format(time.RFC3339Nano)
|
||||
}
|
||||
break
|
||||
}
|
||||
return outcome
|
||||
}
|
||||
|
||||
func buildSummaryAnnotations(conversation *ConversationFile, requestID string) map[string]any {
|
||||
if conversation == nil {
|
||||
return nil
|
||||
}
|
||||
for index := len(conversation.Entries) - 1; index >= 0; index-- {
|
||||
entry := conversation.Entries[index]
|
||||
if strings.TrimSpace(entry.RequestID) != strings.TrimSpace(requestID) || strings.TrimSpace(entry.Kind) != "metadata" {
|
||||
continue
|
||||
}
|
||||
var payload metadataPayload
|
||||
if err := json.Unmarshal(entry.Payload, &payload); err != nil {
|
||||
continue
|
||||
}
|
||||
if strings.TrimSpace(payload.Type) != "thought_annotation" {
|
||||
continue
|
||||
}
|
||||
if strings.TrimSpace(readStringValue(payload.Value["kind"])) != "summary_completed" {
|
||||
continue
|
||||
}
|
||||
thought := strings.TrimSpace(readStringValue(payload.Value["thought"]))
|
||||
if thought == "" {
|
||||
continue
|
||||
}
|
||||
return map[string]any{
|
||||
"summary_completed_thought": thought,
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,238 @@
|
||||
package forwarder
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"log"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
runtimecore "cursor/internal/backend/agent/core"
|
||||
)
|
||||
|
||||
const shellTerminalRecoveryGrace = 1500 * time.Millisecond
|
||||
|
||||
const (
|
||||
shellRecoveryReasonForegroundDeadline = "foreground_deadline_exceeded"
|
||||
shellRecoveryReasonTransportClosed = "transport_closed_without_terminal"
|
||||
)
|
||||
|
||||
func initializePendingExecForTracking(pending runtimecore.PendingExec) runtimecore.PendingExec {
|
||||
if strings.TrimSpace(pending.ExecKind) != "shell" {
|
||||
return pending
|
||||
}
|
||||
now := time.Now().UTC()
|
||||
if pending.OpenedAt.IsZero() {
|
||||
pending.OpenedAt = now
|
||||
}
|
||||
if pending.LastShellActivityAt.IsZero() {
|
||||
pending.LastShellActivityAt = pending.OpenedAt
|
||||
}
|
||||
if pending.ShellForegroundDeadline.IsZero() {
|
||||
pending.ShellForegroundDeadline = pending.OpenedAt.Add(shellForegroundTimeoutDuration(pending.ArgsJSON) + shellTerminalRecoveryGrace)
|
||||
}
|
||||
return pending
|
||||
}
|
||||
|
||||
func shellForegroundTimeoutDuration(argsJSON []byte) time.Duration {
|
||||
timeoutMS := int64(30000)
|
||||
args, err := runtimecore.DecodeArgsMap(argsJSON)
|
||||
if err == nil {
|
||||
if blockUntilMS, found, err := runtimecore.ReadFloat64Arg(args, "block_until_ms", "blockUntilMS"); err == nil && found {
|
||||
if blockUntilMS <= 0 {
|
||||
return 0
|
||||
}
|
||||
timeoutMS = int64(blockUntilMS)
|
||||
}
|
||||
}
|
||||
return time.Duration(timeoutMS) * time.Millisecond
|
||||
}
|
||||
|
||||
func shellForegroundTimeoutMS(argsJSON []byte) int64 {
|
||||
return shellForegroundTimeoutDuration(argsJSON).Milliseconds()
|
||||
}
|
||||
|
||||
func (service *Service) scheduleShellForegroundRecovery(requestID string, pending runtimecore.PendingExec) {
|
||||
if service == nil || strings.TrimSpace(requestID) == "" || strings.TrimSpace(pending.ExecKind) != "shell" || strings.TrimSpace(pending.ExecID) == "" {
|
||||
return
|
||||
}
|
||||
stream, ok := service.broker.Get(requestID)
|
||||
if !ok || stream == nil {
|
||||
return
|
||||
}
|
||||
deadline := pending.ShellForegroundDeadline
|
||||
if deadline.IsZero() {
|
||||
deadline = time.Now().UTC().Add(shellForegroundTimeoutDuration(pending.ArgsJSON) + shellTerminalRecoveryGrace)
|
||||
}
|
||||
service.scheduleStreamTimer(
|
||||
stream,
|
||||
providerTimerKey(streamTimerShellForeground, pending.ExecID),
|
||||
time.Until(deadline),
|
||||
streamTimerShellForeground,
|
||||
pending.ExecID,
|
||||
pending.MessageID,
|
||||
shellRecoveryReasonForegroundDeadline,
|
||||
)
|
||||
}
|
||||
|
||||
func (service *Service) scheduleShellTransportCloseRecovery(requestID string, pending runtimecore.PendingExec) {
|
||||
if service == nil || strings.TrimSpace(requestID) == "" || strings.TrimSpace(pending.ExecKind) != "shell" || strings.TrimSpace(pending.ExecID) == "" {
|
||||
return
|
||||
}
|
||||
stream, ok := service.broker.Get(requestID)
|
||||
if !ok || stream == nil {
|
||||
return
|
||||
}
|
||||
if _, scheduled := markShellRecoveryScheduled(stream, pending.ExecID); !scheduled {
|
||||
return
|
||||
}
|
||||
service.scheduleStreamTimer(
|
||||
stream,
|
||||
providerTimerKey(streamTimerShellTransportClose, pending.ExecID),
|
||||
shellTerminalRecoveryGrace,
|
||||
streamTimerShellTransportClose,
|
||||
pending.ExecID,
|
||||
pending.MessageID,
|
||||
shellRecoveryReasonTransportClosed,
|
||||
)
|
||||
}
|
||||
|
||||
func snapshotPendingExecWithStatus(stream *ActiveStream, execID string) (runtimecore.PendingExec, StreamStatus, bool) {
|
||||
if stream == nil || strings.TrimSpace(execID) == "" {
|
||||
return runtimecore.PendingExec{}, "", false
|
||||
}
|
||||
stream.mu.Lock()
|
||||
defer stream.mu.Unlock()
|
||||
item, ok := stream.PendingExecs[strings.TrimSpace(execID)]
|
||||
if !ok {
|
||||
return runtimecore.PendingExec{}, stream.Status, false
|
||||
}
|
||||
return item, stream.Status, true
|
||||
}
|
||||
|
||||
func markShellRecoveryScheduled(stream *ActiveStream, execID string) (runtimecore.PendingExec, bool) {
|
||||
if stream == nil || strings.TrimSpace(execID) == "" {
|
||||
return runtimecore.PendingExec{}, false
|
||||
}
|
||||
stream.mu.Lock()
|
||||
defer stream.mu.Unlock()
|
||||
current, ok := stream.PendingExecs[strings.TrimSpace(execID)]
|
||||
if !ok || strings.TrimSpace(current.ExecKind) != "shell" || current.ShellRecoveryScheduled {
|
||||
return current, false
|
||||
}
|
||||
current.ShellRecoveryScheduled = true
|
||||
stream.PendingExecs[strings.TrimSpace(execID)] = current
|
||||
stream.UpdatedAt = time.Now().UTC()
|
||||
return current, true
|
||||
}
|
||||
|
||||
func (service *Service) recoverShellWithoutTerminalIfNeeded(stream *ActiveStream, execID string, messageID uint32, reason string) error {
|
||||
if stream == nil || strings.TrimSpace(execID) == "" {
|
||||
return nil
|
||||
}
|
||||
current, status, found := snapshotPendingExecWithStatus(stream, execID)
|
||||
if !found || current.MessageID != messageID || strings.TrimSpace(current.ExecKind) != "shell" || isTerminalStreamStatus(status) {
|
||||
return nil
|
||||
}
|
||||
switch strings.TrimSpace(current.StreamState) {
|
||||
case "exited", "backgrounded", "rejected", "permission_denied":
|
||||
return nil
|
||||
}
|
||||
if reason == shellRecoveryReasonForegroundDeadline && !current.ShellForegroundDeadline.IsZero() && time.Now().UTC().Before(current.ShellForegroundDeadline) {
|
||||
return nil
|
||||
}
|
||||
return service.recoverShellWithoutTerminal(stream, current, reason)
|
||||
}
|
||||
|
||||
func (service *Service) recoverShellWithoutTerminal(stream *ActiveStream, pending runtimecore.PendingExec, reason string) error {
|
||||
if stream == nil {
|
||||
return nil
|
||||
}
|
||||
pending.ShellRecoveryScheduled = true
|
||||
markExecCompleted(stream, pending)
|
||||
resultPayload := buildSyntheticShellResultPayload(pending, reason)
|
||||
log.Printf(
|
||||
"forwarder synthetic shell recovery request_id=%s tool_call_id=%s message_id=%d exec_id=%s reason=%s stream_state=%s",
|
||||
strings.TrimSpace(stream.RequestID),
|
||||
strings.TrimSpace(pending.ToolCallID),
|
||||
pending.MessageID,
|
||||
strings.TrimSpace(pending.ExecID),
|
||||
strings.TrimSpace(reason),
|
||||
strings.TrimSpace(pending.StreamState),
|
||||
)
|
||||
if err := service.appendToolResult(stream, pending.ToolCallID, deriveToolNameFromPendingExec(pending), pending.ArgsJSON, resultPayload, pending.ReasoningContent, nil); err != nil {
|
||||
return err
|
||||
}
|
||||
if _, err := service.appendConversationEntries(stream, stream.ConversationID, []HistoryEntry{
|
||||
newMetadataEntry(stream.TurnSeq, stream.RequestID, "shell_stream_stalled", map[string]any{
|
||||
"tool_call_id": pending.ToolCallID,
|
||||
"message_id": pending.MessageID,
|
||||
"exec_id": pending.ExecID,
|
||||
"exec_kind": pending.ExecKind,
|
||||
"reason": strings.TrimSpace(reason),
|
||||
"recent_stream_state": pending.StreamState,
|
||||
"chunk_count": pending.ChunkCount,
|
||||
"first_chunk_at": pending.FirstChunkAt,
|
||||
"last_activity_at": pending.LastShellActivityAt,
|
||||
"last_heartbeat_at": pending.LastShellHeartbeatAt,
|
||||
"foreground_deadline": pending.ShellForegroundDeadline,
|
||||
"timeout_ms": shellForegroundTimeoutMS(pending.ArgsJSON),
|
||||
"stdout_buffer_bytes": len(pending.StdoutBuffer),
|
||||
"stderr_buffer_bytes": len(pending.StderrBuffer),
|
||||
"shell_recovery_scheduled": pending.ShellRecoveryScheduled,
|
||||
}),
|
||||
}); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := service.syncSummaryCarryForward(stream.ConversationID, stream.RequestID, pending.ModelCallID); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := service.publishToolCallCompleted(stream.RequestID, pending.ToolCallID, pending.ModelCallID, nil); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := service.publishCheckpoint(stream.RequestID, stream.ConversationID); err != nil {
|
||||
return err
|
||||
}
|
||||
return service.reconcileStream(stream)
|
||||
}
|
||||
|
||||
func buildSyntheticShellResultPayload(pending runtimecore.PendingExec, reason string) string {
|
||||
sections := make([]string, 0, 2)
|
||||
if captured := summarizeCapturedShellOutput(pending.StdoutBuffer, pending.StderrBuffer); captured != "" {
|
||||
sections = append(sections, captured)
|
||||
}
|
||||
noteLines := []string{
|
||||
"<shell-incomplete>",
|
||||
"Missing terminal shell stream event (expected exit or backgrounded).",
|
||||
}
|
||||
switch strings.TrimSpace(reason) {
|
||||
case shellRecoveryReasonTransportClosed:
|
||||
noteLines = append(noteLines, "The shell transport closed before a terminal event arrived.")
|
||||
case shellRecoveryReasonForegroundDeadline:
|
||||
noteLines = append(noteLines, fmt.Sprintf("The foreground wait window expired after %dms without a terminal event.", shellForegroundTimeoutMS(pending.ArgsJSON)))
|
||||
default:
|
||||
noteLines = append(noteLines, "The shell stream stopped progressing before a terminal event arrived.")
|
||||
}
|
||||
noteLines = append(noteLines,
|
||||
"The command may still be running in the Cursor app client.",
|
||||
"</shell-incomplete>",
|
||||
)
|
||||
sections = append(sections, strings.Join(noteLines, "\n"))
|
||||
return strings.Join(sections, "\n\n")
|
||||
}
|
||||
|
||||
func summarizeCapturedShellOutput(stdout string, stderr string) string {
|
||||
trimmedStdout := strings.TrimSpace(stdout)
|
||||
trimmedStderr := strings.TrimSpace(stderr)
|
||||
sections := make([]string, 0, 2)
|
||||
if trimmedStdout != "" {
|
||||
sections = append(sections, trimmedStdout)
|
||||
}
|
||||
if trimmedStderr != "" {
|
||||
if trimmedStdout != "" {
|
||||
sections = append(sections, "<stderr>\n"+trimmedStderr+"\n</stderr>")
|
||||
} else {
|
||||
sections = append(sections, trimmedStderr)
|
||||
}
|
||||
}
|
||||
return strings.Join(sections, "\n\n")
|
||||
}
|
||||
@@ -0,0 +1,37 @@
|
||||
package forwarder
|
||||
|
||||
type latestSummaryState struct {
|
||||
RuntimeSnapshot *ConversationFile
|
||||
}
|
||||
|
||||
func (service *Service) loadLatestSummaryState(conversationID string) (*latestSummaryState, bool, error) {
|
||||
if service == nil || service.store == nil {
|
||||
return nil, false, nil
|
||||
}
|
||||
conversation, err := service.store.LoadConversation(conversationID)
|
||||
if err != nil || conversation == nil {
|
||||
return nil, false, err
|
||||
}
|
||||
return &latestSummaryState{RuntimeSnapshot: conversation}, true, nil
|
||||
}
|
||||
|
||||
func (service *Service) loadLatestCarryForwardReplay(_ string) ([][]byte, bool, error) {
|
||||
return nil, false, nil
|
||||
}
|
||||
|
||||
func (service *Service) loadLatestSummaryPromptTokens(conversationID string) (int64, bool, error) {
|
||||
if service == nil || service.store == nil {
|
||||
return 0, false, nil
|
||||
}
|
||||
conversation, err := service.store.LoadConversation(conversationID)
|
||||
if err != nil || conversation == nil {
|
||||
return 0, false, err
|
||||
}
|
||||
if conversation.AutoCompactionPromptTokens > 0 {
|
||||
return conversation.AutoCompactionPromptTokens, true, nil
|
||||
}
|
||||
if conversation.TokenDetailsUsedTokens > 0 {
|
||||
return int64(conversation.TokenDetailsUsedTokens), true, nil
|
||||
}
|
||||
return 0, false, nil
|
||||
}
|
||||
@@ -0,0 +1,871 @@
|
||||
package forwarder
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"sort"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"google.golang.org/protobuf/encoding/protojson"
|
||||
"google.golang.org/protobuf/proto"
|
||||
|
||||
"cursor/gen/agentv1"
|
||||
modeladapter "cursor/internal/backend/agent/model"
|
||||
)
|
||||
|
||||
const todoSectionReminderMessage = "<system_reminder>\nYou are currently under the todo section, be sure to track tasks and do not forget to update.\n</system_reminder>"
|
||||
|
||||
const (
|
||||
promptContextSourceStructuredCurrentPlan = "structured_state/current_plan"
|
||||
promptContextSourceStructuredTodoList = "structured_state/todo_list"
|
||||
promptContextSourceStructuredTodoReminder = "structured_state/todo_reminder"
|
||||
)
|
||||
|
||||
const todoUnsafeReplaceBaseError = "todo update rejected: merge=false may only omit completed or cancelled todos; use merge=true for incremental updates or include every active todo"
|
||||
|
||||
type structuredConversationState struct {
|
||||
PlanText string
|
||||
HasPlan bool
|
||||
Plans map[string]*agentv1.PlanRegistryEntry
|
||||
Todos []*agentv1.TodoItem
|
||||
HasTodos bool
|
||||
}
|
||||
|
||||
func projectConversationStructuredState(conversation *ConversationFile) (structuredConversationState, error) {
|
||||
state := structuredConversationState{}
|
||||
if conversation == nil {
|
||||
return state, nil
|
||||
}
|
||||
for _, entry := range checkpointProjectionEntries(conversation.Entries) {
|
||||
switch strings.TrimSpace(entry.Kind) {
|
||||
case "runtime_state":
|
||||
var payload runtimeStateEntryPayload
|
||||
if err := json.Unmarshal(entry.Payload, &payload); err != nil {
|
||||
return structuredConversationState{}, fmt.Errorf("decode runtime_state entry: %w", err)
|
||||
}
|
||||
if strings.TrimSpace(payload.PlanText) != "" {
|
||||
state.PlanText = strings.TrimSpace(payload.PlanText)
|
||||
state.HasPlan = true
|
||||
}
|
||||
if len(payload.Plans) > 0 {
|
||||
state.Plans = clonePlanRegistryEntries(payload.Plans)
|
||||
}
|
||||
if len(payload.Todos) > 0 {
|
||||
state.Todos = cloneTodoItems(payload.Todos)
|
||||
state.HasTodos = true
|
||||
}
|
||||
continue
|
||||
case "tool_result":
|
||||
default:
|
||||
continue
|
||||
}
|
||||
var payload toolResultEntryPayload
|
||||
if err := json.Unmarshal(entry.Payload, &payload); err != nil {
|
||||
return structuredConversationState{}, fmt.Errorf("decode structured state tool_result: %w", err)
|
||||
}
|
||||
if len(payload.ToolCall) == 0 {
|
||||
continue
|
||||
}
|
||||
toolCall := &agentv1.ToolCall{}
|
||||
if err := protojson.Unmarshal(payload.ToolCall, toolCall); err != nil {
|
||||
return structuredConversationState{}, fmt.Errorf("decode structured state tool_call: %w", err)
|
||||
}
|
||||
entryTime := entry.CreatedAt
|
||||
switch item := toolCall.GetTool().(type) {
|
||||
case *agentv1.ToolCall_CreatePlanToolCall:
|
||||
createPlanToolCall := item.CreatePlanToolCall
|
||||
if createPlanToolCall == nil || createPlanToolCall.GetResult().GetSuccess() == nil {
|
||||
continue
|
||||
}
|
||||
planURI := strings.TrimSpace(createPlanToolCall.GetResult().GetPlanUri())
|
||||
if planURI == "" {
|
||||
continue
|
||||
}
|
||||
state.PlanText = strings.TrimSpace(createPlanToolCall.GetArgs().GetPlan())
|
||||
state.HasPlan = true
|
||||
state.Plans = upsertCurrentPlanRegistryEntry(state.Plans, planURI)
|
||||
todos, err := normalizeTodoItems(flattenCreatePlanTodos(createPlanToolCall.GetArgs()), entryTime, false)
|
||||
if err != nil {
|
||||
return structuredConversationState{}, err
|
||||
}
|
||||
state.Todos = todos
|
||||
state.HasTodos = true
|
||||
case *agentv1.ToolCall_UpdateTodosToolCall:
|
||||
updateTodosToolCall := item.UpdateTodosToolCall
|
||||
if updateTodosToolCall == nil || updateTodosToolCall.GetResult().GetSuccess() == nil {
|
||||
continue
|
||||
}
|
||||
success := updateTodosToolCall.GetResult().GetSuccess()
|
||||
todos, err := normalizeTodoItems(success.GetTodos(), entryTime, false)
|
||||
if err != nil {
|
||||
return structuredConversationState{}, err
|
||||
}
|
||||
if !success.GetWasMerge() && len(missingActiveTodoReplacementIDs(state.Todos, todos)) > 0 {
|
||||
todos, err = mergeTodoItems(state.Todos, todos, entryTime)
|
||||
if err != nil {
|
||||
return structuredConversationState{}, err
|
||||
}
|
||||
}
|
||||
state.Todos = todos
|
||||
state.HasTodos = true
|
||||
}
|
||||
}
|
||||
return state, nil
|
||||
}
|
||||
|
||||
func refreshConversationRuntimeState(conversation *ConversationFile) error {
|
||||
if conversation == nil {
|
||||
return nil
|
||||
}
|
||||
state, err := projectConversationStructuredState(conversation)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
conversation.CurrentPlanText = ""
|
||||
conversation.CurrentPlans = nil
|
||||
conversation.CurrentTodos = nil
|
||||
if state.HasPlan && strings.TrimSpace(state.PlanText) != "" {
|
||||
conversation.CurrentPlanText = strings.TrimSpace(state.PlanText)
|
||||
}
|
||||
if len(state.Plans) > 0 {
|
||||
conversation.CurrentPlans = clonePlanRegistryEntries(state.Plans)
|
||||
}
|
||||
if state.HasTodos && len(state.Todos) > 0 {
|
||||
conversation.CurrentTodos = cloneTodoItems(state.Todos)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func buildStructuredStatePromptContexts(conversation *ConversationFile) ([]PromptContextMessage, []modeladapter.Message, error) {
|
||||
state, err := projectConversationStructuredState(conversation)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
contexts := make([]PromptContextMessage, 0, 2)
|
||||
tailMessages := make([]modeladapter.Message, 0, 1)
|
||||
if state.HasPlan && strings.TrimSpace(state.PlanText) != "" {
|
||||
contexts = append(contexts, newPromptContextMessage(
|
||||
promptContextSourceStructuredCurrentPlan,
|
||||
modeladapter.Message{
|
||||
Role: "user",
|
||||
Content: "<current_plan>\n" + strings.TrimSpace(state.PlanText) + "\n</current_plan>",
|
||||
},
|
||||
false,
|
||||
))
|
||||
}
|
||||
if state.HasTodos {
|
||||
if todoText := renderTodoList(state.Todos); todoText != "" {
|
||||
contexts = append(contexts, newPromptContextMessage(
|
||||
promptContextSourceStructuredTodoList,
|
||||
modeladapter.Message{
|
||||
Role: "user",
|
||||
Content: "<todo_list>\n" + todoText + "\n</todo_list>",
|
||||
},
|
||||
false,
|
||||
))
|
||||
tailMessages = append(tailMessages, modeladapter.Message{
|
||||
Role: "user",
|
||||
Content: todoSectionReminderMessage,
|
||||
})
|
||||
}
|
||||
}
|
||||
return contexts, tailMessages, nil
|
||||
}
|
||||
|
||||
func upsertCurrentPlanRegistryEntry(plans map[string]*agentv1.PlanRegistryEntry, planURI string) map[string]*agentv1.PlanRegistryEntry {
|
||||
trimmedURI := strings.TrimSpace(planURI)
|
||||
if trimmedURI == "" {
|
||||
return clonePlanRegistryEntries(plans)
|
||||
}
|
||||
next := clonePlanRegistryEntries(plans)
|
||||
if next == nil {
|
||||
next = make(map[string]*agentv1.PlanRegistryEntry, 1)
|
||||
}
|
||||
key := "current"
|
||||
if _, ok := next[key]; !ok && len(next) == 1 {
|
||||
for existingKey := range next {
|
||||
key = existingKey
|
||||
}
|
||||
}
|
||||
next[key] = &agentv1.PlanRegistryEntry{
|
||||
Id: key,
|
||||
Path: trimmedURI,
|
||||
}
|
||||
return next
|
||||
}
|
||||
|
||||
func clonePlanRegistryEntries(items map[string]*agentv1.PlanRegistryEntry) map[string]*agentv1.PlanRegistryEntry {
|
||||
if len(items) == 0 {
|
||||
return nil
|
||||
}
|
||||
cloned := make(map[string]*agentv1.PlanRegistryEntry, len(items))
|
||||
for key, item := range items {
|
||||
trimmedKey := strings.TrimSpace(key)
|
||||
if trimmedKey == "" || item == nil {
|
||||
continue
|
||||
}
|
||||
cloned[trimmedKey] = &agentv1.PlanRegistryEntry{
|
||||
Id: strings.TrimSpace(item.GetId()),
|
||||
Path: strings.TrimSpace(item.GetPath()),
|
||||
}
|
||||
if cloned[trimmedKey].GetId() == "" {
|
||||
cloned[trimmedKey].Id = trimmedKey
|
||||
}
|
||||
}
|
||||
if len(cloned) == 0 {
|
||||
return nil
|
||||
}
|
||||
return cloned
|
||||
}
|
||||
|
||||
func encodeConversationPlanBytes(plan string) []byte {
|
||||
text := strings.TrimSpace(plan)
|
||||
if text == "" {
|
||||
return nil
|
||||
}
|
||||
payload, err := proto.Marshal(&agentv1.ConversationPlan{Plan: text})
|
||||
if err != nil {
|
||||
return nil
|
||||
}
|
||||
return payload
|
||||
}
|
||||
|
||||
func encodeConversationTodoBytes(items []*agentv1.TodoItem) [][]byte {
|
||||
if len(items) == 0 {
|
||||
return nil
|
||||
}
|
||||
encoded := make([][]byte, 0, len(items))
|
||||
for _, item := range items {
|
||||
if item == nil {
|
||||
continue
|
||||
}
|
||||
payload, err := proto.Marshal(cloneTodoItem(item))
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
encoded = append(encoded, payload)
|
||||
}
|
||||
if len(encoded) == 0 {
|
||||
return nil
|
||||
}
|
||||
return encoded
|
||||
}
|
||||
|
||||
func flattenCreatePlanTodos(args *agentv1.CreatePlanArgs) []*agentv1.TodoItem {
|
||||
if args == nil {
|
||||
return nil
|
||||
}
|
||||
if len(args.GetTodos()) > 0 {
|
||||
return cloneTodoItems(args.GetTodos())
|
||||
}
|
||||
items := make([]*agentv1.TodoItem, 0, len(args.GetPhases()))
|
||||
for _, phase := range args.GetPhases() {
|
||||
if phase == nil {
|
||||
continue
|
||||
}
|
||||
items = append(items, cloneTodoItems(phase.GetTodos())...)
|
||||
}
|
||||
return items
|
||||
}
|
||||
|
||||
type decodedUpdateTodosArgs struct {
|
||||
Args *agentv1.UpdateTodosArgs
|
||||
MergeSet bool
|
||||
}
|
||||
|
||||
func decodeUpdateTodosArgsJSON(raw []byte) (*agentv1.UpdateTodosArgs, error) {
|
||||
decoded, err := decodeUpdateTodosArgsJSONWithPresence(raw)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return decoded.Args, nil
|
||||
}
|
||||
|
||||
func decodeUpdateTodosArgsJSONWithPresence(raw []byte) (decodedUpdateTodosArgs, error) {
|
||||
payload, err := decodeJSONObject(raw)
|
||||
if err != nil {
|
||||
return decodedUpdateTodosArgs{}, err
|
||||
}
|
||||
todos, err := decodeTodoItemsValue(valueByAlias(payload, "todos"), false)
|
||||
if err != nil {
|
||||
return decodedUpdateTodosArgs{}, err
|
||||
}
|
||||
mergeValue, mergeSet := valueByAliasWithPresence(payload, "merge")
|
||||
return decodedUpdateTodosArgs{
|
||||
Args: &agentv1.UpdateTodosArgs{
|
||||
Todos: todos,
|
||||
Merge: boolValue(mergeValue),
|
||||
},
|
||||
MergeSet: mergeSet,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func decodeReadTodosArgsJSON(raw []byte) (*agentv1.ReadTodosArgs, error) {
|
||||
payload, err := decodeJSONObject(raw)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
statuses, err := decodeTodoStatusesValue(valueByAlias(payload, "status_filter", "statusFilter"))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &agentv1.ReadTodosArgs{
|
||||
StatusFilter: statuses,
|
||||
IdFilter: stringSliceValue(valueByAlias(payload, "id_filter", "idFilter")),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func decodeJSONObject(raw []byte) (map[string]any, error) {
|
||||
if len(raw) == 0 {
|
||||
return map[string]any{}, nil
|
||||
}
|
||||
var payload map[string]any
|
||||
if err := json.Unmarshal(raw, &payload); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if payload == nil {
|
||||
payload = map[string]any{}
|
||||
}
|
||||
return payload, nil
|
||||
}
|
||||
|
||||
func decodeTodoItemsValue(value any, validate bool) ([]*agentv1.TodoItem, error) {
|
||||
items, ok := value.([]any)
|
||||
if !ok || len(items) == 0 {
|
||||
return nil, nil
|
||||
}
|
||||
todos := make([]*agentv1.TodoItem, 0, len(items))
|
||||
for _, rawItem := range items {
|
||||
item, ok := rawItem.(map[string]any)
|
||||
if !ok {
|
||||
return nil, fmt.Errorf("todo item must be an object")
|
||||
}
|
||||
status, err := todoStatusFromValue(valueByAlias(item, "status"))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
todo := &agentv1.TodoItem{
|
||||
Id: strings.TrimSpace(stringValue(valueByAlias(item, "id"))),
|
||||
Content: strings.TrimSpace(stringValue(valueByAlias(item, "content"))),
|
||||
Status: status,
|
||||
CreatedAt: int64Value(valueByAlias(item, "created_at", "createdAt")),
|
||||
UpdatedAt: int64Value(valueByAlias(item, "updated_at", "updatedAt")),
|
||||
Dependencies: stringSliceValue(valueByAlias(item, "dependencies")),
|
||||
}
|
||||
if validate {
|
||||
if todo.GetId() == "" {
|
||||
return nil, fmt.Errorf("todo id is required")
|
||||
}
|
||||
if todo.GetContent() == "" {
|
||||
return nil, fmt.Errorf("todo content is required")
|
||||
}
|
||||
}
|
||||
todos = append(todos, todo)
|
||||
}
|
||||
return todos, nil
|
||||
}
|
||||
|
||||
func decodeTodoStatusesValue(value any) ([]agentv1.TodoStatus, error) {
|
||||
items, ok := value.([]any)
|
||||
if !ok || len(items) == 0 {
|
||||
return nil, nil
|
||||
}
|
||||
statuses := make([]agentv1.TodoStatus, 0, len(items))
|
||||
for _, item := range items {
|
||||
status, err := todoStatusFromValue(item)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if status == agentv1.TodoStatus_TODO_STATUS_UNSPECIFIED {
|
||||
continue
|
||||
}
|
||||
statuses = append(statuses, status)
|
||||
}
|
||||
return statuses, nil
|
||||
}
|
||||
|
||||
func valueByAlias(values map[string]any, aliases ...string) any {
|
||||
value, _ := valueByAliasWithPresence(values, aliases...)
|
||||
return value
|
||||
}
|
||||
|
||||
func valueByAliasWithPresence(values map[string]any, aliases ...string) (any, bool) {
|
||||
for _, alias := range aliases {
|
||||
if alias == "" {
|
||||
continue
|
||||
}
|
||||
if value, ok := values[alias]; ok {
|
||||
return value, true
|
||||
}
|
||||
}
|
||||
return nil, false
|
||||
}
|
||||
|
||||
func stringValue(value any) string {
|
||||
text, ok := value.(string)
|
||||
if !ok {
|
||||
return ""
|
||||
}
|
||||
return strings.TrimSpace(text)
|
||||
}
|
||||
|
||||
func boolValue(value any) bool {
|
||||
switch item := value.(type) {
|
||||
case bool:
|
||||
return item
|
||||
case string:
|
||||
parsed, _ := strconv.ParseBool(strings.TrimSpace(item))
|
||||
return parsed
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func int64Value(value any) int64 {
|
||||
switch item := value.(type) {
|
||||
case float64:
|
||||
return int64(item)
|
||||
case float32:
|
||||
return int64(item)
|
||||
case int64:
|
||||
return item
|
||||
case int32:
|
||||
return int64(item)
|
||||
case int:
|
||||
return int64(item)
|
||||
case json.Number:
|
||||
parsed, _ := item.Int64()
|
||||
return parsed
|
||||
case string:
|
||||
parsed, _ := strconv.ParseInt(strings.TrimSpace(item), 10, 64)
|
||||
return parsed
|
||||
default:
|
||||
return 0
|
||||
}
|
||||
}
|
||||
|
||||
func stringSliceValue(value any) []string {
|
||||
items, ok := value.([]any)
|
||||
if !ok || len(items) == 0 {
|
||||
return nil
|
||||
}
|
||||
result := make([]string, 0, len(items))
|
||||
for _, item := range items {
|
||||
text := stringValue(item)
|
||||
if text == "" {
|
||||
continue
|
||||
}
|
||||
result = append(result, text)
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
func todoStatusFromValue(value any) (agentv1.TodoStatus, error) {
|
||||
switch item := value.(type) {
|
||||
case nil:
|
||||
return agentv1.TodoStatus_TODO_STATUS_UNSPECIFIED, nil
|
||||
case float64:
|
||||
return agentv1.TodoStatus(int32(item)), nil
|
||||
case float32:
|
||||
return agentv1.TodoStatus(int32(item)), nil
|
||||
case int:
|
||||
return agentv1.TodoStatus(item), nil
|
||||
case int32:
|
||||
return agentv1.TodoStatus(item), nil
|
||||
case int64:
|
||||
return agentv1.TodoStatus(item), nil
|
||||
case string:
|
||||
return todoStatusFromString(item)
|
||||
default:
|
||||
return agentv1.TodoStatus_TODO_STATUS_UNSPECIFIED, fmt.Errorf("unsupported todo status type %T", value)
|
||||
}
|
||||
}
|
||||
|
||||
func todoStatusFromString(raw string) (agentv1.TodoStatus, error) {
|
||||
switch strings.ToLower(strings.TrimSpace(raw)) {
|
||||
case "", "unspecified", "todo_status_unspecified":
|
||||
return agentv1.TodoStatus_TODO_STATUS_UNSPECIFIED, nil
|
||||
case "pending", "todo_status_pending":
|
||||
return agentv1.TodoStatus_TODO_STATUS_PENDING, nil
|
||||
case "in_progress", "in-progress", "inprogress", "todo_status_in_progress":
|
||||
return agentv1.TodoStatus_TODO_STATUS_IN_PROGRESS, nil
|
||||
case "completed", "complete", "todo_status_completed":
|
||||
return agentv1.TodoStatus_TODO_STATUS_COMPLETED, nil
|
||||
case "cancelled", "canceled", "todo_status_cancelled":
|
||||
return agentv1.TodoStatus_TODO_STATUS_CANCELLED, nil
|
||||
default:
|
||||
return agentv1.TodoStatus_TODO_STATUS_UNSPECIFIED, fmt.Errorf("unsupported todo status %q", raw)
|
||||
}
|
||||
}
|
||||
|
||||
func todoStatusLabel(status agentv1.TodoStatus) string {
|
||||
switch status {
|
||||
case agentv1.TodoStatus_TODO_STATUS_PENDING:
|
||||
return "pending"
|
||||
case agentv1.TodoStatus_TODO_STATUS_IN_PROGRESS:
|
||||
return "in_progress"
|
||||
case agentv1.TodoStatus_TODO_STATUS_COMPLETED:
|
||||
return "completed"
|
||||
case agentv1.TodoStatus_TODO_STATUS_CANCELLED:
|
||||
return "cancelled"
|
||||
default:
|
||||
return "unspecified"
|
||||
}
|
||||
}
|
||||
|
||||
func normalizeTodoItems(items []*agentv1.TodoItem, baseTime time.Time, validate bool) ([]*agentv1.TodoItem, error) {
|
||||
if len(items) == 0 {
|
||||
return nil, nil
|
||||
}
|
||||
baseMillis := timestampMillis(baseTime)
|
||||
normalized := make([]*agentv1.TodoItem, 0, len(items))
|
||||
for _, item := range items {
|
||||
if item == nil {
|
||||
continue
|
||||
}
|
||||
next := cloneTodoItem(item)
|
||||
next.Id = strings.TrimSpace(next.GetId())
|
||||
next.Content = strings.TrimSpace(next.GetContent())
|
||||
if validate {
|
||||
if next.GetId() == "" {
|
||||
return nil, fmt.Errorf("todo id is required")
|
||||
}
|
||||
if next.GetContent() == "" {
|
||||
return nil, fmt.Errorf("todo content is required")
|
||||
}
|
||||
}
|
||||
if next.GetStatus() == agentv1.TodoStatus_TODO_STATUS_UNSPECIFIED {
|
||||
next.Status = agentv1.TodoStatus_TODO_STATUS_PENDING
|
||||
}
|
||||
if next.GetCreatedAt() <= 0 {
|
||||
next.CreatedAt = baseMillis
|
||||
}
|
||||
if next.GetUpdatedAt() <= 0 {
|
||||
next.UpdatedAt = next.GetCreatedAt()
|
||||
}
|
||||
normalized = append(normalized, next)
|
||||
}
|
||||
return normalized, nil
|
||||
}
|
||||
|
||||
func mergeTodoItems(existing []*agentv1.TodoItem, updates []*agentv1.TodoItem, baseTime time.Time) ([]*agentv1.TodoItem, error) {
|
||||
result, err := normalizeTodoItems(existing, baseTime, false)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
indexByID := make(map[string]int, len(result))
|
||||
for index, item := range result {
|
||||
if item == nil || strings.TrimSpace(item.GetId()) == "" {
|
||||
continue
|
||||
}
|
||||
indexByID[strings.TrimSpace(item.GetId())] = index
|
||||
}
|
||||
nowMillis := timestampMillis(baseTime)
|
||||
for _, update := range updates {
|
||||
if update == nil {
|
||||
continue
|
||||
}
|
||||
id := strings.TrimSpace(update.GetId())
|
||||
if id == "" {
|
||||
return nil, fmt.Errorf("todo id is required")
|
||||
}
|
||||
content := strings.TrimSpace(update.GetContent())
|
||||
if index, ok := indexByID[id]; ok {
|
||||
current := result[index]
|
||||
next := cloneTodoItem(current)
|
||||
next.Id = id
|
||||
if content != "" {
|
||||
next.Content = content
|
||||
}
|
||||
if update.GetStatus() != agentv1.TodoStatus_TODO_STATUS_UNSPECIFIED {
|
||||
next.Status = update.GetStatus()
|
||||
}
|
||||
if len(update.GetDependencies()) > 0 {
|
||||
next.Dependencies = append([]string(nil), update.GetDependencies()...)
|
||||
}
|
||||
if next.GetContent() == "" {
|
||||
return nil, fmt.Errorf("todo content is required")
|
||||
}
|
||||
if next.GetStatus() == agentv1.TodoStatus_TODO_STATUS_UNSPECIFIED {
|
||||
next.Status = agentv1.TodoStatus_TODO_STATUS_PENDING
|
||||
}
|
||||
next.CreatedAt = current.GetCreatedAt()
|
||||
if next.GetCreatedAt() <= 0 {
|
||||
next.CreatedAt = nowMillis
|
||||
}
|
||||
next.UpdatedAt = nowMillis
|
||||
result[index] = next
|
||||
continue
|
||||
}
|
||||
next := cloneTodoItem(update)
|
||||
next.Id = id
|
||||
next.Content = content
|
||||
if next.GetContent() == "" {
|
||||
return nil, fmt.Errorf("todo content is required for new todo")
|
||||
}
|
||||
if next.GetStatus() == agentv1.TodoStatus_TODO_STATUS_UNSPECIFIED {
|
||||
next.Status = agentv1.TodoStatus_TODO_STATUS_PENDING
|
||||
}
|
||||
if next.GetCreatedAt() <= 0 {
|
||||
next.CreatedAt = nowMillis
|
||||
}
|
||||
if next.GetUpdatedAt() <= 0 {
|
||||
next.UpdatedAt = next.GetCreatedAt()
|
||||
}
|
||||
indexByID[id] = len(result)
|
||||
result = append(result, next)
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func missingActiveTodoReplacementIDs(existing []*agentv1.TodoItem, replacement []*agentv1.TodoItem) []string {
|
||||
existingIDs := make(map[string]struct{}, len(existing))
|
||||
for _, item := range existing {
|
||||
if item == nil || isTerminalTodoStatus(item.GetStatus()) {
|
||||
continue
|
||||
}
|
||||
id := strings.TrimSpace(item.GetId())
|
||||
if id == "" {
|
||||
continue
|
||||
}
|
||||
existingIDs[id] = struct{}{}
|
||||
}
|
||||
if len(existingIDs) == 0 {
|
||||
return nil
|
||||
}
|
||||
replacementIDs := make(map[string]struct{}, len(replacement))
|
||||
for _, item := range replacement {
|
||||
id := strings.TrimSpace(item.GetId())
|
||||
if id == "" {
|
||||
continue
|
||||
}
|
||||
replacementIDs[id] = struct{}{}
|
||||
}
|
||||
missing := make([]string, 0)
|
||||
for id := range existingIDs {
|
||||
if _, ok := replacementIDs[id]; !ok {
|
||||
missing = append(missing, id)
|
||||
}
|
||||
}
|
||||
sort.Strings(missing)
|
||||
return missing
|
||||
}
|
||||
|
||||
func isTerminalTodoStatus(status agentv1.TodoStatus) bool {
|
||||
switch status {
|
||||
case agentv1.TodoStatus_TODO_STATUS_COMPLETED,
|
||||
agentv1.TodoStatus_TODO_STATUS_CANCELLED:
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func unsafeTodoReplaceError(missingIDs []string) string {
|
||||
if len(missingIDs) == 0 {
|
||||
return todoUnsafeReplaceBaseError
|
||||
}
|
||||
return fmt.Sprintf("%s; missing active todo ids: %s", todoUnsafeReplaceBaseError, strings.Join(missingIDs, ", "))
|
||||
}
|
||||
|
||||
func filterTodoItems(items []*agentv1.TodoItem, statuses []agentv1.TodoStatus, ids []string) []*agentv1.TodoItem {
|
||||
if len(items) == 0 {
|
||||
return nil
|
||||
}
|
||||
statusFilter := make(map[agentv1.TodoStatus]struct{}, len(statuses))
|
||||
for _, status := range statuses {
|
||||
statusFilter[status] = struct{}{}
|
||||
}
|
||||
idFilter := make(map[string]struct{}, len(ids))
|
||||
for _, id := range ids {
|
||||
trimmed := strings.TrimSpace(id)
|
||||
if trimmed == "" {
|
||||
continue
|
||||
}
|
||||
idFilter[trimmed] = struct{}{}
|
||||
}
|
||||
filtered := make([]*agentv1.TodoItem, 0, len(items))
|
||||
for _, item := range items {
|
||||
if item == nil {
|
||||
continue
|
||||
}
|
||||
if len(statusFilter) > 0 {
|
||||
if _, ok := statusFilter[item.GetStatus()]; !ok {
|
||||
continue
|
||||
}
|
||||
}
|
||||
if len(idFilter) > 0 {
|
||||
if _, ok := idFilter[strings.TrimSpace(item.GetId())]; !ok {
|
||||
continue
|
||||
}
|
||||
}
|
||||
filtered = append(filtered, cloneTodoItem(item))
|
||||
}
|
||||
return filtered
|
||||
}
|
||||
|
||||
func renderTodoList(items []*agentv1.TodoItem) string {
|
||||
if len(items) == 0 {
|
||||
return ""
|
||||
}
|
||||
lines := make([]string, 0, len(items))
|
||||
for _, item := range items {
|
||||
if item == nil {
|
||||
continue
|
||||
}
|
||||
lines = append(lines, fmt.Sprintf("- [%s] %s: %s", todoStatusLabel(item.GetStatus()), strings.TrimSpace(item.GetId()), strings.TrimSpace(item.GetContent())))
|
||||
}
|
||||
return strings.Join(lines, "\n")
|
||||
}
|
||||
|
||||
func buildUpdateTodosToolCall(args *agentv1.UpdateTodosArgs, result *agentv1.UpdateTodosResult) *agentv1.ToolCall {
|
||||
return &agentv1.ToolCall{
|
||||
Tool: &agentv1.ToolCall_UpdateTodosToolCall{
|
||||
UpdateTodosToolCall: &agentv1.UpdateTodosToolCall{
|
||||
Args: args,
|
||||
Result: result,
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func buildUpdateTodosErrorResult(args *agentv1.UpdateTodosArgs, message string) *agentv1.UpdateTodosResult {
|
||||
return &agentv1.UpdateTodosResult{
|
||||
Result: &agentv1.UpdateTodosResult_Error{
|
||||
Error: &agentv1.UpdateTodosError{Error: strings.TrimSpace(message)},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func buildReadTodosToolCall(args *agentv1.ReadTodosArgs, result *agentv1.ReadTodosResult) *agentv1.ToolCall {
|
||||
return &agentv1.ToolCall{
|
||||
Tool: &agentv1.ToolCall_ReadTodosToolCall{
|
||||
ReadTodosToolCall: &agentv1.ReadTodosToolCall{
|
||||
Args: args,
|
||||
Result: result,
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func summarizeUpdateTodosResult(result *agentv1.UpdateTodosResult) string {
|
||||
if result == nil {
|
||||
return "todo update missing"
|
||||
}
|
||||
switch item := result.GetResult().(type) {
|
||||
case *agentv1.UpdateTodosResult_Success:
|
||||
return fmt.Sprintf("todo update success count=%d merge=%t", item.Success.GetTotalCount(), item.Success.GetWasMerge())
|
||||
case *agentv1.UpdateTodosResult_Error:
|
||||
return strings.TrimSpace(item.Error.GetError())
|
||||
default:
|
||||
return "todo update finished"
|
||||
}
|
||||
}
|
||||
|
||||
func summarizeReadTodosResult(result *agentv1.ReadTodosResult) string {
|
||||
if result == nil {
|
||||
return "todo read missing"
|
||||
}
|
||||
switch item := result.GetResult().(type) {
|
||||
case *agentv1.ReadTodosResult_Success:
|
||||
return fmt.Sprintf("todo read success count=%d", item.Success.GetTotalCount())
|
||||
case *agentv1.ReadTodosResult_Error:
|
||||
return strings.TrimSpace(item.Error.GetError())
|
||||
default:
|
||||
return "todo read finished"
|
||||
}
|
||||
}
|
||||
|
||||
func buildTaskArgsFromJSON(raw []byte) *agentv1.TaskArgs {
|
||||
payload, err := decodeJSONObject(raw)
|
||||
if err != nil {
|
||||
return &agentv1.TaskArgs{}
|
||||
}
|
||||
return buildTaskArgsFromMap(payload)
|
||||
}
|
||||
|
||||
func buildTaskArgsFromMap(payload map[string]any) *agentv1.TaskArgs {
|
||||
args := &agentv1.TaskArgs{
|
||||
Description: strings.TrimSpace(stringValue(valueByAlias(payload, "description"))),
|
||||
Prompt: strings.TrimSpace(stringValue(valueByAlias(payload, "prompt"))),
|
||||
SubagentType: subagentTypeFromString(
|
||||
strings.TrimSpace(stringValue(valueByAlias(payload, "subagent_type", "subagentType"))),
|
||||
),
|
||||
Attachments: stringSliceValue(valueByAlias(payload, "attachments")),
|
||||
Mode: taskModeFromReadonly(boolValue(valueByAlias(payload, "readonly", "readOnly"))),
|
||||
}
|
||||
if model := strings.TrimSpace(stringValue(valueByAlias(payload, "model"))); model != "" {
|
||||
args.Model = &model
|
||||
}
|
||||
if resume := strings.TrimSpace(stringValue(valueByAlias(payload, "resume"))); resume != "" {
|
||||
args.Resume = &resume
|
||||
}
|
||||
if agentID := strings.TrimSpace(stringValue(valueByAlias(payload, "agent_id", "agentId"))); agentID != "" {
|
||||
args.AgentId = &agentID
|
||||
}
|
||||
return args
|
||||
}
|
||||
|
||||
func subagentTypeFromString(raw string) *agentv1.SubagentType {
|
||||
switch strings.TrimSpace(raw) {
|
||||
case "explore":
|
||||
return &agentv1.SubagentType{Type: &agentv1.SubagentType_Explore{Explore: &agentv1.SubagentTypeExplore{}}}
|
||||
case "browser-use", "browserUse":
|
||||
return &agentv1.SubagentType{Type: &agentv1.SubagentType_BrowserUse{BrowserUse: &agentv1.SubagentTypeBrowserUse{}}}
|
||||
case "shell":
|
||||
return &agentv1.SubagentType{Type: &agentv1.SubagentType_Shell{Shell: &agentv1.SubagentTypeShell{}}}
|
||||
case "":
|
||||
return &agentv1.SubagentType{Type: &agentv1.SubagentType_Unspecified{Unspecified: &agentv1.SubagentTypeUnspecified{}}}
|
||||
default:
|
||||
return &agentv1.SubagentType{
|
||||
Type: &agentv1.SubagentType_Custom{
|
||||
Custom: &agentv1.SubagentTypeCustom{Name: strings.TrimSpace(raw)},
|
||||
},
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func taskModeFromReadonly(readonly bool) agentv1.TaskMode {
|
||||
if readonly {
|
||||
return agentv1.TaskMode_TASK_MODE_PLAN
|
||||
}
|
||||
return agentv1.TaskMode_TASK_MODE_AGENT
|
||||
}
|
||||
|
||||
func cloneTodoItems(items []*agentv1.TodoItem) []*agentv1.TodoItem {
|
||||
if len(items) == 0 {
|
||||
return nil
|
||||
}
|
||||
cloned := make([]*agentv1.TodoItem, 0, len(items))
|
||||
for _, item := range items {
|
||||
if item == nil {
|
||||
continue
|
||||
}
|
||||
cloned = append(cloned, cloneTodoItem(item))
|
||||
}
|
||||
return cloned
|
||||
}
|
||||
|
||||
func cloneTodoItem(item *agentv1.TodoItem) *agentv1.TodoItem {
|
||||
if item == nil {
|
||||
return nil
|
||||
}
|
||||
return &agentv1.TodoItem{
|
||||
Id: item.GetId(),
|
||||
Content: item.GetContent(),
|
||||
Status: item.GetStatus(),
|
||||
CreatedAt: item.GetCreatedAt(),
|
||||
UpdatedAt: item.GetUpdatedAt(),
|
||||
Dependencies: append([]string(nil), item.GetDependencies()...),
|
||||
}
|
||||
}
|
||||
|
||||
func timestampMillis(baseTime time.Time) int64 {
|
||||
now := baseTime
|
||||
if now.IsZero() {
|
||||
now = time.Now().UTC()
|
||||
}
|
||||
return now.UTC().UnixMilli()
|
||||
}
|
||||
@@ -0,0 +1,205 @@
|
||||
package forwarder
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"strings"
|
||||
"unicode/utf8"
|
||||
|
||||
"google.golang.org/protobuf/encoding/protojson"
|
||||
|
||||
"cursor/gen/agentv1"
|
||||
"cursor/gen/aiserverv1"
|
||||
modeladapter "cursor/internal/backend/agent/model"
|
||||
promptengine "cursor/internal/backend/agent/prompt"
|
||||
)
|
||||
|
||||
const (
|
||||
estimatedTokensPerMessageOverhead = int64(8)
|
||||
estimatedTokensPerToolCallOverhead = int64(6)
|
||||
estimatedTokensPerImagePart = int64(1024)
|
||||
)
|
||||
|
||||
func estimateCompiledPromptTokens(compiled CompiledConversation) int64 {
|
||||
return estimateModelMessagesTokens(compiled.Messages) + estimateToolDescriptorsTokens(compiled.Tools)
|
||||
}
|
||||
|
||||
func estimateModelMessagesTokens(messages []modeladapter.Message) int64 {
|
||||
total := int64(0)
|
||||
for _, item := range messages {
|
||||
total += estimateModelMessageTokens(item)
|
||||
}
|
||||
return total
|
||||
}
|
||||
|
||||
func estimateModelMessageTokens(item modeladapter.Message) int64 {
|
||||
total := estimatedTokensPerMessageOverhead
|
||||
total += estimateTextTokens(item.Role)
|
||||
total += estimateTextTokens(item.Content)
|
||||
total += estimateModelContentPartsTokens(item.Content, item.ContentParts)
|
||||
total += estimateTextTokens(item.ReasoningContent)
|
||||
total += estimateTextTokens(item.ReasoningSignature)
|
||||
total += estimateTextTokens(item.ToolCallID)
|
||||
total += estimateTextTokens(item.Name)
|
||||
for _, toolCall := range item.ToolCalls {
|
||||
total += estimatedTokensPerToolCallOverhead
|
||||
total += estimateTextTokens(toolCall.ID)
|
||||
total += estimateTextTokens(toolCall.Type)
|
||||
total += estimateTextTokens(toolCall.Function.Name)
|
||||
total += estimateTextTokens(toolCall.Function.Arguments)
|
||||
}
|
||||
return total
|
||||
}
|
||||
|
||||
func estimateToolDescriptorsTokens(tools []json.RawMessage) int64 {
|
||||
total := int64(0)
|
||||
for _, item := range tools {
|
||||
total += estimateTextTokens(string(item))
|
||||
}
|
||||
return total
|
||||
}
|
||||
|
||||
func estimateContextItemTokens(item *aiserverv1.ContextItem) int64 {
|
||||
if item == nil {
|
||||
return 0
|
||||
}
|
||||
body, err := protojson.MarshalOptions{EmitUnpopulated: false}.Marshal(item)
|
||||
if err != nil {
|
||||
return estimateTextTokens(item.String())
|
||||
}
|
||||
return estimateTextTokens(string(body))
|
||||
}
|
||||
|
||||
func estimateTextTokens(text string) int64 {
|
||||
trimmed := strings.TrimSpace(text)
|
||||
if trimmed == "" {
|
||||
return 0
|
||||
}
|
||||
runeCount := utf8.RuneCountInString(trimmed)
|
||||
if runeCount <= 0 {
|
||||
return 0
|
||||
}
|
||||
estimated := int64((runeCount + 3) / 4)
|
||||
estimated += int64(strings.Count(trimmed, "\n"))
|
||||
if estimated < 1 {
|
||||
return 1
|
||||
}
|
||||
return estimated
|
||||
}
|
||||
|
||||
func estimateModelContentPartsTokens(content string, parts []modeladapter.ContentPart) int64 {
|
||||
if len(parts) == 0 {
|
||||
return 0
|
||||
}
|
||||
total := int64(0)
|
||||
countText := strings.TrimSpace(content) == ""
|
||||
for _, part := range parts {
|
||||
switch strings.TrimSpace(strings.ToLower(part.Type)) {
|
||||
case "", "text":
|
||||
if countText {
|
||||
total += estimateTextTokens(part.Text)
|
||||
}
|
||||
case "image":
|
||||
total += estimatedTokensPerImagePart
|
||||
if part.Image != nil {
|
||||
total += estimateTextTokens(part.Image.MIMEType)
|
||||
total += estimateTextTokens(part.Image.Path)
|
||||
}
|
||||
}
|
||||
}
|
||||
return total
|
||||
}
|
||||
|
||||
func estimatePromptContentPartsTokens(content string, parts []promptengine.ContentPart) int64 {
|
||||
if len(parts) == 0 {
|
||||
return 0
|
||||
}
|
||||
total := int64(0)
|
||||
countText := strings.TrimSpace(content) == ""
|
||||
for _, part := range parts {
|
||||
switch strings.TrimSpace(strings.ToLower(part.Type)) {
|
||||
case "", "text":
|
||||
if countText {
|
||||
total += estimateTextTokens(part.Text)
|
||||
}
|
||||
case "image":
|
||||
total += estimatedTokensPerImagePart
|
||||
if part.Image != nil {
|
||||
total += estimateTextTokens(part.Image.MIMEType)
|
||||
total += estimateTextTokens(part.Image.Path)
|
||||
}
|
||||
}
|
||||
}
|
||||
return total
|
||||
}
|
||||
|
||||
func estimateCheckpointPromptTokenBreakdown(compiled CompiledConversation, hasCompiled bool, usedTokens uint32, maxTokens uint32) *agentv1.PromptTokenBreakdownSnapshot {
|
||||
if maxTokens == 0 {
|
||||
return nil
|
||||
}
|
||||
categories := make([]*agentv1.PromptTokenBreakdownCategory, 0, 4)
|
||||
if hasCompiled {
|
||||
systemTokens, summaryTokens, conversationTokens := estimateMessageBreakdownTokens(compiled.Messages)
|
||||
categories = appendPromptTokenBreakdownCategory(categories, "system_prompt", "System Prompt", systemTokens)
|
||||
categories = appendPromptTokenBreakdownCategory(categories, "tools", "Tools", estimateToolDescriptorsTokens(compiled.Tools))
|
||||
categories = appendPromptTokenBreakdownCategory(categories, "summarized_conversation", "Summarized Conversation", summaryTokens)
|
||||
categories = appendPromptTokenBreakdownCategory(categories, "conversation", "Conversation", conversationTokens)
|
||||
} else if usedTokens > 0 {
|
||||
categories = appendPromptTokenBreakdownCategory(categories, "conversation", "Conversation", int64(usedTokens))
|
||||
}
|
||||
categoryTotal := int64(0)
|
||||
for _, category := range categories {
|
||||
categoryTotal += int64(category.GetEstimatedTokens())
|
||||
}
|
||||
totalUsedTokens := usedTokens
|
||||
if categoryTotal > int64(totalUsedTokens) {
|
||||
totalUsedTokens = clampInt64ToUint32(categoryTotal)
|
||||
}
|
||||
return &agentv1.PromptTokenBreakdownSnapshot{
|
||||
TotalUsedTokens: totalUsedTokens,
|
||||
MaxTokens: maxTokens,
|
||||
Categories: categories,
|
||||
}
|
||||
}
|
||||
|
||||
func estimateMessageBreakdownTokens(messages []modeladapter.Message) (int64, int64, int64) {
|
||||
systemTokens := int64(0)
|
||||
summaryTokens := int64(0)
|
||||
conversationTokens := int64(0)
|
||||
for _, message := range messages {
|
||||
tokens := estimateModelMessageTokens(message)
|
||||
switch {
|
||||
case strings.TrimSpace(message.Role) == "system":
|
||||
systemTokens += tokens
|
||||
case isConversationSummaryMessage(message):
|
||||
summaryTokens += tokens
|
||||
default:
|
||||
conversationTokens += tokens
|
||||
}
|
||||
}
|
||||
return systemTokens, summaryTokens, conversationTokens
|
||||
}
|
||||
|
||||
func isConversationSummaryMessage(message modeladapter.Message) bool {
|
||||
return strings.Contains(message.Content, "<conversation_summary>") || strings.Contains(message.Content, "</conversation_summary>")
|
||||
}
|
||||
|
||||
func appendPromptTokenBreakdownCategory(categories []*agentv1.PromptTokenBreakdownCategory, id string, label string, estimatedTokens int64) []*agentv1.PromptTokenBreakdownCategory {
|
||||
if estimatedTokens <= 0 {
|
||||
return categories
|
||||
}
|
||||
return append(categories, &agentv1.PromptTokenBreakdownCategory{
|
||||
Id: id,
|
||||
Label: label,
|
||||
EstimatedTokens: clampInt64ToUint32(estimatedTokens),
|
||||
})
|
||||
}
|
||||
|
||||
func clampInt64ToInt32(value int64) int32 {
|
||||
if value <= 0 {
|
||||
return 0
|
||||
}
|
||||
if value > int64(^uint32(0)>>1) {
|
||||
return int32(^uint32(0) >> 1)
|
||||
}
|
||||
return int32(value)
|
||||
}
|
||||
@@ -0,0 +1,459 @@
|
||||
package forwarder
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"cursor/gen/agentv1"
|
||||
modeladapter "cursor/internal/backend/agent/model"
|
||||
promptengine "cursor/internal/backend/agent/prompt"
|
||||
"google.golang.org/protobuf/proto"
|
||||
)
|
||||
|
||||
type turnUsageSnapshot struct {
|
||||
Provider string
|
||||
Model string
|
||||
InputTokens int64
|
||||
OutputTokens int64
|
||||
CacheReadTokens int64
|
||||
CacheWriteTokens int64
|
||||
UsagePresent bool
|
||||
CacheReadPresent bool
|
||||
CacheWritePresent bool
|
||||
}
|
||||
|
||||
func (snapshot turnUsageSnapshot) hasAny() bool {
|
||||
return snapshot.InputTokens > 0 ||
|
||||
snapshot.OutputTokens > 0 ||
|
||||
snapshot.CacheReadTokens > 0 ||
|
||||
snapshot.CacheWriteTokens > 0
|
||||
}
|
||||
|
||||
func (snapshot turnUsageSnapshot) cacheUsageComplete() bool {
|
||||
return snapshot.CacheReadPresent && snapshot.CacheWritePresent
|
||||
}
|
||||
|
||||
func (snapshot turnUsageSnapshot) promptTokensTotal() int64 {
|
||||
return nonNegativeInt64(snapshot.InputTokens) +
|
||||
nonNegativeInt64(snapshot.CacheReadTokens) +
|
||||
nonNegativeInt64(snapshot.CacheWriteTokens)
|
||||
}
|
||||
|
||||
func (snapshot turnUsageSnapshot) requestTokensTotal() int64 {
|
||||
return snapshot.promptTokensTotal() + nonNegativeInt64(snapshot.OutputTokens)
|
||||
}
|
||||
|
||||
func (service *Service) importConversationState(item *ConversationFile, state *agentv1.ConversationStateStructure) ([]HistoryEntry, error) {
|
||||
if item == nil || state == nil {
|
||||
return nil, nil
|
||||
}
|
||||
item.TokenDetailsUsedTokens = state.GetTokenDetails().GetUsedTokens()
|
||||
entries := make([]HistoryEntry, 0, 2)
|
||||
if messages, err := importedConversationStateModelMessages(state); err != nil {
|
||||
return nil, err
|
||||
} else {
|
||||
for _, message := range messages {
|
||||
entry, ok, err := newModelMessageEntry(0, "", message)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if ok {
|
||||
entries = append(entries, entry)
|
||||
}
|
||||
}
|
||||
}
|
||||
if len(entries) == 0 {
|
||||
summary, ok, err := importedConversationStateSummary(state)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if ok {
|
||||
payload, err := json.Marshal(compactionSummaryEntryPayload{
|
||||
Summary: strings.TrimSpace(summary),
|
||||
Trigger: "imported_conversation_state",
|
||||
})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("encode imported summary context: %w", err)
|
||||
}
|
||||
entries = append(entries, HistoryEntry{
|
||||
TurnSeq: 0,
|
||||
Role: "system",
|
||||
Kind: "compacted_summary",
|
||||
Payload: payload,
|
||||
})
|
||||
}
|
||||
}
|
||||
runtimeState, ok, err := runtimeStatePayloadFromConversationState(state)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if ok {
|
||||
payload, err := json.Marshal(runtimeState)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("encode imported runtime state context: %w", err)
|
||||
}
|
||||
entries = append(entries, HistoryEntry{
|
||||
TurnSeq: 0,
|
||||
Role: "system",
|
||||
Kind: "runtime_state",
|
||||
Payload: payload,
|
||||
})
|
||||
}
|
||||
return entries, nil
|
||||
}
|
||||
|
||||
func importedConversationStateModelMessages(state *agentv1.ConversationStateStructure) ([]modeladapter.Message, error) {
|
||||
if state == nil {
|
||||
return nil, nil
|
||||
}
|
||||
if len(state.GetRootPromptMessagesJson()) > 0 {
|
||||
decoded, err := promptengine.DecodeReplayMessages(state.GetRootPromptMessagesJson())
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("decode imported replay messages: %w", err)
|
||||
}
|
||||
decoded = restoreImportedReplayUserMessages(decoded, state.GetTurns())
|
||||
decoded = filterLegacyPlainWriteReplay(decoded)
|
||||
decoded = filterInternalPromptContextReplay(decoded)
|
||||
messages := make([]modeladapter.Message, 0, len(decoded))
|
||||
for _, item := range decoded {
|
||||
messages = append(messages, toModelMessage(item))
|
||||
}
|
||||
return normalizeReplayMessageSequence(messages), nil
|
||||
}
|
||||
if len(state.GetSummary()) > 0 {
|
||||
return nil, nil
|
||||
}
|
||||
if len(state.GetTurns()) == 0 {
|
||||
return nil, nil
|
||||
}
|
||||
messages := make([]modeladapter.Message, 0, len(state.GetTurns())*2)
|
||||
for _, rawTurn := range state.GetTurns() {
|
||||
if len(rawTurn) == 0 {
|
||||
continue
|
||||
}
|
||||
turn := &agentv1.ConversationTurnStructure{}
|
||||
if err := proto.Unmarshal(rawTurn, turn); err != nil {
|
||||
return nil, fmt.Errorf("decode imported turn: %w", err)
|
||||
}
|
||||
agentTurn := turn.GetAgentConversationTurn()
|
||||
if agentTurn == nil {
|
||||
continue
|
||||
}
|
||||
if rawUser := agentTurn.GetUserMessage(); len(rawUser) > 0 {
|
||||
userMessage := &agentv1.UserMessage{}
|
||||
if err := proto.Unmarshal(rawUser, userMessage); err != nil {
|
||||
return nil, fmt.Errorf("decode imported turn user_message: %w", err)
|
||||
}
|
||||
if replay, ok := promptengine.BuildUserMessageReplayMessage(userMessage); ok {
|
||||
messages = append(messages, toModelMessage(replay))
|
||||
}
|
||||
}
|
||||
for _, rawStep := range agentTurn.GetSteps() {
|
||||
if len(rawStep) == 0 {
|
||||
continue
|
||||
}
|
||||
step := &agentv1.ConversationStep{}
|
||||
if err := proto.Unmarshal(rawStep, step); err != nil {
|
||||
return nil, fmt.Errorf("decode imported turn step: %w", err)
|
||||
}
|
||||
for _, replay := range promptengine.BuildLegacyMessagesFromConversationStep(step) {
|
||||
messages = append(messages, toModelMessage(replay))
|
||||
}
|
||||
}
|
||||
}
|
||||
return normalizeReplayMessageSequence(messages), nil
|
||||
}
|
||||
|
||||
func importedConversationStateSummary(state *agentv1.ConversationStateStructure) (string, bool, error) {
|
||||
if state == nil || len(state.GetSummary()) == 0 {
|
||||
return "", false, nil
|
||||
}
|
||||
item := &agentv1.ConversationSummary{}
|
||||
if err := proto.Unmarshal(state.GetSummary(), item); err != nil {
|
||||
return "", false, fmt.Errorf("decode imported summary: %w", err)
|
||||
}
|
||||
text := strings.TrimSpace(item.GetSummary())
|
||||
return text, text != "", nil
|
||||
}
|
||||
|
||||
func newModelMessageEntry(turnSeq int64, requestID string, message modeladapter.Message) (HistoryEntry, bool, error) {
|
||||
message.Role = strings.TrimSpace(message.Role)
|
||||
if message.Role == "" {
|
||||
return HistoryEntry{}, false, nil
|
||||
}
|
||||
if strings.TrimSpace(message.Content) == "" &&
|
||||
len(message.ContentParts) == 0 &&
|
||||
len(message.ToolCalls) == 0 &&
|
||||
strings.TrimSpace(message.ToolCallID) == "" &&
|
||||
!hasReplayableReasoningPayload(message.ReasoningContent, message.ReasoningSignature, message.ReasoningSignatureSource) &&
|
||||
len(message.OpenAIResponsesReasoningSummary) == 0 {
|
||||
return HistoryEntry{}, false, nil
|
||||
}
|
||||
payload, err := json.Marshal(modelMessageEntryPayload{Message: message})
|
||||
if err != nil {
|
||||
return HistoryEntry{}, false, fmt.Errorf("encode imported model message context: %w", err)
|
||||
}
|
||||
return HistoryEntry{
|
||||
TurnSeq: turnSeq,
|
||||
RequestID: strings.TrimSpace(requestID),
|
||||
Role: message.Role,
|
||||
Kind: "model_message",
|
||||
Payload: payload,
|
||||
}, true, nil
|
||||
}
|
||||
|
||||
func runtimeStatePayloadFromConversationState(state *agentv1.ConversationStateStructure) (runtimeStateEntryPayload, bool, error) {
|
||||
if state == nil {
|
||||
return runtimeStateEntryPayload{}, false, nil
|
||||
}
|
||||
payload := runtimeStateEntryPayload{
|
||||
Plans: clonePlanRegistryEntries(state.GetPlans()),
|
||||
}
|
||||
if len(state.GetPlan()) > 0 {
|
||||
plan := &agentv1.ConversationPlan{}
|
||||
if err := proto.Unmarshal(state.GetPlan(), plan); err != nil {
|
||||
return runtimeStateEntryPayload{}, false, fmt.Errorf("decode imported plan: %w", err)
|
||||
}
|
||||
payload.PlanText = strings.TrimSpace(plan.GetPlan())
|
||||
}
|
||||
if len(state.GetTodos()) > 0 {
|
||||
todos := make([]*agentv1.TodoItem, 0, len(state.GetTodos()))
|
||||
for _, raw := range state.GetTodos() {
|
||||
if len(raw) == 0 {
|
||||
continue
|
||||
}
|
||||
item := &agentv1.TodoItem{}
|
||||
if err := proto.Unmarshal(raw, item); err != nil {
|
||||
return runtimeStateEntryPayload{}, false, fmt.Errorf("decode imported todo: %w", err)
|
||||
}
|
||||
todos = append(todos, cloneTodoItem(item))
|
||||
}
|
||||
payload.Todos = todos
|
||||
}
|
||||
ok := strings.TrimSpace(payload.PlanText) != "" || len(payload.Plans) > 0 || len(payload.Todos) > 0
|
||||
return payload, ok, nil
|
||||
}
|
||||
|
||||
func (service *Service) updateConversationTokenState(stream *ActiveStream, conversationID string, usage turnUsageSnapshot, modelCallID string, finalizeAutoCompaction bool) error {
|
||||
if service == nil || !usage.hasAny() {
|
||||
return nil
|
||||
}
|
||||
now := time.Now().UTC()
|
||||
autoCompactionReserveTokens := int64(compactionAutoReserveTokens)
|
||||
if finalizeAutoCompaction {
|
||||
autoCompactionReserveTokens = service.resolveCompactionReserveTokens(activeStreamModelID(stream))
|
||||
if autoCompactionReserveTokens <= 0 {
|
||||
autoCompactionReserveTokens = compactionAutoReserveTokens
|
||||
}
|
||||
}
|
||||
_, err := service.updateConversationMetaAndCheckpoint(stream, conversationID, func(item *ConversationFile) error {
|
||||
if item == nil {
|
||||
return nil
|
||||
}
|
||||
promptTokensTotal := usage.promptTokensTotal()
|
||||
if promptTokensTotal > 0 {
|
||||
item.TokenDetailsUsedTokens = clampInt64ToUint32(promptTokensTotal)
|
||||
}
|
||||
if item.TokenDetailsMaxTokens == 0 {
|
||||
item.TokenDetailsMaxTokens = projectedConversationMaxTokens
|
||||
}
|
||||
if finalizeAutoCompaction {
|
||||
updateConversationAutoCompactionState(item, usage.requestTokensTotal(), autoCompactionReserveTokens, modelCallID, now)
|
||||
}
|
||||
return nil
|
||||
})
|
||||
return err
|
||||
}
|
||||
|
||||
func activeStreamModelID(stream *ActiveStream) string {
|
||||
if stream == nil {
|
||||
return ""
|
||||
}
|
||||
stream.mu.Lock()
|
||||
defer stream.mu.Unlock()
|
||||
return stream.ModelID
|
||||
}
|
||||
|
||||
func updateConversationAutoCompactionState(conversation *ConversationFile, promptTokensTotal int64, reserveTokens int64, modelCallID string, triggeredAt time.Time) {
|
||||
if conversation == nil {
|
||||
return
|
||||
}
|
||||
if promptTokensTotal <= 0 {
|
||||
clearConversationAutoCompactionState(conversation)
|
||||
return
|
||||
}
|
||||
contextWindowTokens := int64(conversation.TokenDetailsMaxTokens)
|
||||
if contextWindowTokens <= 0 {
|
||||
contextWindowTokens = projectedConversationMaxTokens
|
||||
}
|
||||
if reserveTokens <= 0 {
|
||||
reserveTokens = compactionAutoReserveTokens
|
||||
}
|
||||
remainingTokens := contextWindowTokens - promptTokensTotal
|
||||
if remainingTokens > reserveTokens {
|
||||
clearConversationAutoCompactionState(conversation)
|
||||
return
|
||||
}
|
||||
conversation.AutoCompactionPending = true
|
||||
conversation.AutoCompactionPromptTokens = promptTokensTotal
|
||||
conversation.AutoCompactionReserveTokens = reserveTokens
|
||||
conversation.AutoCompactionTriggeredAt = triggeredAt.UTC().Format(time.RFC3339Nano)
|
||||
conversation.AutoCompactionSourceModelCallID = strings.TrimSpace(modelCallID)
|
||||
}
|
||||
|
||||
func clearConversationAutoCompactionState(conversation *ConversationFile) {
|
||||
if conversation == nil {
|
||||
return
|
||||
}
|
||||
conversation.AutoCompactionPending = false
|
||||
conversation.AutoCompactionPromptTokens = 0
|
||||
conversation.AutoCompactionReserveTokens = 0
|
||||
conversation.AutoCompactionTriggeredAt = ""
|
||||
conversation.AutoCompactionSourceModelCallID = ""
|
||||
}
|
||||
|
||||
func (service *Service) recordTurnUsageSnapshot(stream *ActiveStream, conversationID string, turnSeq int64, requestID string, modelCallID string, status string, usage turnUsageSnapshot, errorText string, turnFinalized bool) error {
|
||||
_ = turnFinalized
|
||||
if service == nil || strings.TrimSpace(requestID) == "" {
|
||||
return nil
|
||||
}
|
||||
modelID := ""
|
||||
modelName := ""
|
||||
startedAt := time.Time{}
|
||||
lastEventAt := time.Now().UTC()
|
||||
if stream != nil {
|
||||
stream.mu.Lock()
|
||||
modelID = strings.TrimSpace(stream.ModelID)
|
||||
modelName = strings.TrimSpace(stream.ModelName)
|
||||
startedAt = stream.CreatedAt
|
||||
if !stream.UpdatedAt.IsZero() {
|
||||
lastEventAt = stream.UpdatedAt
|
||||
}
|
||||
stream.mu.Unlock()
|
||||
}
|
||||
if strings.TrimSpace(modelName) == "" {
|
||||
modelName = modelID
|
||||
}
|
||||
provider := strings.TrimSpace(usage.Provider)
|
||||
if strings.TrimSpace(usage.Model) != "" {
|
||||
modelName = strings.TrimSpace(usage.Model)
|
||||
}
|
||||
effectiveModelCallID := firstNonEmpty(strings.TrimSpace(modelCallID), strings.TrimSpace(requestID))
|
||||
if service.usageStore != nil {
|
||||
if err := service.usageStore.UpsertEvent(usageFileEvent{
|
||||
EventID: usageEventID(requestID, effectiveModelCallID),
|
||||
Kind: usageEventKindProvider,
|
||||
At: lastEventAt,
|
||||
InputTokens: usage.InputTokens,
|
||||
OutputTokens: usage.OutputTokens,
|
||||
CacheReadTokens: usage.CacheReadTokens,
|
||||
CacheWriteTokens: usage.CacheWriteTokens,
|
||||
UsagePresent: usage.UsagePresent,
|
||||
}); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
if strings.TrimSpace(conversationID) != "" {
|
||||
_, err := service.updateConversationMetaAndCheckpoint(stream, conversationID, func(item *ConversationFile) error {
|
||||
if item == nil {
|
||||
return nil
|
||||
}
|
||||
item.LastProviderCall = &ConversationProviderCall{
|
||||
RequestID: strings.TrimSpace(requestID),
|
||||
ModelCallID: effectiveModelCallID,
|
||||
Provider: provider,
|
||||
Model: modelName,
|
||||
Status: strings.TrimSpace(status),
|
||||
ErrorText: strings.TrimSpace(errorText),
|
||||
UpdatedAt: lastEventAt,
|
||||
}
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
_ = turnSeq
|
||||
_ = startedAt
|
||||
return nil
|
||||
}
|
||||
|
||||
func (service *Service) recordTurnFinalizedSnapshot(stream *ActiveStream, conversationID string, turnSeq int64, requestID string, status string, errorText string) error {
|
||||
_ = stream
|
||||
_ = errorText
|
||||
if service == nil || service.usageStore == nil || strings.TrimSpace(requestID) == "" {
|
||||
return nil
|
||||
}
|
||||
eventID := turnUsageEventID(conversationID, turnSeq, requestID)
|
||||
usage := usageFileEvent{}
|
||||
if aggregate, ok, err := service.usageStore.LookupEvent(strings.TrimSpace(requestID)); err != nil {
|
||||
return err
|
||||
} else if ok {
|
||||
usage = aggregate
|
||||
}
|
||||
return service.usageStore.UpsertEvent(usageFileEvent{
|
||||
EventID: eventID,
|
||||
Kind: usageEventKindTurn,
|
||||
Status: normalizeUsageTurnStatus(status),
|
||||
At: time.Now().UTC(),
|
||||
InputTokens: usage.InputTokens,
|
||||
OutputTokens: usage.OutputTokens,
|
||||
CacheReadTokens: usage.CacheReadTokens,
|
||||
CacheWriteTokens: usage.CacheWriteTokens,
|
||||
UsagePresent: usage.UsagePresent,
|
||||
})
|
||||
}
|
||||
|
||||
func normalizeUsageTurnStatus(status string) string {
|
||||
if strings.TrimSpace(status) == usageTurnStatusDone {
|
||||
return usageTurnStatusDone
|
||||
}
|
||||
return strings.TrimSpace(status)
|
||||
}
|
||||
|
||||
func turnUsageEventID(conversationID string, turnSeq int64, requestID string) string {
|
||||
conversationID = strings.TrimSpace(conversationID)
|
||||
requestID = strings.TrimSpace(requestID)
|
||||
if conversationID != "" && turnSeq > 0 {
|
||||
return fmt.Sprintf("turn::%s::%d", conversationID, turnSeq)
|
||||
}
|
||||
return "turn::" + requestID
|
||||
}
|
||||
|
||||
func usageEventID(requestID string, modelCallID string) string {
|
||||
requestID = strings.TrimSpace(requestID)
|
||||
modelCallID = strings.TrimSpace(modelCallID)
|
||||
if modelCallID == "" || modelCallID == requestID {
|
||||
return requestID
|
||||
}
|
||||
return requestID + "::" + modelCallID
|
||||
}
|
||||
|
||||
func clampInt64ToUint32(value int64) uint32 {
|
||||
if value <= 0 {
|
||||
return 0
|
||||
}
|
||||
if value > int64(^uint32(0)) {
|
||||
return ^uint32(0)
|
||||
}
|
||||
return uint32(value)
|
||||
}
|
||||
|
||||
func nonNegativeInt64(value int64) int64 {
|
||||
if value < 0 {
|
||||
return 0
|
||||
}
|
||||
return value
|
||||
}
|
||||
|
||||
func maxPositiveInt64(values ...int64) int64 {
|
||||
maxValue := int64(0)
|
||||
for _, value := range values {
|
||||
if value > maxValue {
|
||||
maxValue = value
|
||||
}
|
||||
}
|
||||
return maxValue
|
||||
}
|
||||
@@ -0,0 +1,295 @@
|
||||
// tool_catalog.go 负责从静态 prompt 资产中装载并筛选 canonical tool catalog。
|
||||
package forwarder
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"cursor/gen/agentv1"
|
||||
promptassets "cursor/prompt"
|
||||
)
|
||||
|
||||
type DefaultToolCatalog struct {
|
||||
}
|
||||
|
||||
// NewToolCatalog 创建默认工具目录实现。
|
||||
func NewToolCatalog() *DefaultToolCatalog {
|
||||
return &DefaultToolCatalog{}
|
||||
}
|
||||
|
||||
// Load 按 mode 读取工具资产,并过滤出当前阶段真正允许暴露的工具。
|
||||
func (catalog *DefaultToolCatalog) Load(mode agentv1.AgentMode, subagentTypeName string) ([]json.RawMessage, []string, error) {
|
||||
assetMode, err := toolAssetModeForConversation(mode, subagentTypeName)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
rawTools, err := promptassets.ReadTools(assetMode)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
var items []json.RawMessage
|
||||
if err := json.Unmarshal(rawTools, &items); err != nil {
|
||||
return nil, nil, fmt.Errorf("decode tools asset failed: %w", err)
|
||||
}
|
||||
filtered := make([]json.RawMessage, 0, len(items))
|
||||
names := make([]string, 0, len(items))
|
||||
for _, item := range items {
|
||||
name, err := extractToolName(item)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
if !isToolAllowedInMode(mode, subagentTypeName, name) {
|
||||
continue
|
||||
}
|
||||
filtered = append(filtered, item)
|
||||
names = append(names, name)
|
||||
}
|
||||
return filtered, names, nil
|
||||
}
|
||||
|
||||
var agentModeToolNames = map[string]struct{}{
|
||||
"AskQuestion": {},
|
||||
"CallMcpTool": {},
|
||||
"Delete": {},
|
||||
"FetchMcpResource": {},
|
||||
"GenerateImage": {},
|
||||
"Glob": {},
|
||||
"Grep": {},
|
||||
"Ls": {},
|
||||
"PatchEdit": {},
|
||||
"Read": {},
|
||||
"ReadLints": {},
|
||||
"Shell": {},
|
||||
"AwaitShell": {},
|
||||
"WriteShellStdin": {},
|
||||
"ForceBackgroundShell": {},
|
||||
"SwitchMode": {},
|
||||
"Task": {},
|
||||
"TodoWrite": {},
|
||||
"WebFetch": {},
|
||||
"WebSearch": {},
|
||||
"Write": {},
|
||||
}
|
||||
|
||||
var multitaskModeToolNames = map[string]struct{}{
|
||||
"AskQuestion": {},
|
||||
"CallMcpTool": {},
|
||||
"Delete": {},
|
||||
"FetchMcpResource": {},
|
||||
"GenerateImage": {},
|
||||
"Glob": {},
|
||||
"Grep": {},
|
||||
"Ls": {},
|
||||
"PatchEdit": {},
|
||||
"Read": {},
|
||||
"ReadLints": {},
|
||||
"Shell": {},
|
||||
"AwaitShell": {},
|
||||
"WriteShellStdin": {},
|
||||
"ForceBackgroundShell": {},
|
||||
"SwitchMode": {},
|
||||
"Task": {},
|
||||
"TodoWrite": {},
|
||||
"WebFetch": {},
|
||||
"WebSearch": {},
|
||||
"Write": {},
|
||||
}
|
||||
|
||||
var debugModeToolNames = map[string]struct{}{
|
||||
"AskQuestion": {},
|
||||
"CallMcpTool": {},
|
||||
"Delete": {},
|
||||
"FetchMcpResource": {},
|
||||
"Glob": {},
|
||||
"Grep": {},
|
||||
"Ls": {},
|
||||
"PatchEdit": {},
|
||||
"Read": {},
|
||||
"ReadLints": {},
|
||||
"Shell": {},
|
||||
"AwaitShell": {},
|
||||
"WriteShellStdin": {},
|
||||
"ForceBackgroundShell": {},
|
||||
"Task": {},
|
||||
"TodoWrite": {},
|
||||
"WebFetch": {},
|
||||
"WebSearch": {},
|
||||
"Write": {},
|
||||
}
|
||||
|
||||
var askModeToolNames = map[string]struct{}{
|
||||
"AskQuestion": {},
|
||||
"CallMcpTool": {},
|
||||
"Delete": {},
|
||||
"FetchMcpResource": {},
|
||||
"Glob": {},
|
||||
"Grep": {},
|
||||
"Ls": {},
|
||||
"PatchEdit": {},
|
||||
"Read": {},
|
||||
"ReadLints": {},
|
||||
"Shell": {},
|
||||
"AwaitShell": {},
|
||||
"WriteShellStdin": {},
|
||||
"ForceBackgroundShell": {},
|
||||
"Task": {},
|
||||
"TodoWrite": {},
|
||||
"WebFetch": {},
|
||||
"WebSearch": {},
|
||||
"Write": {},
|
||||
}
|
||||
|
||||
var planModeToolNames = map[string]struct{}{
|
||||
"AskQuestion": {},
|
||||
"CallMcpTool": {},
|
||||
"CreatePlan": {},
|
||||
"FetchMcpResource": {},
|
||||
"Glob": {},
|
||||
"Grep": {},
|
||||
"Ls": {},
|
||||
"Read": {},
|
||||
"ReadLints": {},
|
||||
"Shell": {},
|
||||
"AwaitShell": {},
|
||||
"WriteShellStdin": {},
|
||||
"ForceBackgroundShell": {},
|
||||
"Task": {},
|
||||
"TodoWrite": {},
|
||||
"WebFetch": {},
|
||||
"WebSearch": {},
|
||||
}
|
||||
|
||||
var childConversationDisallowedAgentToolNames = map[string]struct{}{
|
||||
"AskQuestion": {},
|
||||
}
|
||||
|
||||
func supportedToolNamesForMode(mode agentv1.AgentMode) map[string]struct{} {
|
||||
switch normalizeMode(mode) {
|
||||
case agentv1.AgentMode_AGENT_MODE_AGENT:
|
||||
return agentModeToolNames
|
||||
case agentv1.AgentMode_AGENT_MODE_ASK:
|
||||
return askModeToolNames
|
||||
case agentv1.AgentMode_AGENT_MODE_PLAN:
|
||||
return planModeToolNames
|
||||
case agentv1.AgentMode_AGENT_MODE_DEBUG:
|
||||
return debugModeToolNames
|
||||
case agentv1.AgentMode_AGENT_MODE_MULTITASK:
|
||||
return multitaskModeToolNames
|
||||
default:
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
func isToolAllowedInMode(mode agentv1.AgentMode, subagentTypeName string, toolName string) bool {
|
||||
trimmedToolName := strings.TrimSpace(toolName)
|
||||
if trimmedToolName == "" {
|
||||
return false
|
||||
}
|
||||
if isChildConversationSubagentTypeName(subagentTypeName) {
|
||||
if _, disallowed := childConversationDisallowedAgentToolNames[trimmedToolName]; disallowed {
|
||||
return false
|
||||
}
|
||||
_, ok := agentModeToolNames[trimmedToolName]
|
||||
return ok
|
||||
}
|
||||
supported := supportedToolNamesForMode(mode)
|
||||
if supported == nil {
|
||||
return false
|
||||
}
|
||||
_, ok := supported[trimmedToolName]
|
||||
return ok
|
||||
}
|
||||
|
||||
func isChildConversationSubagentTypeName(subagentTypeName string) bool {
|
||||
return strings.TrimSpace(subagentTypeName) != ""
|
||||
}
|
||||
|
||||
func selectToolsByOrderedNames(items []json.RawMessage, orderedNames []string) ([]json.RawMessage, []string, error) {
|
||||
byName := make(map[string]json.RawMessage, len(items))
|
||||
for _, item := range items {
|
||||
name, err := extractToolName(item)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
if _, exists := byName[name]; !exists {
|
||||
byName[name] = item
|
||||
}
|
||||
}
|
||||
filtered := make([]json.RawMessage, 0, len(orderedNames))
|
||||
names := make([]string, 0, len(orderedNames))
|
||||
for _, name := range orderedNames {
|
||||
item, ok := byName[name]
|
||||
if !ok {
|
||||
return nil, nil, fmt.Errorf("tool descriptor %q not found in prompt asset", name)
|
||||
}
|
||||
filtered = append(filtered, item)
|
||||
names = append(names, name)
|
||||
}
|
||||
return filtered, names, nil
|
||||
}
|
||||
|
||||
func toolAssetModeForConversation(mode agentv1.AgentMode, subagentTypeName string) (promptassets.Mode, error) {
|
||||
if isChildConversationSubagentTypeName(subagentTypeName) {
|
||||
return promptassets.ModeAgent, nil
|
||||
}
|
||||
return mapPromptMode(mode)
|
||||
}
|
||||
|
||||
func promptAssetModeForConversation(mode agentv1.AgentMode, subagentTypeName string) (promptassets.Mode, error) {
|
||||
if isChildConversationSubagentTypeName(subagentTypeName) {
|
||||
return promptassets.ModeSubagent, nil
|
||||
}
|
||||
return mapPromptMode(mode)
|
||||
}
|
||||
|
||||
// mapPromptMode 把协议 mode 映射为静态 prompt 资产对应的目录名。
|
||||
func mapPromptMode(mode agentv1.AgentMode) (promptassets.Mode, error) {
|
||||
switch normalizeMode(mode) {
|
||||
case agentv1.AgentMode_AGENT_MODE_AGENT:
|
||||
return promptassets.ModeAgent, nil
|
||||
case agentv1.AgentMode_AGENT_MODE_ASK:
|
||||
return promptassets.ModeAsk, nil
|
||||
case agentv1.AgentMode_AGENT_MODE_PLAN:
|
||||
return promptassets.ModePlan, nil
|
||||
case agentv1.AgentMode_AGENT_MODE_DEBUG:
|
||||
return promptassets.ModeDebug, nil
|
||||
case agentv1.AgentMode_AGENT_MODE_MULTITASK:
|
||||
return promptassets.ModeMultitask, nil
|
||||
default:
|
||||
return "", fmt.Errorf("unsupported prompt asset mode: %s", mode.String())
|
||||
}
|
||||
}
|
||||
|
||||
// extractToolName 从原始 tool descriptor JSON 中提取函数名。
|
||||
func extractToolName(raw json.RawMessage) (string, error) {
|
||||
var wrapper struct {
|
||||
Function struct {
|
||||
Name string `json:"name"`
|
||||
} `json:"function"`
|
||||
}
|
||||
if err := json.Unmarshal(raw, &wrapper); err != nil {
|
||||
return "", fmt.Errorf("decode tool descriptor failed: %w", err)
|
||||
}
|
||||
name := strings.TrimSpace(wrapper.Function.Name)
|
||||
if name == "" {
|
||||
return "", fmt.Errorf("tool descriptor name is required")
|
||||
}
|
||||
return name, nil
|
||||
}
|
||||
|
||||
// sanitizePromptAsset 去掉资产文件中的说明性标题,只保留真正的 prompt 文本。
|
||||
func sanitizePromptAsset(text string, modelName string) string {
|
||||
lines := strings.Split(text, "\n")
|
||||
filtered := make([]string, 0, len(lines))
|
||||
for _, line := range lines {
|
||||
trimmed := strings.TrimSpace(line)
|
||||
switch trimmed {
|
||||
case "# 通用系统提示词", "# 模式静态补充", "---":
|
||||
continue
|
||||
default:
|
||||
filtered = append(filtered, line)
|
||||
}
|
||||
}
|
||||
return promptassets.RenderPromptTemplate(strings.TrimSpace(strings.Join(filtered, "\n")), modelName)
|
||||
}
|
||||
@@ -0,0 +1,111 @@
|
||||
package forwarder
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"google.golang.org/protobuf/encoding/protojson"
|
||||
|
||||
"cursor/gen/agentv1"
|
||||
runtimecore "cursor/internal/backend/agent/core"
|
||||
)
|
||||
|
||||
type recoverableToolInvocationError struct {
|
||||
cause error
|
||||
}
|
||||
|
||||
func (err recoverableToolInvocationError) Error() string {
|
||||
if err.cause == nil {
|
||||
return "recoverable tool invocation error"
|
||||
}
|
||||
return err.cause.Error()
|
||||
}
|
||||
|
||||
func (err recoverableToolInvocationError) Unwrap() error {
|
||||
return err.cause
|
||||
}
|
||||
|
||||
func newRecoverableToolInvocationError(cause error) error {
|
||||
if cause == nil {
|
||||
return nil
|
||||
}
|
||||
return recoverableToolInvocationError{cause: cause}
|
||||
}
|
||||
|
||||
func recoverableToolInvocationCause(err error) (error, bool) {
|
||||
if err == nil {
|
||||
return nil, false
|
||||
}
|
||||
var recoverable recoverableToolInvocationError
|
||||
if !errors.As(err, &recoverable) {
|
||||
return nil, false
|
||||
}
|
||||
if recoverable.cause == nil {
|
||||
return err, true
|
||||
}
|
||||
return recoverable.cause, true
|
||||
}
|
||||
|
||||
func (service *Service) completePreDispatchToolError(
|
||||
stream *ActiveStream,
|
||||
invocation runtimecore.ToolInvocation,
|
||||
startedToolCall *agentv1.ToolCall,
|
||||
startedHistoryAppended bool,
|
||||
startedEmitted bool,
|
||||
cause error,
|
||||
) error {
|
||||
if stream == nil || cause == nil {
|
||||
return cause
|
||||
}
|
||||
if strings.TrimSpace(invocation.CallID) == "" {
|
||||
return cause
|
||||
}
|
||||
if startedToolCall == nil {
|
||||
startedToolCall = buildStartedToolCall(invocation)
|
||||
}
|
||||
if !startedHistoryAppended && startedToolCall != nil {
|
||||
toolCallPayload, err := protojson.Marshal(startedToolCall)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if _, err := service.appendConversationEntries(stream, stream.ConversationID, []HistoryEntry{
|
||||
newToolCallEntryWithProviderMetadata(stream.TurnSeq, stream.RequestID, invocation.CallID, invocation.ToolName, invocation.ReasoningContent, invocation.ReasoningSignature, invocation.ReasoningSignatureSource, invocation.ReasoningProviderItemID, invocation.ReasoningProviderStatus, invocation.ReasoningProviderSummary, invocation.ProviderItemID, invocation.ProviderCallID, invocation.ProviderStatus, toolCallPayload),
|
||||
}); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
if !startedEmitted {
|
||||
if err := service.broker.Publish(stream.RequestID, StreamEvent{
|
||||
Message: buildToolCallStartedMessage(invocation.CallID, invocation.ModelCallID, startedToolCall),
|
||||
}); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
resultText := formatPreDispatchToolError(invocation, cause)
|
||||
if err := service.appendToolResult(stream, invocation.CallID, strings.TrimSpace(invocation.ToolName), invocation.ArgsJSON, resultText, invocation.ReasoningContent, nil); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := service.publishToolCallCompleted(stream.RequestID, invocation.CallID, invocation.ModelCallID, nil); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := service.syncSummaryCarryForward(stream.ConversationID, stream.RequestID, invocation.ModelCallID); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := service.publishCheckpoint(stream.RequestID, stream.ConversationID); err != nil {
|
||||
return err
|
||||
}
|
||||
return service.reconcileStream(stream)
|
||||
}
|
||||
|
||||
func formatPreDispatchToolError(invocation runtimecore.ToolInvocation, cause error) string {
|
||||
toolName := strings.TrimSpace(invocation.ToolName)
|
||||
if toolName == "" {
|
||||
toolName = "Tool"
|
||||
}
|
||||
message := "unknown error"
|
||||
if cause != nil && strings.TrimSpace(cause.Error()) != "" {
|
||||
message = strings.TrimSpace(cause.Error())
|
||||
}
|
||||
return fmt.Sprintf("%s error: %s", toolName, message)
|
||||
}
|
||||
@@ -0,0 +1,394 @@
|
||||
package forwarder
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"strings"
|
||||
"unicode/utf8"
|
||||
)
|
||||
|
||||
const (
|
||||
projectedReplayKiB = 1024
|
||||
|
||||
projectedReadReplayLimit = 64 * projectedReplayKiB
|
||||
projectedShellReplayLimit = 128 * projectedReplayKiB
|
||||
projectedShellStreamLimit = 16 * projectedReplayKiB
|
||||
projectedShellInterleavedLimit = 32 * projectedReplayKiB
|
||||
projectedGrepReplayLimit = 32 * projectedReplayKiB
|
||||
projectedEditReplayLimit = 32 * projectedReplayKiB
|
||||
projectedPatchEditReplayLimit = 4 * projectedReplayKiB
|
||||
projectedWebFetchReplayLimit = 32 * projectedReplayKiB
|
||||
projectedWebSearchReplayLimit = 16 * projectedReplayKiB
|
||||
projectedMcpReplayLimit = 32 * projectedReplayKiB
|
||||
)
|
||||
|
||||
func limitProjectedToolResultReplay(toolName string, content string, resultText string, fromStoredToolCall bool, historical bool) string {
|
||||
if compacted, ok := compactProjectedGenerateImageResultReplay(toolName, content, resultText); ok {
|
||||
return compacted
|
||||
}
|
||||
if historical {
|
||||
if compacted, ok := compactHistoricalEditErrorReplay(toolName, content); ok {
|
||||
content = compacted
|
||||
fromStoredToolCall = false
|
||||
}
|
||||
}
|
||||
if compacted, ok := compactProjectedEditToolResultReplay(toolName, content); ok {
|
||||
content = compacted
|
||||
fromStoredToolCall = false
|
||||
}
|
||||
if compacted, ok := compactProjectedShellToolResultReplay(toolName, content); ok {
|
||||
content = compacted
|
||||
fromStoredToolCall = false
|
||||
}
|
||||
limit, ok := projectedToolReplayLimit(toolName)
|
||||
if !ok {
|
||||
return strings.TrimSpace(content)
|
||||
}
|
||||
content = strings.TrimSpace(content)
|
||||
if len(content) <= limit {
|
||||
return content
|
||||
}
|
||||
if fromStoredToolCall {
|
||||
fallback := strings.TrimSpace(resultText)
|
||||
notice := fmt.Sprintf("[tool result replay truncated: stored ToolCall result exceeded %d bytes]", limit)
|
||||
if fallback == "" {
|
||||
fallback = notice
|
||||
} else {
|
||||
fallback += "\n\n" + notice
|
||||
}
|
||||
return truncateProjectedReplayText(toolName, fallback, limit)
|
||||
}
|
||||
return truncateProjectedReplayText(toolName, content, limit)
|
||||
}
|
||||
|
||||
func projectedToolReplayLimit(toolName string) (int, bool) {
|
||||
switch strings.TrimSpace(toolName) {
|
||||
case "GenerateImage":
|
||||
return projectedWebSearchReplayLimit, true
|
||||
case "Read":
|
||||
return projectedReadReplayLimit, true
|
||||
case "Shell":
|
||||
return projectedShellReplayLimit, true
|
||||
case "Grep":
|
||||
return projectedGrepReplayLimit, true
|
||||
case "PatchEdit", "PatchEditLines", "PatchEditSpan":
|
||||
return projectedPatchEditReplayLimit, true
|
||||
case "Edit", "Write":
|
||||
return projectedEditReplayLimit, true
|
||||
case "WebFetch":
|
||||
return projectedWebFetchReplayLimit, true
|
||||
case "WebSearch":
|
||||
return projectedWebSearchReplayLimit, true
|
||||
case "CallMcpTool", "FetchMcpResource", "ListMcpResources":
|
||||
return projectedMcpReplayLimit, true
|
||||
default:
|
||||
return 0, false
|
||||
}
|
||||
}
|
||||
|
||||
func compactProjectedGenerateImageResultReplay(toolName string, content string, resultText string) (string, bool) {
|
||||
if strings.TrimSpace(toolName) != "GenerateImage" {
|
||||
return "", false
|
||||
}
|
||||
fallback := strings.TrimSpace(resultText)
|
||||
trimmed := strings.TrimSpace(content)
|
||||
if trimmed == "" || trimmed == "{}" || trimmed == "null" {
|
||||
if fallback == "" {
|
||||
return "", false
|
||||
}
|
||||
return fallback, true
|
||||
}
|
||||
var payload any
|
||||
if err := json.Unmarshal([]byte(trimmed), &payload); err != nil {
|
||||
if fallback == "" {
|
||||
return truncateProjectedReplayText("GenerateImage", trimmed, projectedWebSearchReplayLimit), true
|
||||
}
|
||||
return fallback, true
|
||||
}
|
||||
if compactGenerateImagePayload(payload) {
|
||||
encoded, err := json.Marshal(payload)
|
||||
if err == nil {
|
||||
return string(encoded), true
|
||||
}
|
||||
}
|
||||
if fallback != "" {
|
||||
return fallback, true
|
||||
}
|
||||
return trimmed, true
|
||||
}
|
||||
|
||||
func compactGenerateImagePayload(value any) bool {
|
||||
changed := false
|
||||
switch item := value.(type) {
|
||||
case map[string]any:
|
||||
for key, child := range item {
|
||||
if text, ok := child.(string); ok {
|
||||
switch key {
|
||||
case "image_data", "imageData":
|
||||
item[key] = fmt.Sprintf("[base64 image data omitted from replay; bytes=%d]", len(strings.TrimSpace(text)))
|
||||
changed = true
|
||||
continue
|
||||
}
|
||||
}
|
||||
if compactGenerateImagePayload(child) {
|
||||
changed = true
|
||||
}
|
||||
}
|
||||
case []any:
|
||||
for _, child := range item {
|
||||
if compactGenerateImagePayload(child) {
|
||||
changed = true
|
||||
}
|
||||
}
|
||||
}
|
||||
return changed
|
||||
}
|
||||
|
||||
func patchEditReplayLimit(toolName string) int {
|
||||
switch strings.TrimSpace(toolName) {
|
||||
case "PatchEdit", "PatchEditLines", "PatchEditSpan":
|
||||
return projectedPatchEditReplayLimit
|
||||
default:
|
||||
return projectedEditReplayLimit
|
||||
}
|
||||
}
|
||||
|
||||
func compactProjectedShellToolResultReplay(toolName string, content string) (string, bool) {
|
||||
if strings.TrimSpace(toolName) != "Shell" {
|
||||
return "", false
|
||||
}
|
||||
trimmed := strings.TrimSpace(content)
|
||||
if trimmed == "" || trimmed == "{}" || trimmed == "null" {
|
||||
return "", false
|
||||
}
|
||||
var payload any
|
||||
if err := json.Unmarshal([]byte(trimmed), &payload); err != nil {
|
||||
return "", false
|
||||
}
|
||||
if !compactProjectedShellFields(payload) {
|
||||
return "", false
|
||||
}
|
||||
encoded, err := json.Marshal(payload)
|
||||
if err != nil {
|
||||
return "", false
|
||||
}
|
||||
return string(encoded), true
|
||||
}
|
||||
|
||||
func compactProjectedShellFields(value any) bool {
|
||||
changed := false
|
||||
switch item := value.(type) {
|
||||
case map[string]any:
|
||||
for key, child := range item {
|
||||
if text, ok := child.(string); ok {
|
||||
switch key {
|
||||
case "stdout", "stderr":
|
||||
next := truncateProjectedReplayTextMiddle("Shell "+key, text, projectedShellStreamLimit)
|
||||
if next != text {
|
||||
item[key] = next
|
||||
changed = true
|
||||
}
|
||||
continue
|
||||
case "interleaved_output", "interleavedOutput":
|
||||
next := truncateProjectedReplayTextMiddle("Shell interleaved output", text, projectedShellInterleavedLimit)
|
||||
if next != text {
|
||||
item[key] = next
|
||||
changed = true
|
||||
}
|
||||
continue
|
||||
}
|
||||
}
|
||||
if compactProjectedShellFields(child) {
|
||||
changed = true
|
||||
}
|
||||
}
|
||||
case []any:
|
||||
for _, child := range item {
|
||||
if compactProjectedShellFields(child) {
|
||||
changed = true
|
||||
}
|
||||
}
|
||||
}
|
||||
return changed
|
||||
}
|
||||
|
||||
func compactProjectedEditToolResultReplay(toolName string, content string) (string, bool) {
|
||||
switch strings.TrimSpace(toolName) {
|
||||
case "PatchEdit", "PatchEditLines", "PatchEditSpan", "Edit", "Write":
|
||||
default:
|
||||
return "", false
|
||||
}
|
||||
trimmed := strings.TrimSpace(content)
|
||||
if trimmed == "" || trimmed == "{}" || trimmed == "null" {
|
||||
return "", false
|
||||
}
|
||||
var payload map[string]any
|
||||
if err := json.Unmarshal([]byte(trimmed), &payload); err != nil {
|
||||
return "", false
|
||||
}
|
||||
success, ok := payload["success"].(map[string]any)
|
||||
if !ok {
|
||||
return "", false
|
||||
}
|
||||
diffString := firstJSONText(success, "diff_string", "diffString")
|
||||
beforeContent := ""
|
||||
afterContent := ""
|
||||
if diffString == "" {
|
||||
beforeContent = firstJSONText(success, "before_full_file_content", "beforeFullFileContent")
|
||||
afterContent = firstJSONText(success, "after_full_file_content", "afterFullFileContent")
|
||||
if beforeContent != "" {
|
||||
diffString, _, _ = computeEditDiff(beforeContent, afterContent)
|
||||
}
|
||||
}
|
||||
if diffString != "" {
|
||||
diffString = truncateProjectedReplayText(firstNonEmpty(strings.TrimSpace(toolName), "PatchEdit"), diffString, patchEditReplayLimit(toolName))
|
||||
encoded, err := json.Marshal(map[string]any{
|
||||
"success": map[string]any{
|
||||
"diff_string": diffString,
|
||||
},
|
||||
})
|
||||
if err != nil {
|
||||
return "", false
|
||||
}
|
||||
return string(encoded), true
|
||||
}
|
||||
if afterContent != "" {
|
||||
afterContent = truncateProjectedReplayText(firstNonEmpty(strings.TrimSpace(toolName), "Write"), afterContent, patchEditReplayLimit(toolName))
|
||||
encoded, err := json.Marshal(map[string]any{
|
||||
"success": map[string]any{
|
||||
"after_full_file_content": afterContent,
|
||||
},
|
||||
})
|
||||
if err != nil {
|
||||
return "", false
|
||||
}
|
||||
return string(encoded), true
|
||||
}
|
||||
return "", false
|
||||
}
|
||||
|
||||
func compactHistoricalEditErrorReplay(toolName string, content string) (string, bool) {
|
||||
switch strings.TrimSpace(toolName) {
|
||||
case "PatchEdit", "PatchEditLines", "PatchEditSpan", "Edit", "Write":
|
||||
default:
|
||||
return "", false
|
||||
}
|
||||
trimmed := strings.TrimSpace(content)
|
||||
if trimmed == "" {
|
||||
return "", false
|
||||
}
|
||||
var payload map[string]any
|
||||
if err := json.Unmarshal([]byte(trimmed), &payload); err != nil {
|
||||
return "", false
|
||||
}
|
||||
errorPayload, ok := payload["error"].(map[string]any)
|
||||
if !ok {
|
||||
return "", false
|
||||
}
|
||||
if firstJSONText(errorPayload, "error", "modelVisibleError") == "" {
|
||||
return "", false
|
||||
}
|
||||
encoded, err := json.Marshal(map[string]any{
|
||||
"error": map[string]any{
|
||||
"error": "historical edit error omitted from replay; re-read the file before editing",
|
||||
},
|
||||
})
|
||||
if err != nil {
|
||||
return "", false
|
||||
}
|
||||
return string(encoded), true
|
||||
}
|
||||
|
||||
func firstJSONText(payload map[string]any, keys ...string) string {
|
||||
for _, key := range keys {
|
||||
if value, ok := payload[key].(string); ok {
|
||||
return value
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func truncateProjectedReplayText(toolName string, text string, limit int) string {
|
||||
if limit <= 0 || len(text) <= limit {
|
||||
return strings.TrimSpace(text)
|
||||
}
|
||||
original := len(text)
|
||||
notice := fmt.Sprintf("\n\n[truncated: %s result exceeded %d bytes; showing %d of %d bytes]", toolName, limit, limit, original)
|
||||
for {
|
||||
keep := limit - len(notice)
|
||||
if keep <= 0 {
|
||||
return truncateProjectedUTF8(text, limit)
|
||||
}
|
||||
kept := truncateProjectedUTF8(text, keep)
|
||||
nextNotice := fmt.Sprintf("\n\n[truncated: %s result exceeded %d bytes; showing %d of %d bytes]", toolName, limit, len(kept), original)
|
||||
output := strings.TrimRight(kept, "\n") + nextNotice
|
||||
if len(output) <= limit || nextNotice == notice {
|
||||
return output
|
||||
}
|
||||
notice = nextNotice
|
||||
}
|
||||
}
|
||||
|
||||
func truncateProjectedReplayTextMiddle(toolName string, text string, limit int) string {
|
||||
if limit <= 0 {
|
||||
return ""
|
||||
}
|
||||
if len(text) <= limit {
|
||||
return text
|
||||
}
|
||||
original := len(text)
|
||||
notice := fmt.Sprintf("\n\n[truncated: %s result exceeded %d bytes; omitted middle; showing %d of %d bytes]\n\n", toolName, limit, limit, original)
|
||||
for {
|
||||
keep := limit - len(notice)
|
||||
if keep <= 0 {
|
||||
return truncateProjectedUTF8(text, limit)
|
||||
}
|
||||
headLimit := keep / 2
|
||||
tailLimit := keep - headLimit
|
||||
head := truncateProjectedUTF8(text, headLimit)
|
||||
tail := truncateProjectedUTF8Suffix(text, tailLimit)
|
||||
kept := len(head) + len(tail)
|
||||
nextNotice := fmt.Sprintf("\n\n[truncated: %s result exceeded %d bytes; omitted middle; showing %d of %d bytes]\n\n", toolName, limit, kept, original)
|
||||
output := head + nextNotice + tail
|
||||
if len(output) <= limit || nextNotice == notice {
|
||||
return output
|
||||
}
|
||||
notice = nextNotice
|
||||
}
|
||||
}
|
||||
|
||||
func truncateProjectedUTF8(text string, limit int) string {
|
||||
if limit <= 0 {
|
||||
return ""
|
||||
}
|
||||
if len(text) <= limit {
|
||||
return text
|
||||
}
|
||||
if limit > len(text) {
|
||||
limit = len(text)
|
||||
}
|
||||
truncated := text[:limit]
|
||||
for !utf8.ValidString(truncated) && len(truncated) > 0 {
|
||||
truncated = truncated[:len(truncated)-1]
|
||||
}
|
||||
return truncated
|
||||
}
|
||||
|
||||
func truncateProjectedUTF8Suffix(text string, limit int) string {
|
||||
if limit <= 0 {
|
||||
return ""
|
||||
}
|
||||
if len(text) <= limit {
|
||||
return text
|
||||
}
|
||||
start := len(text) - limit
|
||||
if start < 0 {
|
||||
start = 0
|
||||
}
|
||||
suffix := text[start:]
|
||||
for !utf8.ValidString(suffix) && start < len(text) {
|
||||
start++
|
||||
suffix = text[start:]
|
||||
}
|
||||
return suffix
|
||||
}
|
||||
@@ -0,0 +1,510 @@
|
||||
// types.go 定义 forwarder 的核心数据结构与最小接口边界。
|
||||
package forwarder
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"strings"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
"cursor/gen/agentv1"
|
||||
runtimecore "cursor/internal/backend/agent/core"
|
||||
modeladapter "cursor/internal/backend/agent/model"
|
||||
)
|
||||
|
||||
type ConversationFile struct {
|
||||
SchemaVersion int `json:"schema_version,omitempty"`
|
||||
ConversationID string `json:"conversation_id"`
|
||||
RootConversationID string `json:"root_conversation_id"`
|
||||
ParentConversationID string `json:"parent_conversation_id"`
|
||||
ParentToolCallID string `json:"parent_tool_call_id"`
|
||||
SubagentTypeName string `json:"subagent_type_name,omitempty"`
|
||||
Mode string `json:"mode"`
|
||||
ContextVersion int64 `json:"context_version,omitempty"`
|
||||
CurrentLoopID string `json:"current_loop_id,omitempty"`
|
||||
CurrentLoopStatus string `json:"current_loop_status,omitempty"`
|
||||
CurrentRequestID string `json:"current_request_id,omitempty"`
|
||||
CurrentTurnSeq int64 `json:"current_turn_seq,omitempty"`
|
||||
TokenDetailsUsedTokens uint32 `json:"token_details_used_tokens,omitempty"`
|
||||
TokenDetailsMaxTokens uint32 `json:"token_details_max_tokens,omitempty"`
|
||||
AutoCompactionPending bool `json:"auto_compaction_pending,omitempty"`
|
||||
AutoCompactionPromptTokens int64 `json:"auto_compaction_prompt_tokens,omitempty"`
|
||||
AutoCompactionReserveTokens int64 `json:"auto_compaction_reserve_tokens,omitempty"`
|
||||
AutoCompactionTriggeredAt string `json:"auto_compaction_triggered_at,omitempty"`
|
||||
AutoCompactionSourceModelCallID string `json:"auto_compaction_source_model_call_id,omitempty"`
|
||||
CurrentPlanText string `json:"current_plan_text,omitempty"`
|
||||
CurrentPlans map[string]*agentv1.PlanRegistryEntry `json:"current_plans,omitempty"`
|
||||
CurrentTodos []*agentv1.TodoItem `json:"current_todos,omitempty"`
|
||||
LatestRequestPrefix *ConversationRequestPrefix `json:"latest_request_prefix,omitempty"`
|
||||
LastProviderCall *ConversationProviderCall `json:"last_provider_call,omitempty"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
NextTurnSeq int64 `json:"next_turn_seq"`
|
||||
NextEntrySeq int64 `json:"next_entry_seq"`
|
||||
Entries []HistoryEntry `json:"entries,omitempty"`
|
||||
}
|
||||
|
||||
type ConversationRequestPrefix struct {
|
||||
RequestID string `json:"request_id,omitempty"`
|
||||
ModelCallID string `json:"model_call_id,omitempty"`
|
||||
Provider string `json:"provider,omitempty"`
|
||||
OpenAIEndpoint string `json:"openai_endpoint,omitempty"`
|
||||
Model string `json:"model,omitempty"`
|
||||
PromptTokensTotal int64 `json:"prompt_tokens_total,omitempty"`
|
||||
ReplayMessageCount int `json:"replay_message_count,omitempty"`
|
||||
CanonicalBodyHash string `json:"canonical_body_hash,omitempty"`
|
||||
FrontierHash string `json:"frontier_hash,omitempty"`
|
||||
FrontierPath string `json:"frontier_path,omitempty"`
|
||||
BreakpointCount int `json:"breakpoint_count,omitempty"`
|
||||
ExpectedCacheRead bool `json:"expected_cache_read,omitempty"`
|
||||
PreviousFrontierMatched bool `json:"previous_frontier_matched,omitempty"`
|
||||
UpdatedAt time.Time `json:"updated_at,omitempty"`
|
||||
}
|
||||
|
||||
type ConversationProviderCall struct {
|
||||
RequestID string `json:"request_id,omitempty"`
|
||||
ModelCallID string `json:"model_call_id,omitempty"`
|
||||
Provider string `json:"provider,omitempty"`
|
||||
Model string `json:"model,omitempty"`
|
||||
Status string `json:"status,omitempty"`
|
||||
ErrorText string `json:"error_text,omitempty"`
|
||||
UpdatedAt time.Time `json:"updated_at,omitempty"`
|
||||
}
|
||||
|
||||
type HistoryEntry struct {
|
||||
Seq int64 `json:"seq"`
|
||||
TurnSeq int64 `json:"turn_seq"`
|
||||
RequestID string `json:"request_id,omitempty"`
|
||||
Role string `json:"role"`
|
||||
Kind string `json:"kind"`
|
||||
ToolCallID string `json:"tool_call_id,omitempty"`
|
||||
ParentToolCallID string `json:"parent_tool_call_id,omitempty"`
|
||||
Payload json.RawMessage `json:"payload"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
}
|
||||
|
||||
type ConversationSummary struct {
|
||||
ConversationID string `json:"conversation_id"`
|
||||
Mode string `json:"mode"`
|
||||
EntriesCount int `json:"entries_count"`
|
||||
NextTurnSeq int64 `json:"next_turn_seq"`
|
||||
NextEntrySeq int64 `json:"next_entry_seq"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
}
|
||||
|
||||
type StreamStatus string
|
||||
|
||||
const (
|
||||
StreamStatusCreated StreamStatus = "created"
|
||||
StreamStatusStreaming StreamStatus = "streaming"
|
||||
StreamStatusCompleted StreamStatus = "completed"
|
||||
StreamStatusCanceled StreamStatus = "canceled"
|
||||
StreamStatusFailed StreamStatus = "failed"
|
||||
)
|
||||
|
||||
type StreamEvent struct {
|
||||
Message *agentv1.AgentServerMessage
|
||||
End bool
|
||||
TerminalErrorCode string
|
||||
TerminalErrorMessage string
|
||||
}
|
||||
|
||||
type StreamSubscriber struct {
|
||||
Signal chan struct{}
|
||||
}
|
||||
|
||||
type ActiveStream struct {
|
||||
mu sync.Mutex
|
||||
|
||||
RequestID string
|
||||
ConversationID string
|
||||
TurnSeq int64
|
||||
ModelID string
|
||||
ModelName string
|
||||
Mode agentv1.AgentMode
|
||||
LatestUserText string
|
||||
Status StreamStatus
|
||||
ThinkingEffort string
|
||||
SubagentModelOverrides map[string]runtimecore.SubagentModelOverrideSelection
|
||||
|
||||
CurrentModelCallID string
|
||||
ProviderActive bool
|
||||
ProviderCancel func()
|
||||
ProviderPassCount int
|
||||
ToolInvocationCount int
|
||||
ActorMailbox chan streamCommandEnvelope
|
||||
ActorDone chan struct{}
|
||||
Phase TurnPhase
|
||||
PendingProviderAction providerAction
|
||||
PendingProviderCompletion *pendingTurnCompletion
|
||||
CurrentProviderToken uint64
|
||||
CurrentCompactionToken uint64
|
||||
TimerTokens map[string]uint64
|
||||
ProviderAccumulatedText string
|
||||
ProviderAccumulatedReasoning string
|
||||
ProviderAccumulatedReasoningSignature string
|
||||
ProviderAccumulatedReasoningSignatureSource string
|
||||
ProviderAccumulatedReasoningItemID string
|
||||
ProviderAccumulatedReasoningStatus string
|
||||
ProviderAccumulatedReasoningSummary json.RawMessage
|
||||
ProviderSyntheticThinkingStartedAt time.Time
|
||||
ProviderSyntheticThinkingPublished bool
|
||||
ProviderFinishReason string
|
||||
ProviderUsage turnUsageSnapshot
|
||||
ProviderTerminalToolInvocation bool
|
||||
PendingCompaction *PendingCompaction
|
||||
|
||||
Backlog []StreamEvent
|
||||
Subscribers map[string]*StreamSubscriber
|
||||
CheckpointConversation *ConversationFile
|
||||
PendingExecs map[string]runtimecore.PendingExec
|
||||
PendingInteractions map[string]runtimecore.PendingInteraction
|
||||
PartialToolCallIDs map[string]struct{}
|
||||
PatchEditQueues map[string][]queuedPatchEditOperation
|
||||
MCPToolServers map[string]string
|
||||
WorkspacePaths []string
|
||||
TerminalsFolder string
|
||||
RequestFileContents map[string]string
|
||||
RecentCompletedExecs map[uint32]time.Time
|
||||
BackgroundShells map[string]*BackgroundShellState
|
||||
BackgroundShellsByMessageID map[uint32]string
|
||||
BackgroundShellsByExecID map[string]string
|
||||
BackgroundShellActions map[string]time.Time
|
||||
TerminalCleanupTimer *time.Timer
|
||||
TerminalCleanupSeq atomic.Uint64
|
||||
|
||||
CreatedAt time.Time
|
||||
UpdatedAt time.Time
|
||||
}
|
||||
|
||||
type BackgroundShellState struct {
|
||||
ShellID string
|
||||
Command string
|
||||
WorkingDirectory string
|
||||
PID *uint32
|
||||
OriginalToolCallID string
|
||||
OriginalExecID string
|
||||
OriginalMessageID uint32
|
||||
ModelCallID string
|
||||
ArgsJSON []byte
|
||||
Status string
|
||||
ExitCode *int32
|
||||
StdoutBuffer string
|
||||
StderrBuffer string
|
||||
AwaitStdoutOffset int
|
||||
AwaitStderrOffset int
|
||||
CreatedAt time.Time
|
||||
LastActivityAt time.Time
|
||||
CompletedAt time.Time
|
||||
StreamClosed bool
|
||||
}
|
||||
|
||||
type pendingTurnCompletion struct {
|
||||
ConversationID string
|
||||
RequestID string
|
||||
TurnSeq int64
|
||||
ModelCallID string
|
||||
ProviderPass int
|
||||
Usage turnUsageSnapshot
|
||||
Disposition pendingCompletionDisposition
|
||||
}
|
||||
|
||||
type PendingCompaction struct {
|
||||
Trigger string
|
||||
ContextTokens int64
|
||||
ContextWindowSize int64
|
||||
ContextUsagePercent float64
|
||||
ReserveTokens int64
|
||||
MessageCount int32
|
||||
MessagesToCompact int32
|
||||
CompactTurnCount int32
|
||||
IsFirstCompaction bool
|
||||
ExistingSummary string
|
||||
CompactedTurns []compactedTurnSummary
|
||||
ManualInstruction string
|
||||
RequestSource string
|
||||
CurrentTurnSeq int64
|
||||
CurrentRequestID string
|
||||
CurrentUserText string
|
||||
PreserveCurrentTurnInputs bool
|
||||
HookMessage string
|
||||
SummaryModelCallID string
|
||||
StartedAt time.Time
|
||||
}
|
||||
|
||||
type ProviderRequest struct {
|
||||
RequestID string
|
||||
ConversationID string
|
||||
RunID string
|
||||
ModelCallID string
|
||||
ModelID string
|
||||
Mode agentv1.AgentMode
|
||||
ThinkingEffort string
|
||||
Messages []modeladapter.Message
|
||||
StableMessageCount int
|
||||
Tools []json.RawMessage
|
||||
MaxTokens int
|
||||
RequestKnobs map[string]any
|
||||
CompileSummary string
|
||||
Observer modeladapter.LLMArtifactObserver
|
||||
ArtifactPaths *modeladapter.LLMArtifactPaths
|
||||
RequestBodyOverride map[string]any
|
||||
}
|
||||
|
||||
type ProviderGateway interface {
|
||||
StartStream(ctx context.Context, req ProviderRequest, sink func(modeladapter.ModelEvent) error) error
|
||||
}
|
||||
|
||||
type ToolCatalog interface {
|
||||
Load(mode agentv1.AgentMode, subagentTypeName string) ([]json.RawMessage, []string, error)
|
||||
}
|
||||
|
||||
type PromptReminders struct {
|
||||
SystemParts []string
|
||||
TailMessages []modeladapter.Message
|
||||
PromptContexts []PromptContextMessage
|
||||
}
|
||||
|
||||
type ReminderInjector interface {
|
||||
Inject(mode agentv1.AgentMode, conversation *ConversationFile, replayMessages []modeladapter.Message, latestUserText string, toolNames []string) PromptReminders
|
||||
}
|
||||
|
||||
type PromptContextMessage struct {
|
||||
Source string
|
||||
Message modeladapter.Message
|
||||
ContentHash string
|
||||
Persist bool
|
||||
}
|
||||
|
||||
type CompiledConversation struct {
|
||||
Mode agentv1.AgentMode
|
||||
Messages []modeladapter.Message
|
||||
StableMessageCount int
|
||||
Tools []json.RawMessage
|
||||
CompileSummary string
|
||||
}
|
||||
|
||||
type ToolRequestKind string
|
||||
|
||||
const (
|
||||
ToolRequestExec ToolRequestKind = "exec"
|
||||
ToolRequestInteraction ToolRequestKind = "interaction"
|
||||
)
|
||||
|
||||
// providerTerminalError 表示底层 LLM/provider 返回的真实错误。
|
||||
type providerTerminalError struct {
|
||||
cause error
|
||||
}
|
||||
|
||||
// Error 返回 provider 错误的字符串形式。
|
||||
func (err providerTerminalError) Error() string {
|
||||
if err.cause == nil {
|
||||
return "provider error"
|
||||
}
|
||||
return err.cause.Error()
|
||||
}
|
||||
|
||||
// Unwrap 允许调用方继续取到底层原始错误。
|
||||
func (err providerTerminalError) Unwrap() error {
|
||||
return err.cause
|
||||
}
|
||||
|
||||
type toolResultEntryPayload struct {
|
||||
ToolCallID string `json:"tool_call_id"`
|
||||
ToolName string `json:"tool_name"`
|
||||
Arguments string `json:"arguments,omitempty"`
|
||||
ResultText string `json:"result_text,omitempty"`
|
||||
ReasoningContent string `json:"reasoning_content,omitempty"`
|
||||
ReasoningSignature string `json:"reasoning_signature,omitempty"`
|
||||
ReasoningSignatureSource string `json:"reasoning_signature_source,omitempty"`
|
||||
ReasoningItemID string `json:"reasoning_item_id,omitempty"`
|
||||
ReasoningStatus string `json:"reasoning_status,omitempty"`
|
||||
ReasoningSummary json.RawMessage `json:"reasoning_summary,omitempty"`
|
||||
ProviderItemID string `json:"provider_item_id,omitempty"`
|
||||
ProviderCallID string `json:"provider_call_id,omitempty"`
|
||||
ProviderStatus string `json:"provider_status,omitempty"`
|
||||
ToolCall json.RawMessage `json:"tool_call,omitempty"`
|
||||
}
|
||||
|
||||
type toolCallEntryPayload struct {
|
||||
ToolCallID string `json:"tool_call_id"`
|
||||
ToolName string `json:"tool_name"`
|
||||
ReasoningContent string `json:"reasoning_content,omitempty"`
|
||||
ReasoningSignature string `json:"reasoning_signature,omitempty"`
|
||||
ReasoningSignatureSource string `json:"reasoning_signature_source,omitempty"`
|
||||
ReasoningItemID string `json:"reasoning_item_id,omitempty"`
|
||||
ReasoningStatus string `json:"reasoning_status,omitempty"`
|
||||
ReasoningSummary json.RawMessage `json:"reasoning_summary,omitempty"`
|
||||
ProviderItemID string `json:"provider_item_id,omitempty"`
|
||||
ProviderCallID string `json:"provider_call_id,omitempty"`
|
||||
ProviderStatus string `json:"provider_status,omitempty"`
|
||||
ToolCall json.RawMessage `json:"tool_call"`
|
||||
}
|
||||
|
||||
type assistantTextPayload struct {
|
||||
Text string `json:"text"`
|
||||
ReasoningContent string `json:"reasoning_content,omitempty"`
|
||||
ReasoningSignature string `json:"reasoning_signature,omitempty"`
|
||||
ReasoningSignatureSource string `json:"reasoning_signature_source,omitempty"`
|
||||
ReasoningItemID string `json:"reasoning_item_id,omitempty"`
|
||||
ReasoningStatus string `json:"reasoning_status,omitempty"`
|
||||
ReasoningSummary json.RawMessage `json:"reasoning_summary,omitempty"`
|
||||
}
|
||||
|
||||
type metadataPayload struct {
|
||||
Type string `json:"type"`
|
||||
Value map[string]any `json:"value,omitempty"`
|
||||
}
|
||||
|
||||
type promptContextEntryPayload struct {
|
||||
Source string `json:"source"`
|
||||
Role string `json:"role"`
|
||||
Content string `json:"content"`
|
||||
ContentHash string `json:"content_hash,omitempty"`
|
||||
}
|
||||
|
||||
type modelMessageEntryPayload struct {
|
||||
Message modeladapter.Message `json:"message"`
|
||||
}
|
||||
|
||||
type runtimeStateEntryPayload struct {
|
||||
PlanText string `json:"plan_text,omitempty"`
|
||||
Plans map[string]*agentv1.PlanRegistryEntry `json:"plans,omitempty"`
|
||||
Todos []*agentv1.TodoItem `json:"todos,omitempty"`
|
||||
}
|
||||
|
||||
type compactionSummaryEntryPayload struct {
|
||||
Summary string `json:"summary"`
|
||||
Trigger string `json:"trigger,omitempty"`
|
||||
CurrentTurnSeq int64 `json:"current_turn_seq,omitempty"`
|
||||
CurrentRequestID string `json:"current_request_id,omitempty"`
|
||||
CompactTurnCount int32 `json:"compact_turn_count,omitempty"`
|
||||
MessagesToCompact int32 `json:"messages_to_compact,omitempty"`
|
||||
PreserveCurrentTurnInputs bool `json:"preserve_current_turn_inputs,omitempty"`
|
||||
}
|
||||
|
||||
type ModeSource string
|
||||
|
||||
const (
|
||||
ModeSourceUnknown ModeSource = ""
|
||||
ModeSourceUserMessage ModeSource = "user_message"
|
||||
ModeSourceExecutePlanAction ModeSource = "execute_plan_action"
|
||||
ModeSourceConversationState ModeSource = "conversation_state"
|
||||
ModeSourceSwitchModeTool ModeSource = "switch_mode_tool"
|
||||
)
|
||||
|
||||
type InboundIntent struct {
|
||||
Kind string
|
||||
RequestID string
|
||||
ConversationID string
|
||||
ModelID string
|
||||
ModelName string
|
||||
ThinkingEffort string
|
||||
Mode agentv1.AgentMode
|
||||
HasExplicitMode bool
|
||||
ModeSource ModeSource
|
||||
StartsRun bool
|
||||
SubagentTypeName string
|
||||
SubagentModelOverrides map[string]runtimecore.SubagentModelOverrideSelection
|
||||
ConversationState *agentv1.ConversationStateStructure
|
||||
UserMessage *agentv1.UserMessage
|
||||
RequestContext *agentv1.RequestContext
|
||||
ClientMessage *agentv1.AgentClientMessage
|
||||
ExecClientMessage *agentv1.ExecClientMessage
|
||||
ExecClientControlMessage *agentv1.ExecClientControlMessage
|
||||
InteractionResponse *agentv1.InteractionResponse
|
||||
KVClientMessage *agentv1.KvClientMessage
|
||||
CancelReason string
|
||||
IgnoredReason string
|
||||
Prewarm bool
|
||||
}
|
||||
|
||||
// normalizeMode 对外部传入的 mode 做最小归一化,但不再静默降级。
|
||||
func normalizeMode(mode agentv1.AgentMode) agentv1.AgentMode {
|
||||
return mode
|
||||
}
|
||||
|
||||
func isSupportedActiveMode(mode agentv1.AgentMode) bool {
|
||||
switch normalizeMode(mode) {
|
||||
case agentv1.AgentMode_AGENT_MODE_AGENT,
|
||||
agentv1.AgentMode_AGENT_MODE_ASK,
|
||||
agentv1.AgentMode_AGENT_MODE_PLAN,
|
||||
agentv1.AgentMode_AGENT_MODE_DEBUG,
|
||||
agentv1.AgentMode_AGENT_MODE_MULTITASK:
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func validateSupportedActiveMode(mode agentv1.AgentMode) (agentv1.AgentMode, error) {
|
||||
normalized := normalizeMode(mode)
|
||||
if !isSupportedActiveMode(normalized) {
|
||||
return agentv1.AgentMode_AGENT_MODE_UNSPECIFIED, fmt.Errorf("unsupported active mode: %s", normalized.String())
|
||||
}
|
||||
return normalized, nil
|
||||
}
|
||||
|
||||
func resolveExplicitMode(mode agentv1.AgentMode, source ModeSource) (agentv1.AgentMode, ModeSource, bool, error) {
|
||||
normalized, err := validateSupportedActiveMode(mode)
|
||||
if err != nil {
|
||||
return agentv1.AgentMode_AGENT_MODE_UNSPECIFIED, source, true, fmt.Errorf("%w (source=%s)", err, strings.TrimSpace(string(source)))
|
||||
}
|
||||
return normalized, source, true, nil
|
||||
}
|
||||
|
||||
// modeAlias 把协议枚举转换为写入 JSON history 的简短模式名。
|
||||
func modeAlias(mode agentv1.AgentMode) (string, error) {
|
||||
switch normalizeMode(mode) {
|
||||
case agentv1.AgentMode_AGENT_MODE_AGENT:
|
||||
return "agent", nil
|
||||
case agentv1.AgentMode_AGENT_MODE_ASK:
|
||||
return "ask", nil
|
||||
case agentv1.AgentMode_AGENT_MODE_PLAN:
|
||||
return "plan", nil
|
||||
case agentv1.AgentMode_AGENT_MODE_DEBUG:
|
||||
return "debug", nil
|
||||
case agentv1.AgentMode_AGENT_MODE_MULTITASK:
|
||||
return "multitask", nil
|
||||
default:
|
||||
return "", fmt.Errorf("unsupported mode alias: %s", normalizeMode(mode).String())
|
||||
}
|
||||
}
|
||||
|
||||
// parseModeAlias 把写入 JSON history 的模式名恢复为协议枚举。
|
||||
func parseModeAlias(raw string) (agentv1.AgentMode, error) {
|
||||
switch strings.ToLower(strings.TrimSpace(raw)) {
|
||||
case "agent":
|
||||
return agentv1.AgentMode_AGENT_MODE_AGENT, nil
|
||||
case "ask":
|
||||
return agentv1.AgentMode_AGENT_MODE_ASK, nil
|
||||
case "plan":
|
||||
return agentv1.AgentMode_AGENT_MODE_PLAN, nil
|
||||
case "debug":
|
||||
return agentv1.AgentMode_AGENT_MODE_DEBUG, nil
|
||||
case "multitask":
|
||||
return agentv1.AgentMode_AGENT_MODE_MULTITASK, nil
|
||||
default:
|
||||
return agentv1.AgentMode_AGENT_MODE_UNSPECIFIED, fmt.Errorf("unsupported mode alias: %q", strings.TrimSpace(raw))
|
||||
}
|
||||
}
|
||||
|
||||
func parseTargetModeID(raw string) (agentv1.AgentMode, error) {
|
||||
switch strings.ToLower(strings.TrimSpace(raw)) {
|
||||
case "agent":
|
||||
return agentv1.AgentMode_AGENT_MODE_AGENT, nil
|
||||
case "ask":
|
||||
return agentv1.AgentMode_AGENT_MODE_ASK, nil
|
||||
case "plan":
|
||||
return agentv1.AgentMode_AGENT_MODE_PLAN, nil
|
||||
case "debug":
|
||||
return agentv1.AgentMode_AGENT_MODE_DEBUG, nil
|
||||
case "multitask":
|
||||
return agentv1.AgentMode_AGENT_MODE_MULTITASK, nil
|
||||
default:
|
||||
return agentv1.AgentMode_AGENT_MODE_UNSPECIFIED, fmt.Errorf("unsupported target mode id: %q", strings.TrimSpace(raw))
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,233 @@
|
||||
package forwarder
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"connectrpc.com/connect"
|
||||
|
||||
"cursor/gen/aiserverv1"
|
||||
"cursor/internal/logger"
|
||||
)
|
||||
|
||||
const (
|
||||
UploadServiceUploadDocumentationProcedure = "/aiserver.v1.UploadService/UploadDocumentation"
|
||||
UploadServiceGetDocProcedure = "/aiserver.v1.UploadService/GetDoc"
|
||||
UploadServiceGetPagesProcedure = "/aiserver.v1.UploadService/GetPages"
|
||||
UploadServiceUploadedStatusProcedure = "/aiserver.v1.UploadService/UploadedStatus"
|
||||
)
|
||||
|
||||
func newUploadServiceHandler(service *Service) http.Handler {
|
||||
mux := http.NewServeMux()
|
||||
mux.Handle(
|
||||
UploadServiceUploadDocumentationProcedure,
|
||||
connect.NewUnaryHandler(UploadServiceUploadDocumentationProcedure, service.UploadDocumentation),
|
||||
)
|
||||
mux.Handle(
|
||||
UploadServiceGetDocProcedure,
|
||||
connect.NewUnaryHandler(UploadServiceGetDocProcedure, service.GetDoc),
|
||||
)
|
||||
mux.Handle(
|
||||
UploadServiceGetPagesProcedure,
|
||||
connect.NewUnaryHandler(UploadServiceGetPagesProcedure, service.GetPages),
|
||||
)
|
||||
mux.Handle(
|
||||
UploadServiceUploadedStatusProcedure,
|
||||
connect.NewUnaryHandler(UploadServiceUploadedStatusProcedure, service.UploadedStatus),
|
||||
)
|
||||
mux.Handle("/", http.NotFoundHandler())
|
||||
return mux
|
||||
}
|
||||
|
||||
func (service *Service) UploadDocumentation(_ context.Context, req *connect.Request[aiserverv1.UploadDocumentationRequest]) (*connect.Response[aiserverv1.UploadResponse], error) {
|
||||
store, err := service.requireDocsIndexStore()
|
||||
if err != nil {
|
||||
return nil, connect.NewError(connect.CodeInternal, err)
|
||||
}
|
||||
identifier := strings.TrimSpace(req.Msg.GetDocIdentifier())
|
||||
if identifier == "" {
|
||||
identifier = stableDocsIdentifier("upload-documentation", time.Now().UTC().Format(time.RFC3339Nano))
|
||||
}
|
||||
record, err := store.Upsert(uploadedDocsIndexRecord(identifier))
|
||||
if err != nil {
|
||||
return nil, connect.NewError(connect.CodeInternal, err)
|
||||
}
|
||||
logger.Infof("UploadService UploadDocumentation completed doc_identifier=%s", record.Identifier)
|
||||
return connect.NewResponse(&aiserverv1.UploadResponse{
|
||||
Status: aiserverv1.UploadResponse_STATUS_SUCCESS,
|
||||
Progress: 1,
|
||||
UploadedPages: uploadedDocPages(record),
|
||||
DocUuid: uploadDocUUID(record),
|
||||
}), nil
|
||||
}
|
||||
|
||||
func (service *Service) GetDoc(_ context.Context, req *connect.Request[aiserverv1.GetDocRequest]) (*connect.Response[aiserverv1.ProtoDoc], error) {
|
||||
store, err := service.requireDocsIndexStore()
|
||||
if err != nil {
|
||||
return nil, connect.NewError(connect.CodeInternal, err)
|
||||
}
|
||||
identifier := strings.TrimSpace(req.Msg.GetDocIdentifier())
|
||||
if identifier == "" {
|
||||
return connect.NewResponse(uploadedDocNotFound(identifier)), nil
|
||||
}
|
||||
record, ok, err := store.Get(identifier)
|
||||
if err != nil {
|
||||
return nil, connect.NewError(connect.CodeInternal, err)
|
||||
}
|
||||
if !ok {
|
||||
logger.Infof("UploadService GetDoc not_found doc_identifier=%s", identifier)
|
||||
return connect.NewResponse(uploadedDocNotFound(identifier)), nil
|
||||
}
|
||||
logger.Infof("UploadService GetDoc completed doc_identifier=%s", record.Identifier)
|
||||
return connect.NewResponse(uploadedDocsProtoDoc(record)), nil
|
||||
}
|
||||
|
||||
func (service *Service) GetPages(_ context.Context, req *connect.Request[aiserverv1.GetPagesRequest]) (*connect.Response[aiserverv1.Pages], error) {
|
||||
store, err := service.requireDocsIndexStore()
|
||||
if err != nil {
|
||||
return nil, connect.NewError(connect.CodeInternal, err)
|
||||
}
|
||||
identifier := strings.TrimSpace(req.Msg.GetDocIdentifier())
|
||||
if identifier == "" {
|
||||
return connect.NewResponse(&aiserverv1.Pages{}), nil
|
||||
}
|
||||
record, ok, err := store.Get(identifier)
|
||||
if err != nil {
|
||||
return nil, connect.NewError(connect.CodeInternal, err)
|
||||
}
|
||||
if !ok {
|
||||
logger.Infof("UploadService GetPages empty doc_identifier=%s", identifier)
|
||||
return connect.NewResponse(&aiserverv1.Pages{}), nil
|
||||
}
|
||||
pages, pageURLs := uploadedDocsPagesResponse(record)
|
||||
logger.Infof("UploadService GetPages completed doc_identifier=%s pages=%d", record.Identifier, len(pages))
|
||||
return connect.NewResponse(&aiserverv1.Pages{Pages: pages, PageUrls: pageURLs}), nil
|
||||
}
|
||||
|
||||
func (service *Service) UploadedStatus(_ context.Context, req *connect.Request[aiserverv1.UploadedStatusRequest]) (*connect.Response[aiserverv1.UploadedStatus], error) {
|
||||
store, err := service.requireDocsIndexStore()
|
||||
if err != nil {
|
||||
return nil, connect.NewError(connect.CodeInternal, err)
|
||||
}
|
||||
identifier := strings.TrimSpace(req.Msg.GetDocIdentifier())
|
||||
if identifier == "" {
|
||||
return connect.NewResponse(&aiserverv1.UploadedStatus{Status: aiserverv1.UploadedStatus_STATUS_NOT_FOUND}), nil
|
||||
}
|
||||
record, ok, err := store.Get(identifier)
|
||||
if err != nil {
|
||||
return nil, connect.NewError(connect.CodeInternal, err)
|
||||
}
|
||||
if !ok {
|
||||
return connect.NewResponse(&aiserverv1.UploadedStatus{Status: aiserverv1.UploadedStatus_STATUS_NOT_FOUND}), nil
|
||||
}
|
||||
return connect.NewResponse(&aiserverv1.UploadedStatus{
|
||||
Status: aiserverv1.UploadedStatus_STATUS_SUCCEEDED,
|
||||
UploadedPages: uploadedDocPages(record),
|
||||
}), nil
|
||||
}
|
||||
|
||||
func uploadedDocsIndexRecord(identifier string) DocsIndexRecord {
|
||||
url := docsIndexURLCandidate(identifier)
|
||||
title := identifier
|
||||
if url != "" {
|
||||
title = docsTitleFromURL(url)
|
||||
}
|
||||
return DocsIndexRecord{
|
||||
ID: identifier,
|
||||
Identifier: identifier,
|
||||
Title: title,
|
||||
URL: url,
|
||||
Status: docsIndexStatusIndexed,
|
||||
Source: docsIndexSourceLocal,
|
||||
}
|
||||
}
|
||||
|
||||
func uploadedDocsProtoDoc(record DocsIndexRecord) *aiserverv1.ProtoDoc {
|
||||
createdAt := docsIndexTimeString(record.CreatedAt)
|
||||
updatedAt := docsIndexTimeString(record.UpdatedAt)
|
||||
lastUploadedAt := docsIndexTimeString(record.LastIndexAt)
|
||||
if lastUploadedAt == "" {
|
||||
lastUploadedAt = updatedAt
|
||||
}
|
||||
return &aiserverv1.ProtoDoc{
|
||||
Uuid: uploadDocUUID(record),
|
||||
DocIdentifier: record.Identifier,
|
||||
DocName: record.Title,
|
||||
DocUrlRoot: record.URL,
|
||||
DocUrlPrefix: record.URL,
|
||||
CreatedAt: createdAt,
|
||||
UpdatedAt: updatedAt,
|
||||
LastUploadedAt: lastUploadedAt,
|
||||
UploadStatus: &aiserverv1.UploadedStatus{
|
||||
Status: aiserverv1.UploadedStatus_STATUS_SUCCEEDED,
|
||||
UploadedPages: uploadedDocPages(record),
|
||||
},
|
||||
Pages: uploadedDocProtoPages(record),
|
||||
}
|
||||
}
|
||||
|
||||
func uploadedDocNotFound(identifier string) *aiserverv1.ProtoDoc {
|
||||
return &aiserverv1.ProtoDoc{
|
||||
DocIdentifier: strings.TrimSpace(identifier),
|
||||
UploadStatus: &aiserverv1.UploadedStatus{Status: aiserverv1.UploadedStatus_STATUS_NOT_FOUND},
|
||||
}
|
||||
}
|
||||
|
||||
func uploadedDocProtoPages(record DocsIndexRecord) []*aiserverv1.ProtoDocPage {
|
||||
pages, pageURLs := uploadedDocsPagesResponse(record)
|
||||
items := make([]*aiserverv1.ProtoDocPage, 0, len(pages))
|
||||
for i, page := range pages {
|
||||
url := ""
|
||||
if i < len(pageURLs) {
|
||||
url = pageURLs[i]
|
||||
}
|
||||
items = append(items, &aiserverv1.ProtoDocPage{Title: page, Url: url})
|
||||
}
|
||||
return items
|
||||
}
|
||||
|
||||
func uploadedDocsPagesResponse(record DocsIndexRecord) ([]string, []string) {
|
||||
title := strings.TrimSpace(record.Title)
|
||||
url := strings.TrimSpace(record.URL)
|
||||
if title == "" || url == "" {
|
||||
return nil, nil
|
||||
}
|
||||
return []string{title}, []string{url}
|
||||
}
|
||||
|
||||
func uploadedDocPages(record DocsIndexRecord) []string {
|
||||
pages, _ := uploadedDocsPagesResponse(record)
|
||||
return pages
|
||||
}
|
||||
|
||||
func uploadDocUUID(record DocsIndexRecord) string {
|
||||
identifier := strings.TrimSpace(record.Identifier)
|
||||
if identifier == "" {
|
||||
identifier = strings.TrimSpace(record.ID)
|
||||
}
|
||||
if identifier == "" {
|
||||
identifier = fmt.Sprintf("doc-%d", time.Now().UnixNano())
|
||||
}
|
||||
return stableDocsIdentifier("upload-doc-uuid", identifier)
|
||||
}
|
||||
|
||||
func docsIndexTimeString(value time.Time) string {
|
||||
if value.IsZero() {
|
||||
return ""
|
||||
}
|
||||
return value.UTC().Format(time.RFC3339)
|
||||
}
|
||||
|
||||
func docsTitleFromURL(value string) string {
|
||||
trimmed := strings.TrimSpace(value)
|
||||
trimmed = strings.TrimPrefix(trimmed, "https://")
|
||||
trimmed = strings.TrimPrefix(trimmed, "http://")
|
||||
trimmed = strings.Trim(trimmed, "/")
|
||||
if trimmed == "" {
|
||||
return value
|
||||
}
|
||||
return trimmed
|
||||
}
|
||||
@@ -0,0 +1,343 @@
|
||||
package forwarder
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
const (
|
||||
usageFileName = "usage.json"
|
||||
usageFileSchemaVersion = 2
|
||||
usageRecentEventLimit = 500
|
||||
|
||||
usageEventKindProvider = "provider_call"
|
||||
usageEventKindTurn = "turn_finalized"
|
||||
usageTurnStatusDone = "completed"
|
||||
)
|
||||
|
||||
type UsageFileStore struct {
|
||||
path string
|
||||
}
|
||||
|
||||
type usageFileDocument struct {
|
||||
SchemaVersion int `json:"schema_version"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
Totals usageFileTotals `json:"totals"`
|
||||
Daily []usageFileDaily `json:"daily"`
|
||||
RecentEvents []usageFileEvent `json:"recent_events"`
|
||||
EventIndex map[string]usageFileEvent `json:"event_index,omitempty"`
|
||||
}
|
||||
|
||||
type usageFileTotals struct {
|
||||
ProviderCalls int64 `json:"provider_calls"`
|
||||
TurnsTotal int64 `json:"turns_total"`
|
||||
ValidTurnsTotal int64 `json:"valid_turns_total"`
|
||||
InvalidTurnsTotal int64 `json:"invalid_turns_total"`
|
||||
InputTokens int64 `json:"input_tokens"`
|
||||
OutputTokens int64 `json:"output_tokens"`
|
||||
CacheReadTokens int64 `json:"cache_read_tokens"`
|
||||
CacheWriteTokens int64 `json:"cache_write_tokens"`
|
||||
TotalTokens int64 `json:"total_tokens"`
|
||||
}
|
||||
|
||||
type usageFileDaily struct {
|
||||
Date string `json:"date"`
|
||||
ProviderCalls int64 `json:"provider_calls"`
|
||||
TurnsTotal int64 `json:"turns_total"`
|
||||
ValidTurnsTotal int64 `json:"valid_turns_total"`
|
||||
InvalidTurnsTotal int64 `json:"invalid_turns_total"`
|
||||
InputTokens int64 `json:"input_tokens"`
|
||||
OutputTokens int64 `json:"output_tokens"`
|
||||
CacheReadTokens int64 `json:"cache_read_tokens"`
|
||||
CacheWriteTokens int64 `json:"cache_write_tokens"`
|
||||
TotalTokens int64 `json:"total_tokens"`
|
||||
}
|
||||
|
||||
type usageFileEvent struct {
|
||||
EventID string `json:"event_id"`
|
||||
Kind string `json:"kind,omitempty"`
|
||||
Status string `json:"status,omitempty"`
|
||||
At time.Time `json:"at"`
|
||||
InputTokens int64 `json:"input_tokens"`
|
||||
OutputTokens int64 `json:"output_tokens"`
|
||||
CacheReadTokens int64 `json:"cache_read_tokens"`
|
||||
CacheWriteTokens int64 `json:"cache_write_tokens"`
|
||||
TotalTokens int64 `json:"total_tokens"`
|
||||
UsagePresent bool `json:"usage_present"`
|
||||
}
|
||||
|
||||
type usageFileDelta struct {
|
||||
providerCalls int64
|
||||
turnsTotal int64
|
||||
validTurnsTotal int64
|
||||
invalidTurnsTotal int64
|
||||
inputTokens int64
|
||||
outputTokens int64
|
||||
cacheReadTokens int64
|
||||
cacheWriteTokens int64
|
||||
totalTokens int64
|
||||
}
|
||||
|
||||
func NewUsageFileStore(historyRoot string) *UsageFileStore {
|
||||
return &UsageFileStore{path: filepath.Join(strings.TrimSpace(historyRoot), usageFileName)}
|
||||
}
|
||||
|
||||
func (store *UsageFileStore) UpsertEvent(event usageFileEvent) error {
|
||||
if store == nil || strings.TrimSpace(store.path) == "" {
|
||||
return nil
|
||||
}
|
||||
event.EventID = strings.TrimSpace(event.EventID)
|
||||
if event.EventID == "" {
|
||||
return nil
|
||||
}
|
||||
event.Kind = normalizeUsageEventKind(event.Kind)
|
||||
event.Status = strings.TrimSpace(event.Status)
|
||||
if event.At.IsZero() {
|
||||
event.At = time.Now().UTC()
|
||||
} else {
|
||||
event.At = event.At.UTC()
|
||||
}
|
||||
event.InputTokens = nonNegativeInt64(event.InputTokens)
|
||||
event.OutputTokens = nonNegativeInt64(event.OutputTokens)
|
||||
event.CacheReadTokens = nonNegativeInt64(event.CacheReadTokens)
|
||||
event.CacheWriteTokens = nonNegativeInt64(event.CacheWriteTokens)
|
||||
event.TotalTokens = event.InputTokens + event.OutputTokens + event.CacheReadTokens + event.CacheWriteTokens
|
||||
|
||||
if err := os.MkdirAll(filepath.Dir(store.path), 0o755); err != nil {
|
||||
return fmt.Errorf("create usage directory: %w", err)
|
||||
}
|
||||
release, err := acquireConversationLock(store.path + ".lock")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer release()
|
||||
|
||||
doc, err := readUsageFileDocument(store.path)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if doc.EventIndex == nil {
|
||||
doc.EventIndex = make(map[string]usageFileEvent)
|
||||
}
|
||||
oldEvent, found := doc.EventIndex[event.EventID]
|
||||
if found {
|
||||
applyUsageFileDelta(&doc, oldEvent.At, negateUsageFileDelta(usageFileEventDelta(oldEvent)))
|
||||
}
|
||||
applyUsageFileDelta(&doc, event.At, usageFileEventDelta(event))
|
||||
doc.RecentEvents = upsertRecentUsageEvent(doc.RecentEvents, event)
|
||||
doc.RecentEvents = trimRecentUsageEvents(doc.RecentEvents, usageRecentEventLimit)
|
||||
doc.EventIndex = buildUsageEventIndex(doc.RecentEvents)
|
||||
doc.SchemaVersion = usageFileSchemaVersion
|
||||
doc.UpdatedAt = time.Now().UTC()
|
||||
return writeJSONFileAtomic(store.path, doc)
|
||||
}
|
||||
|
||||
func (store *UsageFileStore) LookupEvent(needle string) (usageFileEvent, bool, error) {
|
||||
if store == nil || strings.TrimSpace(store.path) == "" {
|
||||
return usageFileEvent{}, false, nil
|
||||
}
|
||||
doc, err := readUsageFileDocument(store.path)
|
||||
if err != nil {
|
||||
return usageFileEvent{}, false, err
|
||||
}
|
||||
trimmed := strings.TrimSpace(needle)
|
||||
if trimmed == "" {
|
||||
return usageFileEvent{}, false, nil
|
||||
}
|
||||
var aggregate usageFileEvent
|
||||
found := false
|
||||
events := doc.EventIndex
|
||||
if len(events) == 0 {
|
||||
events = make(map[string]usageFileEvent, len(doc.RecentEvents))
|
||||
for _, event := range doc.RecentEvents {
|
||||
if eventID := strings.TrimSpace(event.EventID); eventID != "" {
|
||||
events[eventID] = event
|
||||
}
|
||||
}
|
||||
}
|
||||
for _, event := range events {
|
||||
eventID := strings.TrimSpace(event.EventID)
|
||||
if eventID != trimmed && !strings.HasPrefix(eventID, trimmed+"::") {
|
||||
continue
|
||||
}
|
||||
if !found {
|
||||
aggregate = usageFileEvent{EventID: trimmed, At: event.At}
|
||||
found = true
|
||||
}
|
||||
if event.At.After(aggregate.At) {
|
||||
aggregate.At = event.At
|
||||
}
|
||||
aggregate.InputTokens += nonNegativeInt64(event.InputTokens)
|
||||
aggregate.OutputTokens += nonNegativeInt64(event.OutputTokens)
|
||||
aggregate.CacheReadTokens += nonNegativeInt64(event.CacheReadTokens)
|
||||
aggregate.CacheWriteTokens += nonNegativeInt64(event.CacheWriteTokens)
|
||||
aggregate.TotalTokens += nonNegativeInt64(event.TotalTokens)
|
||||
aggregate.UsagePresent = aggregate.UsagePresent || event.UsagePresent
|
||||
}
|
||||
if found {
|
||||
return aggregate, true, nil
|
||||
}
|
||||
return usageFileEvent{}, false, nil
|
||||
}
|
||||
|
||||
func readUsageFileDocument(path string) (usageFileDocument, error) {
|
||||
body, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
if errors.Is(err, os.ErrNotExist) {
|
||||
return usageFileDocument{
|
||||
SchemaVersion: usageFileSchemaVersion,
|
||||
Daily: make([]usageFileDaily, 0),
|
||||
RecentEvents: make([]usageFileEvent, 0),
|
||||
EventIndex: make(map[string]usageFileEvent),
|
||||
}, nil
|
||||
}
|
||||
return usageFileDocument{}, fmt.Errorf("read usage file: %w", err)
|
||||
}
|
||||
var doc usageFileDocument
|
||||
if err := json.Unmarshal(body, &doc); err != nil {
|
||||
return usageFileDocument{}, fmt.Errorf("decode usage file: %w", err)
|
||||
}
|
||||
if doc.SchemaVersion == 0 {
|
||||
doc.SchemaVersion = 1
|
||||
}
|
||||
doc.RecentEvents = trimRecentUsageEvents(doc.RecentEvents, usageRecentEventLimit)
|
||||
if len(doc.EventIndex) == 0 {
|
||||
doc.EventIndex = buildUsageEventIndex(doc.RecentEvents)
|
||||
}
|
||||
return doc, nil
|
||||
}
|
||||
|
||||
func upsertRecentUsageEvent(items []usageFileEvent, event usageFileEvent) []usageFileEvent {
|
||||
event.EventID = strings.TrimSpace(event.EventID)
|
||||
if event.EventID == "" {
|
||||
return items
|
||||
}
|
||||
next := make([]usageFileEvent, 0, len(items)+1)
|
||||
next = append(next, event)
|
||||
for _, item := range items {
|
||||
if strings.TrimSpace(item.EventID) == event.EventID {
|
||||
continue
|
||||
}
|
||||
next = append(next, item)
|
||||
}
|
||||
return next
|
||||
}
|
||||
|
||||
func trimRecentUsageEvents(items []usageFileEvent, limit int) []usageFileEvent {
|
||||
if limit <= 0 || len(items) <= limit {
|
||||
return items
|
||||
}
|
||||
return items[:limit]
|
||||
}
|
||||
|
||||
func buildUsageEventIndex(items []usageFileEvent) map[string]usageFileEvent {
|
||||
index := make(map[string]usageFileEvent, len(items))
|
||||
for _, event := range items {
|
||||
event.EventID = strings.TrimSpace(event.EventID)
|
||||
if event.EventID == "" {
|
||||
continue
|
||||
}
|
||||
event.Kind = normalizeUsageEventKind(event.Kind)
|
||||
index[event.EventID] = event
|
||||
}
|
||||
return index
|
||||
}
|
||||
|
||||
func normalizeUsageEventKind(kind string) string {
|
||||
switch strings.TrimSpace(kind) {
|
||||
case usageEventKindTurn:
|
||||
return usageEventKindTurn
|
||||
default:
|
||||
return usageEventKindProvider
|
||||
}
|
||||
}
|
||||
|
||||
func usageFileEventDelta(event usageFileEvent) usageFileDelta {
|
||||
switch normalizeUsageEventKind(event.Kind) {
|
||||
case usageEventKindTurn:
|
||||
delta := usageFileDelta{turnsTotal: 1}
|
||||
if strings.TrimSpace(event.Status) == usageTurnStatusDone {
|
||||
delta.validTurnsTotal = 1
|
||||
} else {
|
||||
delta.invalidTurnsTotal = 1
|
||||
}
|
||||
return delta
|
||||
default:
|
||||
return usageFileDelta{
|
||||
providerCalls: 1,
|
||||
inputTokens: nonNegativeInt64(event.InputTokens),
|
||||
outputTokens: nonNegativeInt64(event.OutputTokens),
|
||||
cacheReadTokens: nonNegativeInt64(event.CacheReadTokens),
|
||||
cacheWriteTokens: nonNegativeInt64(event.CacheWriteTokens),
|
||||
totalTokens: nonNegativeInt64(event.TotalTokens),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func negateUsageFileDelta(value usageFileDelta) usageFileDelta {
|
||||
return usageFileDelta{
|
||||
providerCalls: -value.providerCalls,
|
||||
turnsTotal: -value.turnsTotal,
|
||||
validTurnsTotal: -value.validTurnsTotal,
|
||||
invalidTurnsTotal: -value.invalidTurnsTotal,
|
||||
inputTokens: -value.inputTokens,
|
||||
outputTokens: -value.outputTokens,
|
||||
cacheReadTokens: -value.cacheReadTokens,
|
||||
cacheWriteTokens: -value.cacheWriteTokens,
|
||||
totalTokens: -value.totalTokens,
|
||||
}
|
||||
}
|
||||
|
||||
func applyUsageFileDelta(doc *usageFileDocument, at time.Time, delta usageFileDelta) {
|
||||
if doc == nil {
|
||||
return
|
||||
}
|
||||
doc.Totals.ProviderCalls = clampNonNegativeInt64(doc.Totals.ProviderCalls + delta.providerCalls)
|
||||
doc.Totals.TurnsTotal = clampNonNegativeInt64(doc.Totals.TurnsTotal + delta.turnsTotal)
|
||||
doc.Totals.ValidTurnsTotal = clampNonNegativeInt64(doc.Totals.ValidTurnsTotal + delta.validTurnsTotal)
|
||||
doc.Totals.InvalidTurnsTotal = clampNonNegativeInt64(doc.Totals.InvalidTurnsTotal + delta.invalidTurnsTotal)
|
||||
doc.Totals.InputTokens = clampNonNegativeInt64(doc.Totals.InputTokens + delta.inputTokens)
|
||||
doc.Totals.OutputTokens = clampNonNegativeInt64(doc.Totals.OutputTokens + delta.outputTokens)
|
||||
doc.Totals.CacheReadTokens = clampNonNegativeInt64(doc.Totals.CacheReadTokens + delta.cacheReadTokens)
|
||||
doc.Totals.CacheWriteTokens = clampNonNegativeInt64(doc.Totals.CacheWriteTokens + delta.cacheWriteTokens)
|
||||
doc.Totals.TotalTokens = clampNonNegativeInt64(doc.Totals.TotalTokens + delta.totalTokens)
|
||||
|
||||
date := at.UTC().Format("2006-01-02")
|
||||
for index := range doc.Daily {
|
||||
if doc.Daily[index].Date != date {
|
||||
continue
|
||||
}
|
||||
applyUsageDailyDelta(&doc.Daily[index], delta)
|
||||
return
|
||||
}
|
||||
item := usageFileDaily{Date: date}
|
||||
applyUsageDailyDelta(&item, delta)
|
||||
doc.Daily = append(doc.Daily, item)
|
||||
}
|
||||
|
||||
func applyUsageDailyDelta(item *usageFileDaily, delta usageFileDelta) {
|
||||
if item == nil {
|
||||
return
|
||||
}
|
||||
item.ProviderCalls = clampNonNegativeInt64(item.ProviderCalls + delta.providerCalls)
|
||||
item.TurnsTotal = clampNonNegativeInt64(item.TurnsTotal + delta.turnsTotal)
|
||||
item.ValidTurnsTotal = clampNonNegativeInt64(item.ValidTurnsTotal + delta.validTurnsTotal)
|
||||
item.InvalidTurnsTotal = clampNonNegativeInt64(item.InvalidTurnsTotal + delta.invalidTurnsTotal)
|
||||
item.InputTokens = clampNonNegativeInt64(item.InputTokens + delta.inputTokens)
|
||||
item.OutputTokens = clampNonNegativeInt64(item.OutputTokens + delta.outputTokens)
|
||||
item.CacheReadTokens = clampNonNegativeInt64(item.CacheReadTokens + delta.cacheReadTokens)
|
||||
item.CacheWriteTokens = clampNonNegativeInt64(item.CacheWriteTokens + delta.cacheWriteTokens)
|
||||
item.TotalTokens = clampNonNegativeInt64(item.TotalTokens + delta.totalTokens)
|
||||
}
|
||||
|
||||
func clampNonNegativeInt64(value int64) int64 {
|
||||
if value < 0 {
|
||||
return 0
|
||||
}
|
||||
return value
|
||||
}
|
||||
@@ -0,0 +1,607 @@
|
||||
package forwarder
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"errors"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sort"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"connectrpc.com/connect"
|
||||
"github.com/google/uuid"
|
||||
|
||||
"cursor/gen/aiserverv1"
|
||||
"cursor/internal/logger"
|
||||
)
|
||||
|
||||
const sharedUserRuleExtension = ".md"
|
||||
|
||||
type UserRuleRecord struct {
|
||||
ID string
|
||||
Title string
|
||||
Filename string
|
||||
FullPath string
|
||||
Knowledge string
|
||||
CreatedAt string
|
||||
IsGenerated bool
|
||||
ModifiedAt time.Time
|
||||
ContentHash string
|
||||
}
|
||||
|
||||
type userRuleGroup struct {
|
||||
ContentHash string
|
||||
Canonical UserRuleRecord
|
||||
Members []UserRuleRecord
|
||||
}
|
||||
|
||||
type UserRuleStore struct {
|
||||
root string
|
||||
mu sync.Mutex
|
||||
}
|
||||
|
||||
func NewUserRuleStore(root string) *UserRuleStore {
|
||||
return &UserRuleStore{root: strings.TrimSpace(root)}
|
||||
}
|
||||
|
||||
func (store *UserRuleStore) List() ([]UserRuleRecord, error) {
|
||||
if store == nil {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
store.mu.Lock()
|
||||
defer store.mu.Unlock()
|
||||
|
||||
groups, err := store.listGroupsLocked()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
records := canonicalRecordsFromGroups(groups)
|
||||
sort.SliceStable(records, func(i int, j int) bool {
|
||||
if records[i].ModifiedAt.Equal(records[j].ModifiedAt) {
|
||||
return records[i].Filename < records[j].Filename
|
||||
}
|
||||
return records[i].ModifiedAt.After(records[j].ModifiedAt)
|
||||
})
|
||||
return records, nil
|
||||
}
|
||||
|
||||
func (store *UserRuleStore) Add(knowledge string) (UserRuleRecord, error) {
|
||||
if store == nil {
|
||||
return UserRuleRecord{}, fmt.Errorf("user rule store is nil")
|
||||
}
|
||||
if strings.TrimSpace(knowledge) == "" {
|
||||
return UserRuleRecord{}, fmt.Errorf("knowledge is required")
|
||||
}
|
||||
|
||||
store.mu.Lock()
|
||||
defer store.mu.Unlock()
|
||||
|
||||
groups, err := store.listGroupsLocked()
|
||||
if err != nil {
|
||||
return UserRuleRecord{}, err
|
||||
}
|
||||
|
||||
contentHash := hashSharedUserRuleContent([]byte(knowledge))
|
||||
if groupIndex := indexGroupByHash(groups, contentHash); groupIndex >= 0 {
|
||||
if err := store.removeRuleFilesLocked(nonCanonicalMembers(groups[groupIndex].Members)); err != nil {
|
||||
return UserRuleRecord{}, err
|
||||
}
|
||||
return groups[groupIndex].Canonical, nil
|
||||
}
|
||||
|
||||
id := uuid.NewString()
|
||||
if err := store.writeRuleLocked(id, knowledge); err != nil {
|
||||
return UserRuleRecord{}, err
|
||||
}
|
||||
return store.loadRuleByIDLocked(id)
|
||||
}
|
||||
|
||||
func (store *UserRuleStore) Update(id string, knowledge string) (UserRuleRecord, bool, error) {
|
||||
if store == nil {
|
||||
return UserRuleRecord{}, false, fmt.Errorf("user rule store is nil")
|
||||
}
|
||||
if strings.TrimSpace(knowledge) == "" {
|
||||
return UserRuleRecord{}, false, fmt.Errorf("knowledge is required")
|
||||
}
|
||||
|
||||
normalizedID, err := normalizeUserRuleID(id)
|
||||
if err != nil {
|
||||
return UserRuleRecord{}, false, err
|
||||
}
|
||||
|
||||
store.mu.Lock()
|
||||
defer store.mu.Unlock()
|
||||
|
||||
current, err := store.loadRuleByIDLocked(normalizedID)
|
||||
if err != nil {
|
||||
if errors.Is(err, os.ErrNotExist) {
|
||||
return UserRuleRecord{}, false, nil
|
||||
}
|
||||
return UserRuleRecord{}, false, err
|
||||
}
|
||||
|
||||
groups, err := store.listGroupsLocked()
|
||||
if err != nil {
|
||||
return UserRuleRecord{}, false, err
|
||||
}
|
||||
|
||||
currentGroupIndex := indexGroupByHash(groups, current.ContentHash)
|
||||
targetHash := hashSharedUserRuleContent([]byte(knowledge))
|
||||
targetGroupIndex := indexGroupByHash(groups, targetHash)
|
||||
|
||||
if targetGroupIndex >= 0 && groups[targetGroupIndex].ContentHash == current.ContentHash {
|
||||
if currentGroupIndex >= 0 {
|
||||
if err := store.removeRuleFilesLocked(nonCanonicalMembers(groups[currentGroupIndex].Members)); err != nil {
|
||||
return UserRuleRecord{}, false, err
|
||||
}
|
||||
return groups[currentGroupIndex].Canonical, true, nil
|
||||
}
|
||||
return current, true, nil
|
||||
}
|
||||
|
||||
if targetGroupIndex >= 0 {
|
||||
if currentGroupIndex >= 0 {
|
||||
remainingCurrentMembers := groupMembersExcludingIDs(groups[currentGroupIndex], map[string]struct{}{current.ID: {}})
|
||||
if err := store.removeRuleFilesLocked(nonCanonicalMembers(remainingCurrentMembers)); err != nil {
|
||||
return UserRuleRecord{}, false, err
|
||||
}
|
||||
}
|
||||
if err := store.removeRuleFilesLocked(nonCanonicalMembers(groups[targetGroupIndex].Members)); err != nil {
|
||||
return UserRuleRecord{}, false, err
|
||||
}
|
||||
if err := store.removeRuleFilesLocked([]UserRuleRecord{current}); err != nil {
|
||||
return UserRuleRecord{}, false, err
|
||||
}
|
||||
return groups[targetGroupIndex].Canonical, true, nil
|
||||
}
|
||||
|
||||
if err := store.writeRuleLocked(current.ID, knowledge); err != nil {
|
||||
return UserRuleRecord{}, false, err
|
||||
}
|
||||
if currentGroupIndex >= 0 {
|
||||
remainingCurrentMembers := groupMembersExcludingIDs(groups[currentGroupIndex], map[string]struct{}{current.ID: {}})
|
||||
if err := store.removeRuleFilesLocked(nonCanonicalMembers(remainingCurrentMembers)); err != nil {
|
||||
return UserRuleRecord{}, false, err
|
||||
}
|
||||
}
|
||||
record, err := store.loadRuleByIDLocked(current.ID)
|
||||
if err != nil {
|
||||
return UserRuleRecord{}, false, err
|
||||
}
|
||||
return record, true, nil
|
||||
}
|
||||
|
||||
func (store *UserRuleStore) Remove(id string) error {
|
||||
if store == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
normalizedID, err := normalizeUserRuleID(id)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
store.mu.Lock()
|
||||
defer store.mu.Unlock()
|
||||
|
||||
record, err := store.loadRuleByIDLocked(normalizedID)
|
||||
if err != nil {
|
||||
if errors.Is(err, os.ErrNotExist) {
|
||||
return nil
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
groups, err := store.listGroupsLocked()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
groupIndex := indexGroupByHash(groups, record.ContentHash)
|
||||
if groupIndex < 0 {
|
||||
return store.removeRuleFilesLocked([]UserRuleRecord{record})
|
||||
}
|
||||
return store.removeRuleFilesLocked(groups[groupIndex].Members)
|
||||
}
|
||||
|
||||
func (store *UserRuleStore) BuildSystemPromptSection() (string, int, int, error) {
|
||||
if store == nil {
|
||||
return "", 0, 0, nil
|
||||
}
|
||||
|
||||
store.mu.Lock()
|
||||
defer store.mu.Unlock()
|
||||
|
||||
groups, err := store.listGroupsLocked()
|
||||
if err != nil {
|
||||
return "", 0, 0, err
|
||||
}
|
||||
if len(groups) == 0 {
|
||||
return "", 0, 0, nil
|
||||
}
|
||||
|
||||
totalFiles := totalRuleFileCount(groups)
|
||||
records := canonicalRecordsFromGroups(groups)
|
||||
sort.SliceStable(records, func(i int, j int) bool {
|
||||
return records[i].Filename < records[j].Filename
|
||||
})
|
||||
|
||||
lines := []string{
|
||||
`<shared_user_rules description="These shared local rules are loaded from the backend configuration directory and apply to every local conversation. Follow them when relevant.">`,
|
||||
}
|
||||
visibleCount := 0
|
||||
for _, record := range records {
|
||||
if strings.TrimSpace(record.Knowledge) == "" {
|
||||
continue
|
||||
}
|
||||
lines = append(lines,
|
||||
fmt.Sprintf(`<rule file="%s">`, escapeSharedRulePromptText(record.Filename)),
|
||||
escapeSharedRulePromptText(record.Knowledge),
|
||||
"</rule>",
|
||||
)
|
||||
visibleCount++
|
||||
}
|
||||
if visibleCount == 0 {
|
||||
return "", totalFiles, 0, nil
|
||||
}
|
||||
lines = append(lines, "</shared_user_rules>")
|
||||
return strings.Join(lines, "\n"), totalFiles, visibleCount, nil
|
||||
}
|
||||
|
||||
func (store *UserRuleStore) listGroupsLocked() ([]userRuleGroup, error) {
|
||||
records, err := store.scanRuleFilesLocked()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return buildUserRuleGroups(records), nil
|
||||
}
|
||||
|
||||
func (store *UserRuleStore) scanRuleFilesLocked() ([]UserRuleRecord, error) {
|
||||
if err := store.ensureRootLocked(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
entries, err := os.ReadDir(store.root)
|
||||
if err != nil {
|
||||
if errors.Is(err, os.ErrNotExist) {
|
||||
return nil, nil
|
||||
}
|
||||
return nil, fmt.Errorf("read shared user rules directory: %w", err)
|
||||
}
|
||||
|
||||
records := make([]UserRuleRecord, 0, len(entries))
|
||||
for _, entry := range entries {
|
||||
if entry.IsDir() || filepath.Ext(entry.Name()) != sharedUserRuleExtension {
|
||||
continue
|
||||
}
|
||||
record, err := store.readRuleFileLocked(filepath.Join(store.root, entry.Name()))
|
||||
if err != nil {
|
||||
logger.Infof("跳过不可用的共享 user rule 文件 path=%s err=%v", filepath.Join(store.root, entry.Name()), err)
|
||||
continue
|
||||
}
|
||||
records = append(records, record)
|
||||
}
|
||||
return records, nil
|
||||
}
|
||||
|
||||
func (store *UserRuleStore) loadRuleByIDLocked(id string) (UserRuleRecord, error) {
|
||||
return store.readRuleFileLocked(store.rulePathLocked(id))
|
||||
}
|
||||
|
||||
func (store *UserRuleStore) readRuleFileLocked(path string) (UserRuleRecord, error) {
|
||||
data, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
return UserRuleRecord{}, err
|
||||
}
|
||||
if strings.TrimSpace(string(data)) == "" {
|
||||
return UserRuleRecord{}, fmt.Errorf("shared user rule file is empty")
|
||||
}
|
||||
|
||||
info, err := os.Stat(path)
|
||||
if err != nil {
|
||||
return UserRuleRecord{}, err
|
||||
}
|
||||
|
||||
filename := filepath.Base(path)
|
||||
id, err := normalizeUserRuleID(strings.TrimSuffix(filename, filepath.Ext(filename)))
|
||||
if err != nil {
|
||||
return UserRuleRecord{}, err
|
||||
}
|
||||
|
||||
modifiedAt := info.ModTime().UTC()
|
||||
return UserRuleRecord{
|
||||
ID: id,
|
||||
Title: id,
|
||||
Filename: filename,
|
||||
FullPath: path,
|
||||
Knowledge: string(data),
|
||||
CreatedAt: modifiedAt.Format(time.RFC3339Nano),
|
||||
IsGenerated: false,
|
||||
ModifiedAt: modifiedAt,
|
||||
ContentHash: hashSharedUserRuleContent(data),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (store *UserRuleStore) writeRuleLocked(id string, knowledge string) error {
|
||||
normalizedID, err := normalizeUserRuleID(id)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if strings.TrimSpace(knowledge) == "" {
|
||||
return fmt.Errorf("knowledge is required")
|
||||
}
|
||||
if err := store.ensureRootLocked(); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
path := store.rulePathLocked(normalizedID)
|
||||
if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil {
|
||||
return fmt.Errorf("create shared user rules directory: %w", err)
|
||||
}
|
||||
|
||||
tempPath := path + ".tmp"
|
||||
if err := os.WriteFile(tempPath, []byte(knowledge), 0o644); err != nil {
|
||||
return fmt.Errorf("write temp shared user rule: %w", err)
|
||||
}
|
||||
if err := os.Rename(tempPath, path); err != nil {
|
||||
return fmt.Errorf("rename shared user rule file: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (store *UserRuleStore) removeRuleFilesLocked(records []UserRuleRecord) error {
|
||||
for _, record := range records {
|
||||
path := strings.TrimSpace(record.FullPath)
|
||||
if path == "" {
|
||||
path = store.rulePathLocked(record.ID)
|
||||
}
|
||||
if strings.TrimSpace(path) == "" {
|
||||
continue
|
||||
}
|
||||
if err := os.Remove(path); err != nil && !errors.Is(err, os.ErrNotExist) {
|
||||
return fmt.Errorf("remove shared user rule %q: %w", record.ID, err)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (store *UserRuleStore) ensureRootLocked() error {
|
||||
if strings.TrimSpace(store.root) == "" {
|
||||
return fmt.Errorf("user rules root is empty")
|
||||
}
|
||||
if err := os.MkdirAll(store.root, 0o755); err != nil {
|
||||
return fmt.Errorf("create user rules root: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (store *UserRuleStore) rulePathLocked(id string) string {
|
||||
if store == nil || strings.TrimSpace(store.root) == "" {
|
||||
return ""
|
||||
}
|
||||
return filepath.Join(store.root, id+sharedUserRuleExtension)
|
||||
}
|
||||
|
||||
func buildUserRuleGroups(records []UserRuleRecord) []userRuleGroup {
|
||||
if len(records) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
membersByHash := make(map[string][]UserRuleRecord)
|
||||
for _, record := range records {
|
||||
membersByHash[record.ContentHash] = append(membersByHash[record.ContentHash], record)
|
||||
}
|
||||
|
||||
groups := make([]userRuleGroup, 0, len(membersByHash))
|
||||
for contentHash, members := range membersByHash {
|
||||
sort.SliceStable(members, func(i int, j int) bool {
|
||||
return members[i].Filename < members[j].Filename
|
||||
})
|
||||
groups = append(groups, userRuleGroup{
|
||||
ContentHash: contentHash,
|
||||
Canonical: members[0],
|
||||
Members: members,
|
||||
})
|
||||
}
|
||||
|
||||
sort.SliceStable(groups, func(i int, j int) bool {
|
||||
return groups[i].Canonical.Filename < groups[j].Canonical.Filename
|
||||
})
|
||||
return groups
|
||||
}
|
||||
|
||||
func canonicalRecordsFromGroups(groups []userRuleGroup) []UserRuleRecord {
|
||||
if len(groups) == 0 {
|
||||
return nil
|
||||
}
|
||||
records := make([]UserRuleRecord, 0, len(groups))
|
||||
for _, group := range groups {
|
||||
records = append(records, group.Canonical)
|
||||
}
|
||||
return records
|
||||
}
|
||||
|
||||
func totalRuleFileCount(groups []userRuleGroup) int {
|
||||
total := 0
|
||||
for _, group := range groups {
|
||||
total += len(group.Members)
|
||||
}
|
||||
return total
|
||||
}
|
||||
|
||||
func indexGroupByHash(groups []userRuleGroup, contentHash string) int {
|
||||
for index := range groups {
|
||||
if groups[index].ContentHash == contentHash {
|
||||
return index
|
||||
}
|
||||
}
|
||||
return -1
|
||||
}
|
||||
|
||||
func nonCanonicalMembers(records []UserRuleRecord) []UserRuleRecord {
|
||||
if len(records) <= 1 {
|
||||
return nil
|
||||
}
|
||||
output := make([]UserRuleRecord, 0, len(records)-1)
|
||||
output = append(output, records[1:]...)
|
||||
return output
|
||||
}
|
||||
|
||||
func groupMembersExcludingIDs(group userRuleGroup, excludedIDs map[string]struct{}) []UserRuleRecord {
|
||||
if len(group.Members) == 0 {
|
||||
return nil
|
||||
}
|
||||
filtered := make([]UserRuleRecord, 0, len(group.Members))
|
||||
for _, record := range group.Members {
|
||||
if _, excluded := excludedIDs[record.ID]; excluded {
|
||||
continue
|
||||
}
|
||||
filtered = append(filtered, record)
|
||||
}
|
||||
return filtered
|
||||
}
|
||||
|
||||
func normalizeUserRuleID(raw string) (string, error) {
|
||||
id := strings.TrimSpace(raw)
|
||||
id = strings.TrimSuffix(id, sharedUserRuleExtension)
|
||||
switch {
|
||||
case id == "":
|
||||
return "", fmt.Errorf("user rule id is required")
|
||||
case strings.Contains(id, "/"), strings.Contains(id, "\\"), strings.Contains(id, ".."):
|
||||
return "", fmt.Errorf("invalid user rule id %q", raw)
|
||||
default:
|
||||
return id, nil
|
||||
}
|
||||
}
|
||||
|
||||
func hashSharedUserRuleContent(data []byte) string {
|
||||
sum := sha256.Sum256(data)
|
||||
return hex.EncodeToString(sum[:])
|
||||
}
|
||||
|
||||
func escapeSharedRulePromptText(value string) string {
|
||||
replacer := strings.NewReplacer(
|
||||
"&", "&",
|
||||
`"`, """,
|
||||
"<", "<",
|
||||
">", ">",
|
||||
)
|
||||
return replacer.Replace(value)
|
||||
}
|
||||
|
||||
func (service *Service) KnowledgeBaseAdd(_ context.Context, req *connect.Request[aiserverv1.KnowledgeBaseAddRequest]) (*connect.Response[aiserverv1.KnowledgeBaseAddResponse], error) {
|
||||
if service == nil || service.rules == nil {
|
||||
return nil, connect.NewError(connect.CodeInternal, fmt.Errorf("user rule store is not initialized"))
|
||||
}
|
||||
if strings.TrimSpace(req.Msg.GetKnowledge()) == "" {
|
||||
return nil, connect.NewError(connect.CodeInvalidArgument, fmt.Errorf("knowledge is required"))
|
||||
}
|
||||
|
||||
record, err := service.rules.Add(req.Msg.GetKnowledge())
|
||||
if err != nil {
|
||||
return nil, connect.NewError(connect.CodeInternal, err)
|
||||
}
|
||||
if service.docsIndexStore != nil {
|
||||
if _, err := service.docsIndexStore.Upsert(DocsIndexRecord{
|
||||
ID: record.ID,
|
||||
Identifier: record.ID,
|
||||
Title: firstNonEmptyDocs(req.Msg.GetTitle(), record.Title, record.ID),
|
||||
URL: firstNonEmptyDocs(docsIndexURLCandidate(req.Msg.GetKnowledge()), docsIndexURLCandidate(req.Msg.GetTitle())),
|
||||
Content: req.Msg.GetKnowledge(),
|
||||
GitOrigin: req.Msg.GetGitOrigin(),
|
||||
Source: docsIndexSourceLocal,
|
||||
}); err != nil {
|
||||
logger.Errorf("docs index sync failed operation=add id=%s err=%v", record.ID, err)
|
||||
}
|
||||
}
|
||||
return connect.NewResponse(&aiserverv1.KnowledgeBaseAddResponse{
|
||||
Success: true,
|
||||
Id: record.ID,
|
||||
}), nil
|
||||
}
|
||||
|
||||
func (service *Service) KnowledgeBaseList(_ context.Context, req *connect.Request[aiserverv1.KnowledgeBaseListRequest]) (*connect.Response[aiserverv1.KnowledgeBaseListResponse], error) {
|
||||
if service == nil || service.rules == nil {
|
||||
return nil, connect.NewError(connect.CodeInternal, fmt.Errorf("user rule store is not initialized"))
|
||||
}
|
||||
|
||||
records, err := service.rules.List()
|
||||
if err != nil {
|
||||
return nil, connect.NewError(connect.CodeInternal, err)
|
||||
}
|
||||
if limit := req.Msg.GetLimit(); limit > 0 && len(records) > int(limit) {
|
||||
records = records[:limit]
|
||||
}
|
||||
|
||||
items := make([]*aiserverv1.KnowledgeBaseListResponse_Item, 0, len(records))
|
||||
for _, record := range records {
|
||||
items = append(items, &aiserverv1.KnowledgeBaseListResponse_Item{
|
||||
Id: record.ID,
|
||||
Knowledge: record.Knowledge,
|
||||
Title: record.Title,
|
||||
CreatedAt: record.CreatedAt,
|
||||
IsGenerated: false,
|
||||
})
|
||||
}
|
||||
return connect.NewResponse(&aiserverv1.KnowledgeBaseListResponse{
|
||||
Success: true,
|
||||
AllResults: items,
|
||||
}), nil
|
||||
}
|
||||
|
||||
func (service *Service) KnowledgeBaseUpdate(_ context.Context, req *connect.Request[aiserverv1.KnowledgeBaseUpdateRequest]) (*connect.Response[aiserverv1.KnowledgeBaseUpdateResponse], error) {
|
||||
if service == nil || service.rules == nil {
|
||||
return nil, connect.NewError(connect.CodeInternal, fmt.Errorf("user rule store is not initialized"))
|
||||
}
|
||||
if _, err := normalizeUserRuleID(req.Msg.GetId()); err != nil {
|
||||
return nil, connect.NewError(connect.CodeInvalidArgument, err)
|
||||
}
|
||||
if strings.TrimSpace(req.Msg.GetKnowledge()) == "" {
|
||||
return nil, connect.NewError(connect.CodeInvalidArgument, fmt.Errorf("knowledge is required"))
|
||||
}
|
||||
|
||||
_, exists, err := service.rules.Update(req.Msg.GetId(), req.Msg.GetKnowledge())
|
||||
if err != nil {
|
||||
return nil, connect.NewError(connect.CodeInternal, err)
|
||||
}
|
||||
if !exists {
|
||||
return nil, connect.NewError(connect.CodeNotFound, fmt.Errorf("user rule %q not found", strings.TrimSpace(req.Msg.GetId())))
|
||||
}
|
||||
if service.docsIndexStore != nil {
|
||||
if _, err := service.docsIndexStore.Upsert(DocsIndexRecord{
|
||||
ID: strings.TrimSpace(req.Msg.GetId()),
|
||||
Identifier: strings.TrimSpace(req.Msg.GetId()),
|
||||
Title: firstNonEmptyDocs(req.Msg.GetTitle(), strings.TrimSpace(req.Msg.GetId())),
|
||||
URL: firstNonEmptyDocs(docsIndexURLCandidate(req.Msg.GetKnowledge()), docsIndexURLCandidate(req.Msg.GetTitle())),
|
||||
Content: req.Msg.GetKnowledge(),
|
||||
Source: docsIndexSourceLocal,
|
||||
}); err != nil {
|
||||
logger.Errorf("docs index sync failed operation=update id=%s err=%v", strings.TrimSpace(req.Msg.GetId()), err)
|
||||
}
|
||||
}
|
||||
return connect.NewResponse(&aiserverv1.KnowledgeBaseUpdateResponse{Success: true}), nil
|
||||
}
|
||||
|
||||
func (service *Service) KnowledgeBaseRemove(_ context.Context, req *connect.Request[aiserverv1.KnowledgeBaseRemoveRequest]) (*connect.Response[aiserverv1.KnowledgeBaseRemoveResponse], error) {
|
||||
if service == nil || service.rules == nil {
|
||||
return nil, connect.NewError(connect.CodeInternal, fmt.Errorf("user rule store is not initialized"))
|
||||
}
|
||||
if _, err := normalizeUserRuleID(req.Msg.GetId()); err != nil {
|
||||
return nil, connect.NewError(connect.CodeInvalidArgument, err)
|
||||
}
|
||||
if err := service.rules.Remove(req.Msg.GetId()); err != nil {
|
||||
return nil, connect.NewError(connect.CodeInternal, err)
|
||||
}
|
||||
if service.docsIndexStore != nil {
|
||||
if err := service.docsIndexStore.Remove(req.Msg.GetId()); err != nil {
|
||||
logger.Errorf("docs index sync failed operation=remove id=%s err=%v", strings.TrimSpace(req.Msg.GetId()), err)
|
||||
}
|
||||
}
|
||||
return connect.NewResponse(&aiserverv1.KnowledgeBaseRemoveResponse{Success: true}), nil
|
||||
}
|
||||
@@ -0,0 +1,455 @@
|
||||
package forwarder
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"google.golang.org/protobuf/encoding/protojson"
|
||||
|
||||
"cursor/gen/agentv1"
|
||||
execbridge "cursor/internal/backend/agent/bridge/exec"
|
||||
runtimecore "cursor/internal/backend/agent/core"
|
||||
)
|
||||
|
||||
const (
|
||||
writeReadExecKind = "write_read"
|
||||
writeWriteExecKind = "write_write"
|
||||
writePostReadExecKind = "write_post_read"
|
||||
)
|
||||
|
||||
type writeOperationArgs struct {
|
||||
Path string `json:"path"`
|
||||
Contents string `json:"contents"`
|
||||
}
|
||||
|
||||
type pendingWritePayload struct {
|
||||
VisibleArgs writeOperationArgs `json:"visible_args"`
|
||||
ResolvedPath string `json:"resolved_path"`
|
||||
BeforeContent string `json:"before_content,omitempty"`
|
||||
AfterContent string `json:"after_content,omitempty"`
|
||||
}
|
||||
|
||||
func isHiddenWriteExecKind(kind string) bool {
|
||||
switch strings.TrimSpace(kind) {
|
||||
case writeReadExecKind, writeWriteExecKind, writePostReadExecKind:
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func (service *Service) handleWriteToolInvocation(stream *ActiveStream, invocation runtimecore.ToolInvocation) error {
|
||||
if stream == nil {
|
||||
return fmt.Errorf("write stream is required")
|
||||
}
|
||||
|
||||
writeArgs, err := decodeWriteOperationArgs(invocation.ArgsJSON)
|
||||
if err != nil {
|
||||
return newRecoverableToolInvocationError(err)
|
||||
}
|
||||
|
||||
resolvedPath := strings.TrimSpace(writeArgs.Path)
|
||||
writeArgs.Path = resolvedPath
|
||||
|
||||
visibleArgsJSON, err := writeArgs.MarshalJSON()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
invocation.ArgsJSON = visibleArgsJSON
|
||||
|
||||
startedToolCall := buildStartedToolCall(invocation)
|
||||
if startedToolCall != nil {
|
||||
historyStartedToolCall := buildStartedWriteHistoryToolCall(writeArgs.Path)
|
||||
toolCallPayload, err := protojson.Marshal(historyStartedToolCall)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
_, err = service.appendConversationEntries(stream, stream.ConversationID, []HistoryEntry{
|
||||
newToolCallEntryWithProviderMetadata(stream.TurnSeq, stream.RequestID, invocation.CallID, invocation.ToolName, invocation.ReasoningContent, invocation.ReasoningSignature, invocation.ReasoningSignatureSource, invocation.ReasoningProviderItemID, invocation.ReasoningProviderStatus, invocation.ReasoningProviderSummary, invocation.ProviderItemID, invocation.ProviderCallID, invocation.ProviderStatus, toolCallPayload),
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
if err := service.broker.Publish(stream.RequestID, StreamEvent{
|
||||
Message: buildToolCallStartedMessage(invocation.CallID, invocation.ModelCallID, startedToolCall),
|
||||
}); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
return service.startHiddenWriteRead(stream, invocation.CallID, invocation.ModelCallID, currentProviderPass(stream), invocation.ReasoningContent, invocation.ReasoningSignature, invocation.ReasoningSignatureSource, pendingWritePayload{
|
||||
VisibleArgs: writeArgs,
|
||||
ResolvedPath: resolvedPath,
|
||||
})
|
||||
}
|
||||
|
||||
func (service *Service) startHiddenWriteRead(stream *ActiveStream, toolCallID string, modelCallID string, providerPass int, reasoningContent string, reasoningSignature string, reasoningSignatureSource string, payload pendingWritePayload) error {
|
||||
readArgsJSON, err := json.Marshal(map[string]any{
|
||||
"path": strings.TrimSpace(payload.ResolvedPath),
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
serverMessage, pendingExec, err := service.execBridge.OpenExec(execbridge.OpenExecContext{
|
||||
ConversationID: stream.ConversationID,
|
||||
}, runtimecore.ToolInvocation{
|
||||
CallID: strings.TrimSpace(toolCallID),
|
||||
ToolName: "Read",
|
||||
ArgsJSON: readArgsJSON,
|
||||
ModelCallID: strings.TrimSpace(modelCallID),
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
pendingArgsJSON, err := payload.MarshalJSON()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
pendingExec.ModelCallID = strings.TrimSpace(modelCallID)
|
||||
pendingExec.ProviderPass = providerPass
|
||||
pendingExec.ToolCallID = strings.TrimSpace(toolCallID)
|
||||
pendingExec.ReasoningContent = reasoningContent
|
||||
pendingExec.ReasoningSignature = strings.TrimSpace(reasoningSignature)
|
||||
pendingExec.ReasoningSignatureSource = strings.TrimSpace(reasoningSignatureSource)
|
||||
pendingExec.ExecKind = writeReadExecKind
|
||||
pendingExec.ArgsJSON = pendingArgsJSON
|
||||
stream.mu.Lock()
|
||||
stream.PendingExecs[pendingExec.ExecID] = pendingExec
|
||||
stream.mu.Unlock()
|
||||
if err := service.publishCheckpoint(stream.RequestID, stream.ConversationID); err != nil {
|
||||
return err
|
||||
}
|
||||
return service.broker.Publish(stream.RequestID, StreamEvent{Message: serverMessage})
|
||||
}
|
||||
|
||||
func (service *Service) startHiddenWriteExec(stream *ActiveStream, toolCallID string, modelCallID string, providerPass int, reasoningContent string, reasoningSignature string, reasoningSignatureSource string, payload pendingWritePayload, beforeContent string) error {
|
||||
transportContents := prepareWriteContentsForClient(payload.ResolvedPath, payload.VisibleArgs.Contents)
|
||||
writeArgsJSON, err := json.Marshal(map[string]any{
|
||||
"path": strings.TrimSpace(payload.ResolvedPath),
|
||||
"contents": transportContents,
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
serverMessage, pendingExec, err := service.execBridge.OpenExec(execbridge.OpenExecContext{
|
||||
ConversationID: stream.ConversationID,
|
||||
}, runtimecore.ToolInvocation{
|
||||
CallID: strings.TrimSpace(toolCallID),
|
||||
ToolName: "Write",
|
||||
ArgsJSON: writeArgsJSON,
|
||||
ModelCallID: strings.TrimSpace(modelCallID),
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
payload.BeforeContent = beforeContent
|
||||
payload.AfterContent = payload.VisibleArgs.Contents
|
||||
pendingArgsJSON, err := payload.MarshalJSON()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
pendingExec.ModelCallID = strings.TrimSpace(modelCallID)
|
||||
pendingExec.ProviderPass = providerPass
|
||||
pendingExec.ToolCallID = strings.TrimSpace(toolCallID)
|
||||
pendingExec.ReasoningContent = reasoningContent
|
||||
pendingExec.ReasoningSignature = strings.TrimSpace(reasoningSignature)
|
||||
pendingExec.ReasoningSignatureSource = strings.TrimSpace(reasoningSignatureSource)
|
||||
pendingExec.ExecKind = writeWriteExecKind
|
||||
pendingExec.ArgsJSON = pendingArgsJSON
|
||||
stream.mu.Lock()
|
||||
stream.PendingExecs[pendingExec.ExecID] = pendingExec
|
||||
stream.mu.Unlock()
|
||||
if err := service.publishCheckpoint(stream.RequestID, stream.ConversationID); err != nil {
|
||||
return err
|
||||
}
|
||||
return service.broker.Publish(stream.RequestID, StreamEvent{Message: serverMessage})
|
||||
}
|
||||
|
||||
func (service *Service) startHiddenWritePostRead(stream *ActiveStream, toolCallID string, modelCallID string, providerPass int, reasoningContent string, reasoningSignature string, reasoningSignatureSource string, payload pendingWritePayload) error {
|
||||
readArgsJSON, err := json.Marshal(map[string]any{
|
||||
"path": strings.TrimSpace(payload.ResolvedPath),
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
serverMessage, pendingExec, err := service.execBridge.OpenExec(execbridge.OpenExecContext{
|
||||
ConversationID: stream.ConversationID,
|
||||
}, runtimecore.ToolInvocation{
|
||||
CallID: strings.TrimSpace(toolCallID),
|
||||
ToolName: "Read",
|
||||
ArgsJSON: readArgsJSON,
|
||||
ModelCallID: strings.TrimSpace(modelCallID),
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
pendingArgsJSON, err := payload.MarshalJSON()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
pendingExec.ModelCallID = strings.TrimSpace(modelCallID)
|
||||
pendingExec.ProviderPass = providerPass
|
||||
pendingExec.ToolCallID = strings.TrimSpace(toolCallID)
|
||||
pendingExec.ReasoningContent = reasoningContent
|
||||
pendingExec.ReasoningSignature = strings.TrimSpace(reasoningSignature)
|
||||
pendingExec.ReasoningSignatureSource = strings.TrimSpace(reasoningSignatureSource)
|
||||
pendingExec.ExecKind = writePostReadExecKind
|
||||
pendingExec.ArgsJSON = pendingArgsJSON
|
||||
stream.mu.Lock()
|
||||
stream.PendingExecs[pendingExec.ExecID] = pendingExec
|
||||
stream.mu.Unlock()
|
||||
if err := service.publishCheckpoint(stream.RequestID, stream.ConversationID); err != nil {
|
||||
return err
|
||||
}
|
||||
return service.broker.Publish(stream.RequestID, StreamEvent{Message: serverMessage})
|
||||
}
|
||||
|
||||
func (service *Service) handleHiddenWriteExecResult(stream *ActiveStream, pending runtimecore.PendingExec, message *agentv1.ExecClientMessage) error {
|
||||
if stream == nil || message == nil {
|
||||
return fmt.Errorf("write exec result requires stream and message")
|
||||
}
|
||||
payload, err := decodePendingWritePayload(pending.ArgsJSON)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
switch strings.TrimSpace(pending.ExecKind) {
|
||||
case writeReadExecKind:
|
||||
markExecCompleted(stream, pending)
|
||||
readResult := message.GetReadResult()
|
||||
if readResult == nil {
|
||||
return service.finishWriteOperation(stream, pending.ToolCallID, pending.ModelCallID, pending.ProviderPass, pending.ReasoningContent, payload.VisibleArgs, buildEditErrorResult(payload.ResolvedPath, "write read result missing"))
|
||||
}
|
||||
switch readResult.GetResult().(type) {
|
||||
case *agentv1.ReadResult_FileNotFound:
|
||||
return service.startHiddenWriteExec(stream, pending.ToolCallID, pending.ModelCallID, pending.ProviderPass, pending.ReasoningContent, pending.ReasoningSignature, pending.ReasoningSignatureSource, payload, "")
|
||||
default:
|
||||
content, ok := extractReadContentForEdit(readResult)
|
||||
if !ok {
|
||||
return service.finishWriteOperation(stream, pending.ToolCallID, pending.ModelCallID, pending.ProviderPass, pending.ReasoningContent, payload.VisibleArgs, buildEditResultFromReadResult(payload.ResolvedPath, readResult))
|
||||
}
|
||||
return service.startHiddenWriteExec(stream, pending.ToolCallID, pending.ModelCallID, pending.ProviderPass, pending.ReasoningContent, pending.ReasoningSignature, pending.ReasoningSignatureSource, payload, content)
|
||||
}
|
||||
case writeWriteExecKind:
|
||||
markExecCompleted(stream, pending)
|
||||
writeResult := message.GetWriteResult()
|
||||
if writeResult == nil {
|
||||
return service.finishWriteOperation(stream, pending.ToolCallID, pending.ModelCallID, pending.ProviderPass, pending.ReasoningContent, payload.VisibleArgs, buildEditErrorResult(payload.ResolvedPath, "write result missing"))
|
||||
}
|
||||
switch item := writeResult.GetResult().(type) {
|
||||
case *agentv1.WriteResult_Success:
|
||||
payload.ResolvedPath = firstNonEmpty(strings.TrimSpace(item.Success.GetPath()), strings.TrimSpace(payload.ResolvedPath), strings.TrimSpace(payload.VisibleArgs.Path))
|
||||
if item.Success.GetFileContentAfterWrite() != "" {
|
||||
payload.AfterContent, _ = reconcilePostWriteObservedContent(payload.AfterContent, item.Success.GetFileContentAfterWrite())
|
||||
}
|
||||
payload.VisibleArgs.Path = payload.ResolvedPath
|
||||
return service.startHiddenWritePostRead(stream, pending.ToolCallID, pending.ModelCallID, pending.ProviderPass, pending.ReasoningContent, pending.ReasoningSignature, pending.ReasoningSignatureSource, payload)
|
||||
default:
|
||||
return service.finishWriteOperation(stream, pending.ToolCallID, pending.ModelCallID, pending.ProviderPass, pending.ReasoningContent, payload.VisibleArgs, buildEditResultFromWriteResult(payload.ResolvedPath, writeResult))
|
||||
}
|
||||
case writePostReadExecKind:
|
||||
markExecCompleted(stream, pending)
|
||||
writeArgs := payload.VisibleArgs
|
||||
writeArgs.Path = firstNonEmpty(strings.TrimSpace(payload.ResolvedPath), strings.TrimSpace(writeArgs.Path))
|
||||
readResult := message.GetReadResult()
|
||||
if readResult == nil {
|
||||
return service.finishWriteOperation(stream, pending.ToolCallID, pending.ModelCallID, pending.ProviderPass, pending.ReasoningContent, writeArgs, buildSuccessfulWriteResult(writeArgs.Path, payload.BeforeContent, payload.AfterContent))
|
||||
}
|
||||
if content, ok := extractReadContentForEdit(readResult); ok {
|
||||
if success := readResult.GetSuccess(); success != nil {
|
||||
writeArgs.Path = firstNonEmpty(strings.TrimSpace(success.GetPath()), writeArgs.Path)
|
||||
}
|
||||
finalAfterContent, _ := reconcilePostWriteObservedContent(payload.AfterContent, content)
|
||||
return service.finishWriteOperation(stream, pending.ToolCallID, pending.ModelCallID, pending.ProviderPass, pending.ReasoningContent, writeArgs, buildSuccessfulWriteResult(writeArgs.Path, payload.BeforeContent, finalAfterContent))
|
||||
}
|
||||
return service.finishWriteOperation(stream, pending.ToolCallID, pending.ModelCallID, pending.ProviderPass, pending.ReasoningContent, writeArgs, buildSuccessfulWriteResult(writeArgs.Path, payload.BeforeContent, payload.AfterContent))
|
||||
default:
|
||||
return fmt.Errorf("unsupported hidden write exec kind: %s", pending.ExecKind)
|
||||
}
|
||||
}
|
||||
|
||||
func (service *Service) handleHiddenWriteExecControl(stream *ActiveStream, pending runtimecore.PendingExec, message *agentv1.ExecClientControlMessage) error {
|
||||
if stream == nil || message == nil {
|
||||
return fmt.Errorf("write exec control requires stream and message")
|
||||
}
|
||||
if _, ok := message.GetMessage().(*agentv1.ExecClientControlMessage_Heartbeat); ok {
|
||||
return nil
|
||||
}
|
||||
if _, ok := message.GetMessage().(*agentv1.ExecClientControlMessage_StreamClose); ok {
|
||||
return nil
|
||||
}
|
||||
payload, err := decodePendingWritePayload(pending.ArgsJSON)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
markExecCompleted(stream, pending)
|
||||
if strings.TrimSpace(pending.ExecKind) == writePostReadExecKind {
|
||||
return service.finishWriteOperation(stream, pending.ToolCallID, pending.ModelCallID, pending.ProviderPass, pending.ReasoningContent, payload.VisibleArgs, buildSuccessfulWriteResult(payload.ResolvedPath, payload.BeforeContent, payload.AfterContent))
|
||||
}
|
||||
return service.finishWriteOperation(stream, pending.ToolCallID, pending.ModelCallID, pending.ProviderPass, pending.ReasoningContent, payload.VisibleArgs, buildEditErrorResult(payload.ResolvedPath, hiddenWriteControlError(message)))
|
||||
}
|
||||
|
||||
func (service *Service) finishWriteOperation(stream *ActiveStream, toolCallID string, modelCallID string, providerPass int, reasoningContent string, writeArgs writeOperationArgs, result *agentv1.EditResult) error {
|
||||
if stream == nil {
|
||||
return nil
|
||||
}
|
||||
if result == nil {
|
||||
result = buildEditErrorResult(writeArgs.Path, "write result missing")
|
||||
}
|
||||
writeArgs.Path = firstNonEmpty(resultPath(result), writeArgs.Path)
|
||||
historyArgsJSON, err := writeArgs.MarshalJSON()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
toolCall := buildCompletedWriteToolCall(writeArgs.Path, writeArgs.Contents, result)
|
||||
historyToolCall := buildCompletedWriteHistoryToolCall(writeArgs.Path, result)
|
||||
if err := service.appendToolResult(stream, strings.TrimSpace(toolCallID), "Write", historyArgsJSON, summarizeWriteHistoryResult(writeArgs.Path, result), reasoningContent, historyToolCall); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := service.publishToolCallCompleted(stream.RequestID, strings.TrimSpace(toolCallID), strings.TrimSpace(modelCallID), toolCall); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := service.syncSummaryCarryForward(stream.ConversationID, stream.RequestID, strings.TrimSpace(modelCallID)); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := service.publishCheckpoint(stream.RequestID, stream.ConversationID); err != nil {
|
||||
return err
|
||||
}
|
||||
return service.reconcileStream(stream)
|
||||
}
|
||||
|
||||
func decodeWriteOperationArgs(raw []byte) (writeOperationArgs, error) {
|
||||
var args map[string]any
|
||||
if err := json.Unmarshal(raw, &args); err != nil {
|
||||
return writeOperationArgs{}, fmt.Errorf("decode write args failed: %w", err)
|
||||
}
|
||||
result := writeOperationArgs{
|
||||
Path: firstNonEmpty(readStringAny(args["path"]), readStringAny(args["file_path"])),
|
||||
Contents: firstNonEmpty(readStringAny(args["contents"]), readStringAny(args["content"]), readStringAny(args["stream_content"]), readStringAny(args["streamContent"])),
|
||||
}
|
||||
if strings.TrimSpace(result.Path) == "" {
|
||||
return writeOperationArgs{}, fmt.Errorf("write path is required")
|
||||
}
|
||||
if !isAbsoluteToolPath(result.Path) {
|
||||
return writeOperationArgs{}, fmt.Errorf("write path must be absolute")
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func (args writeOperationArgs) MarshalJSON() ([]byte, error) {
|
||||
return json.Marshal(map[string]any{
|
||||
"path": strings.TrimSpace(args.Path),
|
||||
"contents": args.Contents,
|
||||
})
|
||||
}
|
||||
|
||||
func (payload pendingWritePayload) MarshalJSON() ([]byte, error) {
|
||||
type alias pendingWritePayload
|
||||
return json.Marshal(alias(payload))
|
||||
}
|
||||
|
||||
func decodePendingWritePayload(raw []byte) (pendingWritePayload, error) {
|
||||
var payload pendingWritePayload
|
||||
if err := json.Unmarshal(raw, &payload); err != nil {
|
||||
return pendingWritePayload{}, fmt.Errorf("decode write pending payload failed: %w", err)
|
||||
}
|
||||
return payload, nil
|
||||
}
|
||||
|
||||
func buildSuccessfulWriteResult(path string, beforeContent string, afterContent string) *agentv1.EditResult {
|
||||
diffString, linesAdded, linesRemoved := computeEditDiff(beforeContent, afterContent)
|
||||
return buildSuccessfulEditResult(path, beforeContent, afterContent, diffString, linesAdded, linesRemoved, "")
|
||||
}
|
||||
|
||||
func buildCompletedWriteToolCall(path string, streamContent string, result *agentv1.EditResult) *agentv1.ToolCall {
|
||||
return &agentv1.ToolCall{
|
||||
Tool: &agentv1.ToolCall_EditToolCall{
|
||||
EditToolCall: &agentv1.EditToolCall{
|
||||
Args: &agentv1.EditArgs{
|
||||
Path: strings.TrimSpace(path),
|
||||
StreamContent: literalStringPtr(streamContent),
|
||||
},
|
||||
Result: result,
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func buildStartedWriteHistoryToolCall(path string) *agentv1.ToolCall {
|
||||
return &agentv1.ToolCall{
|
||||
Tool: &agentv1.ToolCall_EditToolCall{
|
||||
EditToolCall: &agentv1.EditToolCall{
|
||||
Args: &agentv1.EditArgs{
|
||||
Path: strings.TrimSpace(path),
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func buildCompletedWriteHistoryToolCall(path string, result *agentv1.EditResult) *agentv1.ToolCall {
|
||||
return buildCompletedEditToolCall(path, compactWriteHistoryEditResult(path, result))
|
||||
}
|
||||
|
||||
func summarizeWriteHistoryResult(path string, result *agentv1.EditResult) string {
|
||||
if result == nil {
|
||||
return "write result missing"
|
||||
}
|
||||
success := result.GetSuccess()
|
||||
if success == nil {
|
||||
return summarizeEditResultWithoutPath(result)
|
||||
}
|
||||
encoded, err := json.Marshal(map[string]any{
|
||||
"success": map[string]any{
|
||||
"path": firstNonEmpty(success.GetPath(), path),
|
||||
},
|
||||
})
|
||||
if err != nil {
|
||||
return fmt.Sprintf(`{"success":{"path":%q}}`, firstNonEmpty(success.GetPath(), path))
|
||||
}
|
||||
return string(encoded)
|
||||
}
|
||||
|
||||
func compactWriteHistoryEditResult(path string, result *agentv1.EditResult) *agentv1.EditResult {
|
||||
if result == nil {
|
||||
return buildEditErrorResult("", "write result missing")
|
||||
}
|
||||
success := result.GetSuccess()
|
||||
if success == nil {
|
||||
return editResultWithoutPath(result)
|
||||
}
|
||||
return &agentv1.EditResult{
|
||||
Result: &agentv1.EditResult_Success{
|
||||
Success: &agentv1.EditSuccess{
|
||||
Path: firstNonEmpty(success.GetPath(), path),
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func boundedWriteDiffString(diffString string) string {
|
||||
return truncateProjectedReplayText("Write", diffString, projectedEditReplayLimit)
|
||||
}
|
||||
|
||||
func resultPath(result *agentv1.EditResult) string {
|
||||
if result == nil {
|
||||
return ""
|
||||
}
|
||||
switch item := result.GetResult().(type) {
|
||||
case *agentv1.EditResult_Success:
|
||||
return strings.TrimSpace(item.Success.GetPath())
|
||||
case *agentv1.EditResult_FileNotFound:
|
||||
return strings.TrimSpace(item.FileNotFound.GetPath())
|
||||
case *agentv1.EditResult_ReadPermissionDenied:
|
||||
return strings.TrimSpace(item.ReadPermissionDenied.GetPath())
|
||||
case *agentv1.EditResult_WritePermissionDenied:
|
||||
return strings.TrimSpace(item.WritePermissionDenied.GetPath())
|
||||
case *agentv1.EditResult_Rejected:
|
||||
return strings.TrimSpace(item.Rejected.GetPath())
|
||||
case *agentv1.EditResult_Error:
|
||||
return strings.TrimSpace(item.Error.GetPath())
|
||||
default:
|
||||
return ""
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user