From 2c650b37659880c8f18f551bee946a16d76d820d Mon Sep 17 00:00:00 2001 From: warelik <54947489+warelik@users.noreply.github.com> Date: Sat, 1 Aug 2026 10:58:19 +0300 Subject: [PATCH 01/13] fix(backend): cover cursor-agent CLI local-mode endpoint surface Desktop Cursor works through the local proxy because it drives the agent over BidiAppend/RunSSE plus the already-mocked unary endpoints. The cursor-agent CLI speaks the same agent protocol but calls additional unary endpoints that had no local handlers, so every request fell into a wildcard route and came back as HTTP 404, which the Connect client maps to '[unimplemented] HTTP 404'. Three independent breaks, one visible symptom: 1. Startup: ServerConfigService/GetServerConfig (only the AiService variant was mocked), DashboardService/GetTeamAdminSettingsOrEmptyIfNotInTeam and DashboardService/ListMarketplaces were missing, so the CLI aborted during session init. 2. Git workspaces: the CLI resolves the repo path-encryption key from indexingConfig.default{User,Team}PathEncryptionKey in GetServerConfig when no IDE-stored repo keys exist. The mock returned no indexingConfig, so repository identity init failed with 'No encryption key found'. 3. Tool execution: every fs tool executor (Ls/Grep/Glob/Shell) consults the ignore service, which calls getRepoBlockExcludeGlobs() -> DashboardService/GetTeamReposOrEmptyIfNotInTeam. The 404 propagated as the tool result error, so model answers arrived but every tool call returned '[unimplemented] HTTP 404'. Non-git workspaces skip the repo-block path, which is why tools only failed inside git repos. Also close the remaining tolerated 404 noise so a CLI session produces zero unimplemented responses: model listing (GetUsableModels, GetDefaultModelForCli, GetDefaultModel), dashboard/plugin housekeeping (GetGlobalCommands, GetEffectiveUserPlugins, RegisterMarketplaceAndPlugins, GetCliDownloadUrl) and telemetry (AnalyticsService/SubmitLogs, AnalyticsService/TrackEvents, OTLP /v1/traces). Additional logging: PolicyMiddleware now includes the request path in the per-request log line, which is what made this diagnosable from app.log. --- internal/backend/host.go | 180 +++++++++++++++++++++ internal/backend/server/middleware.go | 6 +- internal/backend/server/upstream/action.go | 24 +++ internal/backend/server/upstream/client.go | 24 +++ internal/backend/server/upstream/mocks.go | 51 ++++++ 5 files changed, 284 insertions(+), 1 deletion(-) diff --git a/internal/backend/host.go b/internal/backend/host.go index 945d6f2..29587d0 100644 --- a/internal/backend/host.go +++ b/internal/backend/host.go @@ -330,6 +330,19 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error { Name: "server_config", })), ), + server.POST("/aiserver.v1.ServerConfigService/GetServerConfig", + server.Name("server_config_service_get_server_config"), + server.ConnectUnary(), + server.Local(upstream.MockProtoAction(routeDeps, upstream.CompatRouteConfig{ + Name: "server_config_service_get_server_config", + StatusCode: http.StatusOK, + MockProtoType: "aiserver.v1.GetServerConfigResponse", + MockBuilder: upstream.ServerConfigMockBuilder, + })), + server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{ + Name: "server_config_service_get_server_config", + })), + ), server.POST("/aiserver.v1.AiService/AvailableModels", server.Name("available_models"), server.ConnectUnary(), @@ -343,6 +356,45 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error { Name: "available_models", })), ), + server.POST("/aiserver.v1.AiService/GetUsableModels", + server.Name("usable_models"), + server.ConnectUnary(), + server.Local(upstream.MockProtoAction(routeDeps, upstream.CompatRouteConfig{ + Name: "usable_models", + StatusCode: http.StatusOK, + MockProtoType: "aiserver.v1.GetUsableModelsResponse", + MockBuilder: upstream.UsableModelsMockBuilder, + })), + server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{ + Name: "usable_models", + })), + ), + server.POST("/aiserver.v1.AiService/GetDefaultModelForCli", + server.Name("default_model_for_cli"), + server.ConnectUnary(), + server.Local(upstream.MockProtoAction(routeDeps, upstream.CompatRouteConfig{ + Name: "default_model_for_cli", + StatusCode: http.StatusOK, + MockProtoType: "aiserver.v1.GetDefaultModelForCliResponse", + MockBuilder: upstream.DefaultModelForCliMockBuilder, + })), + server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{ + Name: "default_model_for_cli", + })), + ), + server.POST("/aiserver.v1.AiService/GetDefaultModel", + server.Name("default_model"), + server.ConnectUnary(), + server.Local(upstream.MockProtoAction(routeDeps, upstream.CompatRouteConfig{ + Name: "default_model", + StatusCode: http.StatusOK, + MockProtoType: "aiserver.v1.GetDefaultModelResponse", + MockBuilder: upstream.DefaultModelMockBuilder, + })), + server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{ + Name: "default_model", + })), + ), server.POST("/aiserver.v1.AiService/GetDefaultModelNudgeData", server.Name("default_model_nudge"), server.ConnectUnary(), @@ -382,6 +434,43 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error { Name: "first_window_statsig_decision", })), ), + server.POST("/aiserver.v1.AnalyticsService/SubmitLogs", + server.Name("analytics_submit_logs"), + server.ConnectUnary(), + server.Local(upstream.MockProtoAction(routeDeps, upstream.CompatRouteConfig{ + Name: "analytics_submit_logs", + StatusCode: http.StatusOK, + MockProtoType: "aiserver.v1.SubmitLogsResponse", + MockBuilder: upstream.SubmitLogsMockBuilder, + })), + server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{ + Name: "analytics_submit_logs", + })), + ), + server.POST("/aiserver.v1.AnalyticsService/TrackEvents", + server.Name("analytics_track_events"), + server.ConnectUnary(), + server.Local(upstream.MockProtoAction(routeDeps, upstream.CompatRouteConfig{ + Name: "analytics_track_events", + StatusCode: http.StatusOK, + MockProtoType: "aiserver.v1.TrackEventsResponse", + MockBuilder: upstream.EmptyMockBuilder, + })), + server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{ + Name: "analytics_track_events", + })), + ), + server.POST("/v1/traces", + server.Name("otlp_traces"), + server.HTTP(), + server.Local(upstream.FixedStatusAction(routeDeps, upstream.CompatRouteConfig{ + Name: "otlp_traces", + StatusCode: http.StatusOK, + })), + server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{ + Name: "otlp_traces", + })), + ), server.POST("/oauth/token", server.Name("oauth_token"), server.HTTP(), @@ -525,6 +614,97 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error { Name: "dashboard_get_managed_skills", })), ), + server.POST("/aiserver.v1.DashboardService/GetTeamAdminSettingsOrEmptyIfNotInTeam", + server.Name("dashboard_get_team_admin_settings_or_empty"), + server.ConnectUnary(), + server.Local(upstream.MockProtoAction(routeDeps, upstream.CompatRouteConfig{ + Name: "dashboard_get_team_admin_settings_or_empty", + StatusCode: http.StatusOK, + MockProtoType: "aiserver.v1.GetTeamAdminSettingsResponse", + MockBuilder: upstream.EmptyMockBuilder, + })), + server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{ + Name: "dashboard_get_team_admin_settings_or_empty", + })), + ), + server.POST("/aiserver.v1.DashboardService/GetTeamReposOrEmptyIfNotInTeam", + server.Name("dashboard_get_team_repos_or_empty"), + server.ConnectUnary(), + server.Local(upstream.MockProtoAction(routeDeps, upstream.CompatRouteConfig{ + Name: "dashboard_get_team_repos_or_empty", + StatusCode: http.StatusOK, + MockProtoType: "aiserver.v1.GetTeamReposResponse", + MockBuilder: upstream.EmptyMockBuilder, + })), + server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{ + Name: "dashboard_get_team_repos_or_empty", + })), + ), + server.POST("/aiserver.v1.DashboardService/ListMarketplaces", + server.Name("dashboard_list_marketplaces"), + server.ConnectUnary(), + server.Local(upstream.MockProtoAction(routeDeps, upstream.CompatRouteConfig{ + Name: "dashboard_list_marketplaces", + StatusCode: http.StatusOK, + MockProtoType: "aiserver.v1.ListMarketplacesResponse", + MockBuilder: upstream.EmptyMockBuilder, + })), + server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{ + Name: "dashboard_list_marketplaces", + })), + ), + server.POST("/aiserver.v1.DashboardService/GetGlobalCommands", + server.Name("dashboard_get_global_commands"), + server.ConnectUnary(), + server.Local(upstream.MockProtoAction(routeDeps, upstream.CompatRouteConfig{ + Name: "dashboard_get_global_commands", + StatusCode: http.StatusOK, + MockProtoType: "aiserver.v1.GetGlobalCommandsResponse", + MockBuilder: upstream.EmptyMockBuilder, + })), + server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{ + Name: "dashboard_get_global_commands", + })), + ), + server.POST("/aiserver.v1.DashboardService/GetEffectiveUserPlugins", + server.Name("dashboard_get_effective_user_plugins"), + server.ConnectUnary(), + server.Local(upstream.MockProtoAction(routeDeps, upstream.CompatRouteConfig{ + Name: "dashboard_get_effective_user_plugins", + StatusCode: http.StatusOK, + MockProtoType: "aiserver.v1.GetEffectiveUserPluginsResponse", + MockBuilder: upstream.EmptyMockBuilder, + })), + server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{ + Name: "dashboard_get_effective_user_plugins", + })), + ), + server.POST("/aiserver.v1.DashboardService/RegisterMarketplaceAndPlugins", + server.Name("dashboard_register_marketplace_and_plugins"), + server.ConnectUnary(), + server.Local(upstream.MockProtoAction(routeDeps, upstream.CompatRouteConfig{ + Name: "dashboard_register_marketplace_and_plugins", + StatusCode: http.StatusOK, + MockProtoType: "aiserver.v1.RegisterMarketplaceAndPluginsResponse", + MockBuilder: upstream.EmptyMockBuilder, + })), + server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{ + Name: "dashboard_register_marketplace_and_plugins", + })), + ), + server.POST("/aiserver.v1.DashboardService/GetCliDownloadUrl", + server.Name("dashboard_get_cli_download_url"), + server.ConnectUnary(), + server.Local(upstream.MockProtoAction(routeDeps, upstream.CompatRouteConfig{ + Name: "dashboard_get_cli_download_url", + StatusCode: http.StatusOK, + MockProtoType: "aiserver.v1.GetCliDownloadUrlResponse", + MockBuilder: upstream.EmptyMockBuilder, + })), + server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{ + Name: "dashboard_get_cli_download_url", + })), + ), server.POST("/aiserver.v1.DashboardService/GetMe", server.Name("dashboard_get_me"), server.ConnectUnary(), diff --git a/internal/backend/server/middleware.go b/internal/backend/server/middleware.go index 2ab33c1..bdb3d07 100644 --- a/internal/backend/server/middleware.go +++ b/internal/backend/server/middleware.go @@ -43,7 +43,11 @@ func PolicyMiddleware(configs *serverconfig.Manager) Middleware { return func(next HandlerFunc) HandlerFunc { return func(ctx *Context) error { ctx.Mode = parseExecutionMode(configs.RouteMode(ctx.UpstreamURL != nil)) - logger.Infof("ctx.Mode=%s upstream=%t", ctx.Mode, ctx.UpstreamURL != nil) + path := "" + if ctx.Request != nil && ctx.Request.URL != nil { + path = ctx.Request.URL.Path + } + logger.Infof("ctx.Mode=%s upstream=%t path=%s", ctx.Mode, ctx.UpstreamURL != nil, path) return next(ctx) } } diff --git a/internal/backend/server/upstream/action.go b/internal/backend/server/upstream/action.go index c0ecff3..500234a 100644 --- a/internal/backend/server/upstream/action.go +++ b/internal/backend/server/upstream/action.go @@ -165,6 +165,18 @@ func DefaultModelNudgeMockBuilder(reqCtx *RequestContext) (map[string]any, error return buildDefaultModelNudgeDataPayload(reqCtx) } +func UsableModelsMockBuilder(reqCtx *RequestContext) (map[string]any, error) { + return buildUsableModelsPayload(reqCtx) +} + +func DefaultModelForCliMockBuilder(reqCtx *RequestContext) (map[string]any, error) { + return buildDefaultModelForCliPayload(reqCtx) +} + +func DefaultModelMockBuilder(reqCtx *RequestContext) (map[string]any, error) { + return buildDefaultModelPayload(reqCtx) +} + func BootstrapStatsigMockBuilder(reqCtx *RequestContext) (map[string]any, error) { return buildBootstrapStatsigPayload(reqCtx) } @@ -185,6 +197,18 @@ func DashboardManagedSkillsMockBuilder(reqCtx *RequestContext) (map[string]any, return buildDashboardManagedSkillsPayload(reqCtx) } +// EmptyMockBuilder возвращает пустой proto-ответ для ручек, где клиенту +// достаточно успешного "пусто": нет team-настроек, нет репозиториев, +// нет маркетплейсов/плагинов/команд, телеметрия принята без обработки. +func EmptyMockBuilder(reqCtx *RequestContext) (map[string]any, error) { + return map[string]any{}, nil +} + +// SubmitLogsMockBuilder подтверждает приём логов телеметрии без обработки. +func SubmitLogsMockBuilder(reqCtx *RequestContext) (map[string]any, error) { + return map[string]any{"success": true}, nil +} + func DashboardGetMeMockBuilder(reqCtx *RequestContext) (map[string]any, error) { return buildDashboardGetMePayload(reqCtx) } diff --git a/internal/backend/server/upstream/client.go b/internal/backend/server/upstream/client.go index 218c234..ba2c87e 100644 --- a/internal/backend/server/upstream/client.go +++ b/internal/backend/server/upstream/client.go @@ -427,6 +427,30 @@ func newProtoMessage(typeName string) (proto.Message, error) { return &aiserverv1.GetUsageLimitStatusAndActiveGrantsResponse{}, nil case "aiserver.v1.IsOnNewPricingResponse": return &aiserverv1.IsOnNewPricingResponse{}, nil + case "aiserver.v1.GetTeamAdminSettingsResponse": + return &aiserverv1.GetTeamAdminSettingsResponse{}, nil + case "aiserver.v1.GetTeamReposResponse": + return &aiserverv1.GetTeamReposResponse{}, nil + case "aiserver.v1.ListMarketplacesResponse": + return &aiserverv1.ListMarketplacesResponse{}, nil + case "aiserver.v1.GetUsableModelsResponse": + return &aiserverv1.GetUsableModelsResponse{}, nil + case "aiserver.v1.GetDefaultModelForCliResponse": + return &aiserverv1.GetDefaultModelForCliResponse{}, nil + case "aiserver.v1.GetDefaultModelResponse": + return &aiserverv1.GetDefaultModelResponse{}, nil + case "aiserver.v1.GetGlobalCommandsResponse": + return &aiserverv1.GetGlobalCommandsResponse{}, nil + case "aiserver.v1.GetEffectiveUserPluginsResponse": + return &aiserverv1.GetEffectiveUserPluginsResponse{}, nil + case "aiserver.v1.RegisterMarketplaceAndPluginsResponse": + return &aiserverv1.RegisterMarketplaceAndPluginsResponse{}, nil + case "aiserver.v1.GetCliDownloadUrlResponse": + return &aiserverv1.GetCliDownloadUrlResponse{}, nil + case "aiserver.v1.SubmitLogsResponse": + return &aiserverv1.SubmitLogsResponse{}, nil + case "aiserver.v1.TrackEventsResponse": + return &aiserverv1.TrackEventsResponse{}, nil default: return nil, fmt.Errorf("unsupported proto message type %q", typeName) } diff --git a/internal/backend/server/upstream/mocks.go b/internal/backend/server/upstream/mocks.go index c33dcfc..c299104 100644 --- a/internal/backend/server/upstream/mocks.go +++ b/internal/backend/server/upstream/mocks.go @@ -17,6 +17,13 @@ const ( modelRuntimeThinkingEffortParameterID = "thinking_effort" + // localPathEncryptionKey — стабильный ключ шифрования путей для индексации + // репозитория, который cursor-agent CLI запрашивает через GetServerConfig + // (indexingConfig.default{User,Team}PathEncryptionKey). Без него CLI + // не может инициализировать repo identity в git-воркспейсе и вызовы + // файловых инструментов падают с "[unimplemented] HTTP 404". + localPathEncryptionKey = "6f6e63652d6c6f63616c2d706174682d656e6372797074696f6e2d6b6579" + localUltraMembershipType = "ultra" localUltraPaymentID = "local_ultra" localUltraSubscriptionStatus = "active" @@ -427,6 +434,10 @@ func buildServerConfigPayload(*RequestContext) (map[string]any, error) { "configVersion": "local_cli_sandbox_defaults_disabled_v2", // "http2Config": "HTTP2_CONFIG_FORCE_ALL_DISABLED", "cliSandboxDefaultEnabled": true, + "indexingConfig": map[string]any{ + "defaultUserPathEncryptionKey": localPathEncryptionKey, + "defaultTeamPathEncryptionKey": localPathEncryptionKey, + }, }, nil } @@ -486,6 +497,36 @@ func buildDefaultModelNudgeDataPayload(reqCtx *RequestContext) (map[string]any, }, nil } +func buildUsableModelsPayload(reqCtx *RequestContext) (map[string]any, error) { + adapters, err := loadConfiguredModelAdapters(reqCtx) + if err != nil { + return nil, err + } + refs := collectModelAdapterRefs(adapters) + models := make([]map[string]any, 0, len(refs)) + for _, ref := range refs { + models = append(models, map[string]any{"modelName": ref}) + } + return map[string]any{"models": models}, nil +} + +func buildDefaultModelForCliPayload(reqCtx *RequestContext) (map[string]any, error) { + adapters, err := loadConfiguredModelAdapters(reqCtx) + if err != nil { + return nil, err + } + return map[string]any{"model": map[string]any{"modelName": firstModelAdapterRef(adapters)}}, nil +} + +func buildDefaultModelPayload(reqCtx *RequestContext) (map[string]any, error) { + adapters, err := loadConfiguredModelAdapters(reqCtx) + if err != nil { + return nil, err + } + defaultModel := firstModelAdapterRef(adapters) + return map[string]any{"model": defaultModel, "thinkingModel": defaultModel}, nil +} + func buildBootstrapStatsigPayload(reqCtx *RequestContext) (map[string]any, error) { generatedAtMs := uint64(time.Now().UnixMilli()) authID := resolveBootstrapStatsigAuthID(reqCtx) @@ -823,6 +864,16 @@ func collectModelAdapterRefs(adapters []legacyruntime.ModelAdapterConfig) []stri return output } +// firstModelAdapterRef возвращает канал первого адаптера или пустую строку, +// если ни один адаптер не сконфигурирован. +func firstModelAdapterRef(adapters []legacyruntime.ModelAdapterConfig) string { + refs := collectModelAdapterRefs(adapters) + if len(refs) == 0 { + return "" + } + return refs[0] +} + func resolveBootstrapStatsigAuthID(reqCtx *RequestContext) string { if reqCtx != nil { if authID := authIDFromBearer(reqCtx.Headers.Get("authorization")); authID != "" { From d1d56bd6ebc54809aefcc7047874d4a680d52192 Mon Sep 17 00:00:00 2001 From: leokun <131544788+leookun@users.noreply.github.com> Date: Mon, 3 Aug 2026 12:31:41 +0800 Subject: [PATCH 02/13] =?UTF-8?q?=E6=9B=B4=E6=96=B0=20README.md?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- README.md | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/README.md b/README.md index 16d7a2f..8ae8e8b 100644 --- a/README.md +++ b/README.md @@ -20,7 +20,7 @@ https://t.me/cursor_byok ## 路线图 [正式版路线图](https://github.com/leookun/cursor-byok/discussions/32) -[详细使用教程](https://dcne38qm5vlg.feishu.cn/wiki/JeP7wdGnziBXuikNaF5czWbrn8c) +[详细使用教程](https://docs.leokun.cn) ## 后续 From c64911e8c6a4c14eff7d6fd47a3c0b1f45e8fb1c Mon Sep 17 00:00:00 2001 From: leokun <131544788+leookun@users.noreply.github.com> Date: Mon, 3 Aug 2026 18:53:30 +0800 Subject: [PATCH 03/13] Update README.md --- README.md | 8 +++++--- 1 file changed, 5 insertions(+), 3 deletions(-) diff --git a/README.md b/README.md index 8ae8e8b..199b883 100644 --- a/README.md +++ b/README.md @@ -38,12 +38,14 @@ https://t.me/cursor_byok +## Star History + ## Star History - - - Star History Chart + + + Star History Chart From 6b7c07744886c34d6d946934d7f0f7849b0b40ec Mon Sep 17 00:00:00 2001 From: leokun <131544788+leookun@users.noreply.github.com> Date: Mon, 3 Aug 2026 18:54:04 +0800 Subject: [PATCH 04/13] Update README.md --- README.md | 2 -- 1 file changed, 2 deletions(-) diff --git a/README.md b/README.md index 199b883..16c84df 100644 --- a/README.md +++ b/README.md @@ -38,8 +38,6 @@ https://t.me/cursor_byok -## Star History - ## Star History From a2e3a0777364988487862aa5ff9fccb52426ef59 Mon Sep 17 00:00:00 2001 From: leokun <131544788+leookun@users.noreply.github.com> Date: Mon, 3 Aug 2026 18:54:56 +0800 Subject: [PATCH 05/13] Update README.md --- README.md | 10 ---------- 1 file changed, 10 deletions(-) diff --git a/README.md b/README.md index 16c84df..84f0be1 100644 --- a/README.md +++ b/README.md @@ -37,13 +37,3 @@ https://t.me/cursor_byok - -## Star History - - - - - - Star History Chart - - From 5df3e8f5120fc08842fa506be1f8c894e6d72825 Mon Sep 17 00:00:00 2001 From: leookun Date: Mon, 3 Aug 2026 22:56:17 +0800 Subject: [PATCH 06/13] feat(cursor): update model handling to use agentv1 responses and enhance CLI model details - Refactored response handling in newProtoMessage to utilize agentv1 for GetUsableModelsResponse and GetDefaultModelForCliResponse. - Added new tests for buildCLIModelDetails and encodeCLIModels to ensure correct model details and encoding. - Updated buildUsableModelsPayload and buildDefaultModelForCliPayload to leverage buildCLIModelDetails for improved model data structure. --- internal/backend/server/upstream/client.go | 5 +- internal/backend/server/upstream/mocks.go | 39 ++++++++++++---- .../backend/server/upstream/mocks_test.go | 46 +++++++++++++++++++ 3 files changed, 78 insertions(+), 12 deletions(-) diff --git a/internal/backend/server/upstream/client.go b/internal/backend/server/upstream/client.go index 47bc314..dfedcd8 100644 --- a/internal/backend/server/upstream/client.go +++ b/internal/backend/server/upstream/client.go @@ -14,6 +14,7 @@ import ( "strings" "time" + "cursor/gen/agentv1" "cursor/gen/aiserverv1" "cursor/internal/logger" "cursor/internal/netproxy" @@ -433,9 +434,9 @@ func newProtoMessage(typeName string) (proto.Message, error) { case "aiserver.v1.ListMarketplacesResponse": return &aiserverv1.ListMarketplacesResponse{}, nil case "aiserver.v1.GetUsableModelsResponse": - return &aiserverv1.GetUsableModelsResponse{}, nil + return &agentv1.GetUsableModelsResponse{}, nil case "aiserver.v1.GetDefaultModelForCliResponse": - return &aiserverv1.GetDefaultModelForCliResponse{}, nil + return &agentv1.GetDefaultModelForCliResponse{}, nil case "aiserver.v1.GetDefaultModelResponse": return &aiserverv1.GetDefaultModelResponse{}, nil case "aiserver.v1.GetGlobalCommandsResponse": diff --git a/internal/backend/server/upstream/mocks.go b/internal/backend/server/upstream/mocks.go index c299104..cc61898 100644 --- a/internal/backend/server/upstream/mocks.go +++ b/internal/backend/server/upstream/mocks.go @@ -431,8 +431,8 @@ func buildServerTimePayload(*RequestContext) (map[string]any, error) { func buildServerConfigPayload(*RequestContext) (map[string]any, error) { return map[string]any{ - "configVersion": "local_cli_sandbox_defaults_disabled_v2", - // "http2Config": "HTTP2_CONFIG_FORCE_ALL_DISABLED", + "configVersion": "local_cli_sandbox_defaults_disabled_v2", + "http2Config": "HTTP2_CONFIG_FORCE_ALL_DISABLED", "cliSandboxDefaultEnabled": true, "indexingConfig": map[string]any{ "defaultUserPathEncryptionKey": localPathEncryptionKey, @@ -451,6 +451,7 @@ func buildAvailableModelsPayload(reqCtx *RequestContext) (map[string]any, error) if len(modelRefs) > 0 { defaultModel = modelRefs[0] } + modelEntries := buildAvailableModelEntries(adapters) return map[string]any{ "backgroundComposerModelConfig": map[string]any{ "bestOfNDefaultModels": append([]string(nil), modelRefs...), @@ -470,7 +471,7 @@ func buildAvailableModelsPayload(reqCtx *RequestContext) (map[string]any, error) "defaultModel": defaultModel, }, "disableUnusedModelsAfterNHours": availableModelsDisableUnusedHours, - "models": buildAvailableModelEntries(adapters), + "models": modelEntries, "planExecutionModelConfig": map[string]any{ "defaultModel": defaultModel, "fallbackModels": append([]string(nil), modelRefs...), @@ -502,12 +503,7 @@ func buildUsableModelsPayload(reqCtx *RequestContext) (map[string]any, error) { if err != nil { return nil, err } - refs := collectModelAdapterRefs(adapters) - models := make([]map[string]any, 0, len(refs)) - for _, ref := range refs { - models = append(models, map[string]any{"modelName": ref}) - } - return map[string]any{"models": models}, nil + return map[string]any{"models": buildCLIModelDetails(adapters)}, nil } func buildDefaultModelForCliPayload(reqCtx *RequestContext) (map[string]any, error) { @@ -515,7 +511,11 @@ func buildDefaultModelForCliPayload(reqCtx *RequestContext) (map[string]any, err if err != nil { return nil, err } - return map[string]any{"model": map[string]any{"modelName": firstModelAdapterRef(adapters)}}, nil + models := buildCLIModelDetails(adapters) + if len(models) == 0 { + return map[string]any{"model": map[string]any{}}, nil + } + return map[string]any{"model": models[0]}, nil } func buildDefaultModelPayload(reqCtx *RequestContext) (map[string]any, error) { @@ -720,6 +720,25 @@ func buildAvailableModelEntries(adapters []legacyruntime.ModelAdapterConfig) []m return output } +func buildCLIModelDetails(adapters []legacyruntime.ModelAdapterConfig) []map[string]any { + models := make([]map[string]any, 0, len(adapters)) + for _, adapter := range adapters { + channelID := strings.TrimSpace(adapter.ID) + if channelID == "" { + continue + } + models = append(models, map[string]any{ + "modelId": channelID, + "displayModelId": channelID, + "apiKeyCredentials": map[string]any{ + "apiKey": strings.TrimSpace(adapter.APIKey), + "baseUrl": strings.TrimSpace(adapter.BaseURL), + }, + }) + } + return models +} + func buildThinkingEffortParameterDefinitions(adapterType string) []map[string]any { values := thinkingEffortValuesForAdapter(adapterType) options := make([]map[string]any, 0, len(values)) diff --git a/internal/backend/server/upstream/mocks_test.go b/internal/backend/server/upstream/mocks_test.go index 5c878b9..48c4fa6 100644 --- a/internal/backend/server/upstream/mocks_test.go +++ b/internal/backend/server/upstream/mocks_test.go @@ -2,9 +2,55 @@ package upstream import ( "encoding/json" + "reflect" "testing" + + "cursor/gen/agentv1" + legacyruntime "cursor/internal/runtime" + + "google.golang.org/protobuf/proto" ) +func TestBuildCLIModelDetailsPreservesChannelCredentials(t *testing.T) { + adapters := []legacyruntime.ModelAdapterConfig{ + {ID: " channel-a ", ModelID: "model-a", APIKey: "provider-secret-a", BaseURL: "https://provider-a.example/v1"}, + {ID: "channel-b", ModelID: "model-a"}, + {ID: "", ModelID: "model-c"}, + } + + got := buildCLIModelDetails(adapters) + want := []map[string]any{ + {"modelId": "channel-a", "displayModelId": "channel-a", "apiKeyCredentials": map[string]any{"apiKey": "provider-secret-a", "baseUrl": "https://provider-a.example/v1"}}, + {"modelId": "channel-b", "displayModelId": "channel-b", "apiKeyCredentials": map[string]any{"apiKey": "", "baseUrl": ""}}, + } + if !reflect.DeepEqual(got, want) { + t.Fatalf("build CLI model details: got %v, want %v", got, want) + } +} + +func TestEncodeCLIModelsUsesAgentModelDetailsWireFormat(t *testing.T) { + payload := map[string]any{"models": buildCLIModelDetails([]legacyruntime.ModelAdapterConfig{{ID: "channel-a", APIKey: "provider-secret", BaseURL: "https://provider.example/v1"}})} + encoded, err := encodeMockProto("aiserver.v1.GetUsableModelsResponse", payload) + if err != nil { + t.Fatalf("encode CLI models: %v", err) + } + + response := &agentv1.GetUsableModelsResponse{} + if err := proto.Unmarshal(encoded, response); err != nil { + t.Fatalf("decode CLI models with agent proto: %v", err) + } + if len(response.Models) != 1 { + t.Fatalf("decoded model count: got %d, want 1", len(response.Models)) + } + model := response.Models[0] + if model.GetModelId() != "channel-a" || model.GetDisplayModelId() != "channel-a" { + t.Fatalf("decoded channel IDs: model=%q display=%q", model.GetModelId(), model.GetDisplayModelId()) + } + if credentials := model.GetApiKeyCredentials(); credentials == nil || credentials.GetApiKey() != "provider-secret" || credentials.GetBaseUrl() != "https://provider.example/v1" { + t.Fatalf("decoded relay credentials: %#v", credentials) + } +} + func TestBuildBootstrapStatsigConfigJSONDisablesAlwaysLocalDecompositionGate(t *testing.T) { payload, err := buildBootstrapStatsigConfigJSON(12345, "test-auth-id") if err != nil { From 0c9a07e0180de0ba1a00c9737fdd8e84693bfdef Mon Sep 17 00:00:00 2001 From: leookun Date: Mon, 3 Aug 2026 23:17:44 +0800 Subject: [PATCH 07/13] chore(release): update version to 0.0.44 and enhance release notes - Updated version number to 0.0.44 across all relevant files. - Added support for cursor-cli in release notes. --- build/config.yml | 2 +- build/darwin/Info.dev.plist | 4 ++-- build/darwin/Info.plist | 4 ++-- build/linux/nfpm/nfpm.yaml | 2 +- build/windows/info.json | 4 ++-- build/windows/nsis/wails_tools.nsh | 2 +- build/windows/wails.exe.manifest | 2 +- release-notes.md | 8 +------- 8 files changed, 11 insertions(+), 17 deletions(-) diff --git a/build/config.yml b/build/config.yml index 60eb77c..c72bd69 100644 --- a/build/config.yml +++ b/build/config.yml @@ -8,7 +8,7 @@ info: description: "Cursor助手" copyright: "© 2026, Cursor助手" comments: "Cursor助手" - version: "0.0.43" + version: "0.0.44" dev_mode: root_path: . diff --git a/build/darwin/Info.dev.plist b/build/darwin/Info.dev.plist index 1aa1487..4ca9f2e 100644 --- a/build/darwin/Info.dev.plist +++ b/build/darwin/Info.dev.plist @@ -17,9 +17,9 @@ CFBundlePackageType APPL CFBundleShortVersionString - 0.0.43 + 0.0.44 CFBundleVersion - 0.0.43 + 0.0.44 LSMinimumSystemVersion 12.0.0 LSUIElement diff --git a/build/darwin/Info.plist b/build/darwin/Info.plist index 1363dd9..3591e80 100644 --- a/build/darwin/Info.plist +++ b/build/darwin/Info.plist @@ -17,9 +17,9 @@ CFBundlePackageType APPL CFBundleShortVersionString - 0.0.43 + 0.0.44 CFBundleVersion - 0.0.43 + 0.0.44 LSMinimumSystemVersion 12.0.0 LSUIElement diff --git a/build/linux/nfpm/nfpm.yaml b/build/linux/nfpm/nfpm.yaml index 96a9d63..7389f1f 100644 --- a/build/linux/nfpm/nfpm.yaml +++ b/build/linux/nfpm/nfpm.yaml @@ -6,7 +6,7 @@ name: "Cursor助手" arch: ${GOARCH} platform: "linux" -version: "0.0.43" +version: "0.0.44" section: "default" priority: "extra" maintainer: ${GIT_COMMITTER_NAME} <${GIT_COMMITTER_EMAIL}> diff --git a/build/windows/info.json b/build/windows/info.json index 6ad5d43..d8c9a4a 100644 --- a/build/windows/info.json +++ b/build/windows/info.json @@ -1,10 +1,10 @@ { "fixed": { - "file_version": "0.0.43" + "file_version": "0.0.44" }, "info": { "0000": { - "ProductVersion": "0.0.43", + "ProductVersion": "0.0.44", "CompanyName": "Cursor助手", "FileDescription": "Cursor助手", "LegalCopyright": "© 2026, Cursor助手", diff --git a/build/windows/nsis/wails_tools.nsh b/build/windows/nsis/wails_tools.nsh index 4c944bb..f011a54 100644 --- a/build/windows/nsis/wails_tools.nsh +++ b/build/windows/nsis/wails_tools.nsh @@ -14,7 +14,7 @@ !define INFO_PRODUCTNAME "Cursor助手" !endif !ifndef INFO_PRODUCTVERSION - !define INFO_PRODUCTVERSION "0.0.43" + !define INFO_PRODUCTVERSION "0.0.44" !endif !ifndef INFO_COPYRIGHT !define INFO_COPYRIGHT "© 2026, Cursor助手" diff --git a/build/windows/wails.exe.manifest b/build/windows/wails.exe.manifest index 8b50f75..3ed917e 100644 --- a/build/windows/wails.exe.manifest +++ b/build/windows/wails.exe.manifest @@ -1,6 +1,6 @@ - + diff --git a/release-notes.md b/release-notes.md index 1c3a92a..1c11759 100644 --- a/release-notes.md +++ b/release-notes.md @@ -8,11 +8,5 @@ QQ交流群: Tg群组: https://t.me/cursor_byok -- 修复commands无法识别的问题 @DedSecer -- 修复 summarize 指令(主动压缩) @DedSecer -- 修复 MiniMax 禁止 thinking 的问题 @Octopus -- 支持Fork message (需要新开对话) @DedSecer -- 支持 @关联对话 @DedSecer -- 支持插件市场(须登录你的任意账号) @aike1202 -- 新增支持 cursor debugger 调试器, 用法请看.agents/skills/coding-guidance/SKILL.md +- 支持cursor-cli From a9b406ec306bbfa81da2738483d04f080ea658b1 Mon Sep 17 00:00:00 2001 From: Phoeky <61970859+Phoeky@users.noreply.github.com> Date: Tue, 4 Aug 2026 17:39:14 +0800 Subject: [PATCH 08/13] =?UTF-8?q?=E8=A7=A3=E5=86=B3=20VDI=20=E7=8E=AF?= =?UTF-8?q?=E5=A2=83=EF=BC=88=E5=A6=82=20VMware=20Horizon=EF=BC=89?= =?UTF-8?q?=E5=AF=BC=E8=87=B4=E4=B8=BB=E7=AA=97=E5=8F=A3=E7=99=BD=E5=B1=8F?= =?UTF-8?q?=E3=80=82?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- internal/app/runner.go | 5 +++++ 1 file changed, 5 insertions(+) diff --git a/internal/app/runner.go b/internal/app/runner.go index d4ee32f..2ab3174 100644 --- a/internal/app/runner.go +++ b/internal/app/runner.go @@ -134,6 +134,11 @@ func Run(resources EmbeddedResources) error { Assets: application.AssetOptions{ Handler: application.AssetFileServerFS(resources.Assets), }, + Windows: application.WindowsOptions{ + // VDI 环境(如 VMware Horizon)会阻止 WebView2 创建沙箱受限令牌的renderer 进程,导致主窗口白屏。 + // --no-sandbox 跳过受限沙箱创建,使渲染进程可以正常启动。 + AdditionalBrowserArgs: []string{"--no-sandbox"}, + }, Mac: application.MacOptions{ ActivationPolicy: application.ActivationPolicyAccessory, ApplicationShouldTerminateAfterLastWindowClosed: false, From f5a7347f44f25e1291710e8e8d4dc5aeb17b7c05 Mon Sep 17 00:00:00 2001 From: leookun Date: Tue, 4 Aug 2026 21:49:26 +0800 Subject: [PATCH 09/13] fix: gate WebView sandbox opt-out behind environment variable --- internal/app/runner.go | 16 +++++++++++++--- internal/app/runner_test.go | 29 +++++++++++++++++++++++++++++ 2 files changed, 42 insertions(+), 3 deletions(-) create mode 100644 internal/app/runner_test.go diff --git a/internal/app/runner.go b/internal/app/runner.go index 2ab3174..6ab7666 100644 --- a/internal/app/runner.go +++ b/internal/app/runner.go @@ -8,7 +8,9 @@ import ( "encoding/pem" "io/fs" "net" + "os" goruntime "runtime" + "strconv" "strings" "time" @@ -37,6 +39,8 @@ const ( appName = "Cursor助手" // adRefreshInterval 表示后台广告拉取间隔。 adRefreshInterval = 3 * time.Minute + // disableWebViewSandboxEnv allows affected VDI users to opt out of the WebView2 sandbox. + disableWebViewSandboxEnv = "CURSOR_BYOK_DISABLE_WEBVIEW_SANDBOX" ) // EmbeddedResources 定义了当前模块中的 EmbeddedResources 类型。 @@ -135,9 +139,7 @@ func Run(resources EmbeddedResources) error { Handler: application.AssetFileServerFS(resources.Assets), }, Windows: application.WindowsOptions{ - // VDI 环境(如 VMware Horizon)会阻止 WebView2 创建沙箱受限令牌的renderer 进程,导致主窗口白屏。 - // --no-sandbox 跳过受限沙箱创建,使渲染进程可以正常启动。 - AdditionalBrowserArgs: []string{"--no-sandbox"}, + AdditionalBrowserArgs: windowsAdditionalBrowserArgs(), }, Mac: application.MacOptions{ ActivationPolicy: application.ActivationPolicyAccessory, @@ -420,6 +422,14 @@ func Run(resources EmbeddedResources) error { return app.Run() } +func windowsAdditionalBrowserArgs() []string { + disableSandbox, err := strconv.ParseBool(strings.TrimSpace(os.Getenv(disableWebViewSandboxEnv))) + if err != nil || !disableSandbox { + return nil + } + return []string{"--no-sandbox"} +} + func browserReachableLoopbackBaseURL(listenAddr string) string { host, port, err := net.SplitHostPort(strings.TrimSpace(listenAddr)) if err != nil || strings.TrimSpace(port) == "" { diff --git a/internal/app/runner_test.go b/internal/app/runner_test.go new file mode 100644 index 0000000..1010b38 --- /dev/null +++ b/internal/app/runner_test.go @@ -0,0 +1,29 @@ +package app + +import ( + "reflect" + "testing" +) + +func TestWindowsAdditionalBrowserArgs(t *testing.T) { + tests := []struct { + name string + env string + want []string + }{ + {name: "unset"}, + {name: "enabled with one", env: "1", want: []string{"--no-sandbox"}}, + {name: "enabled with true", env: " true ", want: []string{"--no-sandbox"}}, + {name: "disabled", env: "false"}, + {name: "invalid", env: "yes"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Setenv(disableWebViewSandboxEnv, tt.env) + if got := windowsAdditionalBrowserArgs(); !reflect.DeepEqual(got, tt.want) { + t.Fatalf("windowsAdditionalBrowserArgs() = %v, want %v", got, tt.want) + } + }) + } +} From 622a57fed7cc70045dcab6b88e187b59008cb427 Mon Sep 17 00:00:00 2001 From: leookun Date: Tue, 4 Aug 2026 23:49:49 +0800 Subject: [PATCH 10/13] Revert "Merge pull request #239 from DedSecer/fix/cursor-fork-context" This reverts commit 697fa99245b77d24c9fa1aaf2612f9a8cb98d44f, reversing changes made to 5742073be5d14e19042a0e1a5a1b09e4c9c596a9. --- internal/backend/forwarder/actor.go | 7 - internal/backend/forwarder/blob_sync.go | 383 ------------ internal/backend/forwarder/blob_sync_test.go | 546 ------------------ internal/backend/forwarder/broker.go | 46 +- internal/backend/forwarder/events.go | 16 - internal/backend/forwarder/file_store.go | 2 - internal/backend/forwarder/imported_blobs.go | 167 ------ internal/backend/forwarder/projector.go | 459 +++++++-------- internal/backend/forwarder/projector_test.go | 433 -------------- internal/backend/forwarder/rewind.go | 25 +- internal/backend/forwarder/runtime_summary.go | 3 +- internal/backend/forwarder/service.go | 108 +--- internal/backend/forwarder/token_usage.go | 55 +- internal/backend/forwarder/types.go | 33 -- 14 files changed, 304 insertions(+), 1979 deletions(-) delete mode 100644 internal/backend/forwarder/blob_sync.go delete mode 100644 internal/backend/forwarder/blob_sync_test.go delete mode 100644 internal/backend/forwarder/imported_blobs.go delete mode 100644 internal/backend/forwarder/projector_test.go diff --git a/internal/backend/forwarder/actor.go b/internal/backend/forwarder/actor.go index a712976..2614cc1 100644 --- a/internal/backend/forwarder/actor.go +++ b/internal/backend/forwarder/actor.go @@ -23,7 +23,6 @@ const ( TurnPhaseWaitingExternal TurnPhase = "waiting_external" TurnPhaseAwaitingUser TurnPhase = "awaiting_user" TurnPhaseCompacting TurnPhase = "compacting" - TurnPhaseCheckpointing TurnPhase = "checkpointing" TurnPhaseCompleted TurnPhase = "completed" TurnPhaseFailed TurnPhase = "failed" TurnPhaseCanceled TurnPhase = "canceled" @@ -67,7 +66,6 @@ const ( streamTimerNonStreamingRecovery streamTimerKind = "non_streaming_recovery" streamTimerShellForeground streamTimerKind = "shell_foreground" streamTimerShellTransportClose streamTimerKind = "shell_transport_close" - streamTimerCheckpointBlobs streamTimerKind = "checkpoint_blobs" streamTimerOrphanCancel streamTimerKind = "orphan_cancel" ) @@ -320,9 +318,6 @@ func (service *Service) handleStreamCommand(stream *ActiveStream, command stream case streamCommandCancel: return service.handleCancelIntent(command.Intent) case streamCommandMetadata: - if strings.TrimSpace(command.Intent.Kind) == "kv_result" { - return service.handleCheckpointBlobResult(stream, command.Intent.KVClientMessage) - } return service.handleMetadataIntent(command.Intent) case streamCommandExecResult: return service.handleExecResult(command.Intent) @@ -1008,8 +1003,6 @@ func (service *Service) handleTimerEvent(stream *ActiveStream, payload *streamTi return nil } return service.recoverShellWithoutTerminal(stream, current, shellRecoveryReasonTransportClosed) - case streamTimerCheckpointBlobs: - return service.handleCheckpointBlobTimeout(stream) case streamTimerOrphanCancel: stream.mu.Lock() subscriberCount := len(stream.Subscribers) diff --git a/internal/backend/forwarder/blob_sync.go b/internal/backend/forwarder/blob_sync.go deleted file mode 100644 index eb79b33..0000000 --- a/internal/backend/forwarder/blob_sync.go +++ /dev/null @@ -1,383 +0,0 @@ -package forwarder - -import ( - "encoding/hex" - "fmt" - "log" - "strings" - "time" - - "google.golang.org/protobuf/proto" - - "cursor/gen/agentv1" -) - -const ( - checkpointBlobWriteTimeout = 10 * time.Second - checkpointBlobCacheIdleTTL = 6 * time.Hour - checkpointBlobCacheMaxConversations = 256 -) - -type checkpointBlobCacheEntry struct { - Confirmed map[string]struct{} - LastAccess time.Time -} - -func checkpointBlobKey(id []byte) string { - return string(id) -} - -func checkpointBlobHex(key string) string { - return hex.EncodeToString([]byte(key)) -} - -func (service *Service) confirmedCheckpointBlob(conversationID string, key string) bool { - if service == nil || key == "" { - return false - } - service.checkpointBlobMu.Lock() - defer service.checkpointBlobMu.Unlock() - conversationID = strings.TrimSpace(conversationID) - entry := service.checkpointBlobs[conversationID] - if entry == nil { - return false - } - entry.LastAccess = time.Now().UTC() - _, ok := entry.Confirmed[key] - return ok -} - -func (service *Service) confirmCheckpointBlob(conversationID string, key string) { - if service == nil || key == "" { - return - } - service.checkpointBlobMu.Lock() - defer service.checkpointBlobMu.Unlock() - if service.checkpointBlobs == nil { - service.checkpointBlobs = make(map[string]*checkpointBlobCacheEntry) - } - conversationID = strings.TrimSpace(conversationID) - now := time.Now().UTC() - entry := service.checkpointBlobs[conversationID] - if entry == nil { - entry = &checkpointBlobCacheEntry{Confirmed: make(map[string]struct{})} - service.checkpointBlobs[conversationID] = entry - } - entry.Confirmed[key] = struct{}{} - entry.LastAccess = now - service.pruneCheckpointBlobCacheLocked(now) -} - -func (service *Service) pruneCheckpointBlobCacheLocked(now time.Time) { - if service == nil || len(service.checkpointBlobs) == 0 { - return - } - cutoff := now.Add(-checkpointBlobCacheIdleTTL) - for conversationID, entry := range service.checkpointBlobs { - if entry == nil || entry.LastAccess.Before(cutoff) { - delete(service.checkpointBlobs, conversationID) - } - } - for len(service.checkpointBlobs) > checkpointBlobCacheMaxConversations { - oldestConversationID := "" - oldestAccess := now - for conversationID, entry := range service.checkpointBlobs { - if entry == nil || oldestConversationID == "" || entry.LastAccess.Before(oldestAccess) { - oldestConversationID = conversationID - if entry != nil { - oldestAccess = entry.LastAccess - } - } - } - if oldestConversationID == "" { - return - } - delete(service.checkpointBlobs, oldestConversationID) - } -} - -func checkpointCompletionAction(completion *pendingTurnCompletion) checkpointTerminalAction { - if completion == nil { - return checkpointTerminalAction{kind: checkpointTerminalActionNone} - } - return checkpointTerminalAction{ - kind: checkpointTerminalActionComplete, - completion: *clonePendingTurnCompletion(completion), - } -} - -func checkpointCancellationAction(message string) checkpointTerminalAction { - return checkpointTerminalAction{ - kind: checkpointTerminalActionCancel, - cancelMessage: firstNonEmpty(strings.TrimSpace(message), "[canceled] User aborted request"), - } -} - -func mergeCheckpointTerminalAction(current checkpointTerminalAction, incoming checkpointTerminalAction) checkpointTerminalAction { - switch { - case incoming.kind == checkpointTerminalActionCancel: - return incoming - case current.kind == checkpointTerminalActionCancel: - return current - case incoming.kind == checkpointTerminalActionComplete: - return incoming - default: - return current - } -} - -func (action checkpointTerminalAction) completionValue() *pendingTurnCompletion { - if action.kind != checkpointTerminalActionComplete { - return nil - } - return clonePendingTurnCompletion(&action.completion) -} - -func (service *Service) queueCheckpointProjection(stream *ActiveStream, projection *CheckpointProjection, terminalAction checkpointTerminalAction) error { - if service == nil || stream == nil || projection == nil || projection.State == nil { - return nil - } - state, ok := proto.Clone(projection.State).(*agentv1.ConversationStateStructure) - if !ok || state == nil { - return fmt.Errorf("clone checkpoint state") - } - - stream.mu.Lock() - if stream.PendingCheckpointBlobWrites == nil { - stream.PendingCheckpointBlobWrites = make(map[uint32]pendingCheckpointBlobWrite) - } - if stream.PendingCheckpointBlobRequests == nil { - stream.PendingCheckpointBlobRequests = make(map[string]uint32) - } - stream.NextCheckpointRevision++ - if stream.NextCheckpointRevision == 0 { - stream.NextCheckpointRevision++ - } - revision := stream.NextCheckpointRevision - if stream.PendingCheckpoint != nil { - terminalAction = mergeCheckpointTerminalAction(stream.PendingCheckpoint.TerminalAction, terminalAction) - } - required := make(map[string]struct{}, len(projection.Blobs)) - toWrite := make([]struct { - requestID uint32 - blob CheckpointBlob - }, 0, len(projection.Blobs)) - for _, blob := range projection.Blobs { - key := checkpointBlobKey(blob.ID) - if key == "" { - continue - } - required[key] = struct{}{} - if service.confirmedCheckpointBlob(stream.ConversationID, key) { - continue - } - if _, pending := stream.PendingCheckpointBlobRequests[key]; pending { - continue - } - stream.NextCheckpointBlobRequestID++ - if stream.NextCheckpointBlobRequestID == 0 { - stream.NextCheckpointBlobRequestID++ - } - requestID := stream.NextCheckpointBlobRequestID - stream.PendingCheckpointBlobWrites[requestID] = pendingCheckpointBlobWrite{ - Key: key, - Revision: revision, - } - stream.PendingCheckpointBlobRequests[key] = requestID - toWrite = append(toWrite, struct { - requestID uint32 - blob CheckpointBlob - }{requestID: requestID, blob: blob}) - } - stream.PendingCheckpoint = &pendingCheckpointPublish{ - Revision: revision, - State: state, - Required: required, - TerminalAction: terminalAction, - } - if terminalAction.kind != checkpointTerminalActionNone { - stream.Phase = TurnPhaseCheckpointing - } - for requestID, write := range stream.PendingCheckpointBlobWrites { - if _, stillRequired := required[write.Key]; stillRequired { - continue - } - delete(stream.PendingCheckpointBlobWrites, requestID) - delete(stream.PendingCheckpointBlobRequests, write.Key) - } - stream.UpdatedAt = time.Now().UTC() - stream.mu.Unlock() - - for _, item := range toWrite { - if err := service.broker.Publish(stream.RequestID, StreamEvent{ - Message: buildSetCheckpointBlobMessage(item.requestID, item.blob), - }); err != nil { - service.discardPendingCheckpoint(stream, fmt.Errorf("publish checkpoint blob write: %w", err)) - return err - } - } - if len(toWrite) > 0 { - service.scheduleStreamTimer( - stream, - providerTimerKey(streamTimerCheckpointBlobs, ""), - checkpointBlobWriteTimeout, - streamTimerCheckpointBlobs, - "", - 0, - "checkpoint blob write timeout", - ) - } - return service.publishReadyCheckpoint(stream) -} - -func clonePendingTurnCompletion(completion *pendingTurnCompletion) *pendingTurnCompletion { - if completion == nil { - return nil - } - cloned := *completion - return &cloned -} - -func (service *Service) handleCheckpointBlobResult(stream *ActiveStream, message *agentv1.KvClientMessage) error { - if service == nil || stream == nil || message == nil { - return nil - } - result := message.GetSetBlobResult() - if result == nil { - return nil - } - stream.mu.Lock() - write, ok := stream.PendingCheckpointBlobWrites[message.GetId()] - pendingRequiresBlob := false - if ok { - delete(stream.PendingCheckpointBlobWrites, message.GetId()) - delete(stream.PendingCheckpointBlobRequests, write.Key) - if stream.PendingCheckpoint != nil { - _, pendingRequiresBlob = stream.PendingCheckpoint.Required[write.Key] - } - } - conversationID := stream.ConversationID - stream.UpdatedAt = time.Now().UTC() - stream.mu.Unlock() - if !ok { - return nil - } - if result.GetError() != nil { - if !pendingRequiresBlob { - return service.publishReadyCheckpoint(stream) - } - return service.abandonPendingCheckpoint(stream, fmt.Errorf( - "write checkpoint blob %s: %s", - checkpointBlobHex(write.Key), - firstNonEmpty(result.GetError().GetMessage(), "client blob store rejected write"), - )) - } - service.confirmCheckpointBlob(conversationID, write.Key) - return service.publishReadyCheckpoint(stream) -} - -func (service *Service) publishReadyCheckpoint(stream *ActiveStream) error { - if service == nil || stream == nil { - return nil - } - stream.mu.Lock() - pending := stream.PendingCheckpoint - if pending == nil { - stream.mu.Unlock() - return nil - } - for key := range pending.Required { - if !service.confirmedCheckpointBlob(stream.ConversationID, key) { - stream.mu.Unlock() - return nil - } - } - stream.PendingCheckpoint = nil - state := pending.State - terminalAction := pending.TerminalAction - stream.UpdatedAt = time.Now().UTC() - stream.mu.Unlock() - clearStreamTimer(stream, providerTimerKey(streamTimerCheckpointBlobs, "")) - if err := service.broker.Publish(stream.RequestID, StreamEvent{Message: buildCheckpointMessage(state)}); err != nil { - return err - } - switch terminalAction.kind { - case checkpointTerminalActionComplete: - if completion := terminalAction.completionValue(); completion != nil { - return service.finishSuccessfulTurnAfterCheckpoint(stream, *completion) - } - case checkpointTerminalActionCancel: - return service.finishCanceledTurnAfterCheckpoint(stream, terminalAction.cancelMessage) - } - return nil -} - -func (service *Service) discardPendingCheckpoint(stream *ActiveStream, cause error) { - if service == nil || stream == nil { - return - } - stream.mu.Lock() - stream.PendingCheckpoint = nil - stream.PendingCheckpointBlobWrites = make(map[uint32]pendingCheckpointBlobWrite) - stream.PendingCheckpointBlobRequests = make(map[string]uint32) - stream.UpdatedAt = time.Now().UTC() - stream.mu.Unlock() - clearStreamTimer(stream, providerTimerKey(streamTimerCheckpointBlobs, "")) - if cause != nil { - log.Printf("forwarder pending checkpoint discarded request_id=%s conversation_id=%s err=%v", stream.RequestID, stream.ConversationID, cause) - } -} - -func (service *Service) abandonPendingCheckpoint(stream *ActiveStream, cause error) error { - if service == nil || stream == nil { - return nil - } - stream.mu.Lock() - pending := stream.PendingCheckpoint - stream.PendingCheckpoint = nil - stream.PendingCheckpointBlobWrites = make(map[uint32]pendingCheckpointBlobWrite) - stream.PendingCheckpointBlobRequests = make(map[string]uint32) - stream.UpdatedAt = time.Now().UTC() - stream.mu.Unlock() - clearStreamTimer(stream, providerTimerKey(streamTimerCheckpointBlobs, "")) - if cause != nil { - log.Printf("forwarder checkpoint blob sync abandoned request_id=%s conversation_id=%s err=%v", stream.RequestID, stream.ConversationID, cause) - } - if pending != nil { - return service.failTerminalCheckpointSync(stream, cause) - } - return nil -} - -func (service *Service) finishCanceledTurnAfterCheckpoint(stream *ActiveStream, message string) error { - if stream == nil { - return nil - } - service.setTurnPhase(stream, TurnPhaseCanceled) - return service.broker.Cancel(stream.RequestID, firstNonEmpty(strings.TrimSpace(message), "[canceled] User aborted request")) -} - -func (service *Service) failTerminalCheckpointSync(stream *ActiveStream, cause error) error { - if stream == nil { - return nil - } - message := "checkpoint synchronization failed" - if cause != nil && strings.TrimSpace(cause.Error()) != "" { - message = strings.TrimSpace(cause.Error()) - } - service.setTurnPhase(stream, TurnPhaseFailed) - return service.broker.Fail(stream.RequestID, "checkpoint_sync_error", message) -} - -func (service *Service) handleCheckpointBlobTimeout(stream *ActiveStream) error { - if stream == nil { - return nil - } - stream.mu.Lock() - pendingCount := len(stream.PendingCheckpointBlobWrites) - stream.mu.Unlock() - if pendingCount == 0 { - return service.publishReadyCheckpoint(stream) - } - return service.abandonPendingCheckpoint(stream, fmt.Errorf("%d checkpoint blob writes timed out", pendingCount)) -} diff --git a/internal/backend/forwarder/blob_sync_test.go b/internal/backend/forwarder/blob_sync_test.go deleted file mode 100644 index d101a92..0000000 --- a/internal/backend/forwarder/blob_sync_test.go +++ /dev/null @@ -1,546 +0,0 @@ -package forwarder - -import ( - "os" - "path/filepath" - "testing" - - "cursor/gen/agentv1" -) - -func TestCheckpointBlobSyncPublishesCheckpointAfterAllWrites(t *testing.T) { - service, stream := testCheckpointBlobService(t) - projection, err := service.projector.ProjectCheckpointProjection(testConversation([]HistoryEntry{ - testUserMessageEntry(t, 1, "request-1", "hello"), - newAssistantTextEntry(1, "request-1", "hi", "", ""), - })) - if err != nil { - t.Fatalf("ProjectCheckpointProjection() error = %v", err) - } - - if err := service.queueCheckpointProjection(stream, projection, checkpointTerminalAction{kind: checkpointTerminalActionNone}); err != nil { - t.Fatalf("queueCheckpointProjection() error = %v", err) - } - events, err := service.broker.ReadFromCursor(stream.RequestID, 0) - if err != nil { - t.Fatalf("ReadFromCursor() error = %v", err) - } - if len(events) != len(projection.Blobs) { - t.Fatalf("events before ACK = %d, want %d blob writes", len(events), len(projection.Blobs)) - } - for _, event := range events { - if event.Message.GetKvServerMessage().GetSetBlobArgs() == nil { - t.Fatalf("event before ACK = %#v, want set_blob_args", event.Message) - } - } - - for index, event := range events { - requestID := event.Message.GetKvServerMessage().GetId() - if err := service.handleCheckpointBlobResult(stream, &agentv1.KvClientMessage{ - Id: requestID, - Message: &agentv1.KvClientMessage_SetBlobResult{ - SetBlobResult: &agentv1.SetBlobResult{}, - }, - }); err != nil { - t.Fatalf("handleCheckpointBlobResult(%d) error = %v", index, err) - } - } - events, err = service.broker.ReadFromCursor(stream.RequestID, 0) - if err != nil { - t.Fatalf("ReadFromCursor() after ACK error = %v", err) - } - if len(events) != len(projection.Blobs)+1 { - t.Fatalf("events after ACK = %d, want %d", len(events), len(projection.Blobs)+1) - } - checkpoint := events[len(events)-1].Message.GetConversationCheckpointUpdate() - if checkpoint == nil || len(checkpoint.GetTurns()) != 1 { - t.Fatalf("last event checkpoint = %#v, want one Blob-backed turn", checkpoint) - } -} - -func TestCheckpointBlobSyncRejectDoesNotPublishDanglingCheckpoint(t *testing.T) { - service, stream := testCheckpointBlobService(t) - projection, err := service.projector.ProjectCheckpointProjection(testConversation([]HistoryEntry{ - testUserMessageEntry(t, 1, "request-1", "hello"), - })) - if err != nil { - t.Fatalf("ProjectCheckpointProjection() error = %v", err) - } - if err := service.queueCheckpointProjection(stream, projection, checkpointTerminalAction{kind: checkpointTerminalActionNone}); err != nil { - t.Fatalf("queueCheckpointProjection() error = %v", err) - } - events, err := service.broker.ReadFromCursor(stream.RequestID, 0) - if err != nil || len(events) == 0 { - t.Fatalf("blob write events = %d, err = %v", len(events), err) - } - requestID := events[0].Message.GetKvServerMessage().GetId() - if err := service.handleCheckpointBlobResult(stream, &agentv1.KvClientMessage{ - Id: requestID, - Message: &agentv1.KvClientMessage_SetBlobResult{ - SetBlobResult: &agentv1.SetBlobResult{Error: &agentv1.Error{Message: "disk full"}}, - }, - }); err != nil { - t.Fatalf("handleCheckpointBlobResult() error = %v", err) - } - events, err = service.broker.ReadFromCursor(stream.RequestID, 0) - if err != nil { - t.Fatalf("ReadFromCursor() after rejection error = %v", err) - } - for _, event := range events { - if event.Message.GetConversationCheckpointUpdate() != nil { - t.Fatal("rejected Blob write published a dangling checkpoint") - } - } -} - -func TestTerminalCheckpointRejectionFailsInsteadOfCompleting(t *testing.T) { - service, stream := testCheckpointBlobService(t) - projection, err := service.projector.ProjectCheckpointProjection(testConversation([]HistoryEntry{ - testUserMessageEntry(t, 1, stream.RequestID, "hello"), - })) - if err != nil { - t.Fatalf("projection: %v", err) - } - completion := &pendingTurnCompletion{RequestID: stream.RequestID} - if err := service.queueCheckpointProjection(stream, projection, checkpointCompletionAction(completion)); err != nil { - t.Fatalf("queue checkpoint: %v", err) - } - stream.mu.Lock() - var requestID uint32 - for pendingID := range stream.PendingCheckpointBlobWrites { - requestID = pendingID - break - } - stream.mu.Unlock() - if requestID == 0 { - t.Fatal("test did not queue a Blob write") - } - if err := rejectCheckpointBlob(service, stream, requestID, "disk full"); err != nil { - t.Fatalf("reject terminal checkpoint: %v", err) - } - events, err := service.broker.ReadFromCursor(stream.RequestID, 0) - if err != nil { - t.Fatalf("ReadFromCursor() error = %v", err) - } - var failed, completed bool - for _, event := range events { - if event.End && event.TerminalErrorCode == "checkpoint_sync_error" { - failed = true - } - if event.Message.GetInteractionUpdate().GetTurnEnded() != nil { - completed = true - } - } - if !failed || completed { - t.Fatalf("terminal checkpoint events failed=%v completed=%v", failed, completed) - } -} - -func TestCheckpointBlobSyncMergesRevisionsWithoutObsoleteFailure(t *testing.T) { - service, stream := testCheckpointBlobService(t) - first, err := service.projector.ProjectCheckpointProjection(testConversation([]HistoryEntry{ - testUserMessageEntry(t, 1, "request-1", "hello"), - })) - if err != nil { - t.Fatalf("first ProjectCheckpointProjection() error = %v", err) - } - latest, err := service.projector.ProjectCheckpointProjection(testConversation([]HistoryEntry{ - testUserMessageEntry(t, 1, "request-1", "hello"), - newAssistantTextEntry(1, "request-1", "latest answer", "", ""), - })) - if err != nil { - t.Fatalf("latest ProjectCheckpointProjection() error = %v", err) - } - if err := service.queueCheckpointProjection(stream, first, checkpointTerminalAction{kind: checkpointTerminalActionNone}); err != nil { - t.Fatalf("queue first projection: %v", err) - } - firstEvents, err := service.broker.ReadFromCursor(stream.RequestID, 0) - if err != nil { - t.Fatalf("read first events: %v", err) - } - if err := service.queueCheckpointProjection(stream, latest, checkpointTerminalAction{kind: checkpointTerminalActionNone}); err != nil { - t.Fatalf("queue latest projection: %v", err) - } - - latestRequired := make(map[string]struct{}, len(latest.Blobs)) - for _, blob := range latest.Blobs { - latestRequired[string(blob.ID)] = struct{}{} - } - var obsoleteRequestID uint32 - for _, event := range firstEvents { - message := event.Message.GetKvServerMessage() - if message == nil || message.GetSetBlobArgs() == nil { - continue - } - if _, required := latestRequired[string(message.GetSetBlobArgs().GetBlobId())]; !required { - obsoleteRequestID = message.GetId() - break - } - } - if obsoleteRequestID == 0 { - t.Fatal("test did not find an obsolete first-revision Blob write") - } - if err := rejectCheckpointBlob(service, stream, obsoleteRequestID, "obsolete write rejected"); err != nil { - t.Fatalf("reject obsolete Blob: %v", err) - } - if err := acknowledgePendingCheckpointBlobs(service, stream); err != nil { - t.Fatalf("acknowledge latest Blob writes: %v", err) - } - - events, err := service.broker.ReadFromCursor(stream.RequestID, 0) - if err != nil { - t.Fatalf("read merged events: %v", err) - } - checkpoints := 0 - for _, event := range events { - if checkpoint := event.Message.GetConversationCheckpointUpdate(); checkpoint != nil { - checkpoints++ - if len(checkpoint.GetTurns()) != len(latest.State.GetTurns()) || string(checkpoint.GetTurns()[0]) != string(latest.State.GetTurns()[0]) { - t.Fatalf("published checkpoint is not the latest revision: %#v", checkpoint) - } - } - } - if checkpoints != 1 { - t.Fatalf("published checkpoints = %d, want exactly latest revision", checkpoints) - } -} - -func TestCheckpointBlobSyncCarriesCompletionIntoLatestRevision(t *testing.T) { - service, stream := testCheckpointBlobService(t) - first, err := service.projector.ProjectCheckpointProjection(testConversation([]HistoryEntry{ - testUserMessageEntry(t, 1, "request-1", "hello"), - })) - if err != nil { - t.Fatalf("first projection: %v", err) - } - latest, err := service.projector.ProjectCheckpointProjection(testConversation([]HistoryEntry{ - testUserMessageEntry(t, 1, "request-1", "hello"), - newAssistantTextEntry(1, "request-1", "done", "", ""), - })) - if err != nil { - t.Fatalf("latest projection: %v", err) - } - completion := &pendingTurnCompletion{ - RequestID: stream.RequestID, - Usage: turnUsageSnapshot{InputTokens: 11, OutputTokens: 7}, - } - if err := service.queueCheckpointProjection(stream, first, checkpointCompletionAction(completion)); err != nil { - t.Fatalf("queue completion projection: %v", err) - } - if err := service.queueCheckpointProjection(stream, latest, checkpointTerminalAction{kind: checkpointTerminalActionNone}); err != nil { - t.Fatalf("queue latest projection: %v", err) - } - if err := acknowledgePendingCheckpointBlobs(service, stream); err != nil { - t.Fatalf("acknowledge latest Blob writes: %v", err) - } - - events, err := service.broker.ReadFromCursor(stream.RequestID, 0) - if err != nil { - t.Fatalf("read completion events: %v", err) - } - checkpointIndex, turnEndedIndex, endIndex := -1, -1, -1 - for index, event := range events { - switch { - case event.Message.GetConversationCheckpointUpdate() != nil: - checkpointIndex = index - case event.Message.GetInteractionUpdate().GetTurnEnded() != nil: - turnEndedIndex = index - case event.End: - endIndex = index - } - } - if checkpointIndex < 0 || turnEndedIndex <= checkpointIndex || endIndex <= turnEndedIndex { - t.Fatalf("terminal order checkpoint=%d turn_ended=%d end=%d", checkpointIndex, turnEndedIndex, endIndex) - } -} - -func TestCheckpointBlobSyncTimeoutFailsStreamWithoutPublishingDanglingCheckpoint(t *testing.T) { - service, stream := testCheckpointBlobService(t) - projection, err := service.projector.ProjectCheckpointProjection(testConversation([]HistoryEntry{ - testUserMessageEntry(t, 1, "request-1", "hello"), - })) - if err != nil { - t.Fatalf("projection: %v", err) - } - if err := service.queueCheckpointProjection(stream, projection, checkpointTerminalAction{kind: checkpointTerminalActionNone}); err != nil { - t.Fatalf("queue checkpoint: %v", err) - } - if err := service.handleCheckpointBlobTimeout(stream); err != nil { - t.Fatalf("timeout checkpoint: %v", err) - } - events, err := service.broker.ReadFromCursor(stream.RequestID, 0) - if err != nil { - t.Fatalf("read timeout events: %v", err) - } - var failed bool - for _, event := range events { - if event.Message.GetConversationCheckpointUpdate() != nil { - t.Fatal("timed-out Blob dependency published a dangling checkpoint") - } - if event.End && event.TerminalErrorCode == "checkpoint_sync_error" { - failed = true - } - } - if !failed { - t.Fatal("timed-out checkpoint did not fail the stream explicitly") - } -} - -func TestCheckpointBlobSyncReusesConversationCacheAcrossRequests(t *testing.T) { - service, firstStream := testCheckpointBlobService(t) - projection, err := service.projector.ProjectCheckpointProjection(testConversation([]HistoryEntry{ - testUserMessageEntry(t, 1, "request-1", "hello"), - })) - if err != nil { - t.Fatalf("projection: %v", err) - } - if err := service.queueCheckpointProjection(firstStream, projection, checkpointTerminalAction{kind: checkpointTerminalActionNone}); err != nil { - t.Fatalf("queue first request: %v", err) - } - if err := acknowledgePendingCheckpointBlobs(service, firstStream); err != nil { - t.Fatalf("acknowledge first request: %v", err) - } - - secondStream, err := service.broker.OpenStream( - "request-2", firstStream.ConversationID, 2, "default", "default", - agentv1.AgentMode_AGENT_MODE_AGENT, "continue", - ) - if err != nil { - t.Fatalf("OpenStream() second request error = %v", err) - } - if err := service.queueCheckpointProjection(secondStream, projection, checkpointTerminalAction{kind: checkpointTerminalActionNone}); err != nil { - t.Fatalf("queue second request: %v", err) - } - events, err := service.broker.ReadFromCursor(secondStream.RequestID, 0) - if err != nil { - t.Fatalf("read second request events: %v", err) - } - if len(events) != 1 || events[0].Message.GetConversationCheckpointUpdate() == nil { - t.Fatalf("second request events = %#v, want cached immediate checkpoint", events) - } -} - -func TestCancellationReplacesUnconfirmedCheckpointBeforeEnding(t *testing.T) { - service, stream := testCheckpointBlobService(t) - conversation := testConversation([]HistoryEntry{ - testUserMessageEntry(t, 1, stream.RequestID, "hello"), - }) - if err := service.replaceCheckpointConversation(stream, conversation); err != nil { - t.Fatalf("replaceCheckpointConversation() error = %v", err) - } - projection, err := service.projector.ProjectCheckpointProjection(conversation) - if err != nil { - t.Fatalf("ProjectCheckpointProjection() error = %v", err) - } - if err := service.queueCheckpointProjection(stream, projection, checkpointTerminalAction{kind: checkpointTerminalActionNone}); err != nil { - t.Fatalf("queueCheckpointProjection() error = %v", err) - } - stream.mu.Lock() - var staleRequestID uint32 - for requestID := range stream.PendingCheckpointBlobWrites { - staleRequestID = requestID - break - } - stream.mu.Unlock() - if staleRequestID == 0 { - t.Fatal("test did not queue an unconfirmed Blob write") - } - - if err := service.handleCancelIntent(InboundIntent{ - Kind: "cancel", - RequestID: stream.RequestID, - CancelReason: "user stopped", - }); err != nil { - t.Fatalf("handleCancelIntent() error = %v", err) - } - stream.mu.Lock() - phase := stream.Phase - status := stream.Status - pendingCheckpoint := stream.PendingCheckpoint - pendingWrites := len(stream.PendingCheckpointBlobWrites) - stream.mu.Unlock() - if phase != TurnPhaseCheckpointing || status != StreamStatusCreated { - t.Fatalf("before checkpoint ACK phase=%s status=%s, want checkpointing/created", phase, status) - } - if pendingCheckpoint == nil || pendingWrites == 0 { - t.Fatalf("before checkpoint ACK pending_checkpoint=%v pending_writes=%d", pendingCheckpoint != nil, pendingWrites) - } - if err := service.handleCheckpointBlobResult(stream, &agentv1.KvClientMessage{ - Id: staleRequestID, - Message: &agentv1.KvClientMessage_SetBlobResult{ - SetBlobResult: &agentv1.SetBlobResult{}, - }, - }); err != nil { - t.Fatalf("stale Blob ACK error = %v", err) - } - if err := acknowledgePendingCheckpointBlobs(service, stream); err != nil { - t.Fatalf("acknowledge cancellation checkpoint: %v", err) - } - stream.mu.Lock() - phase = stream.Phase - status = stream.Status - stream.mu.Unlock() - if phase != TurnPhaseCanceled || status != StreamStatusCanceled { - t.Fatalf("after checkpoint ACK phase=%s status=%s, want canceled", phase, status) - } - assertCanceledEndEvent(t, service, stream) -} - -func TestCancellationMetadataFailureStillEndsStream(t *testing.T) { - service, stream := testCheckpointBlobService(t) - blockingPath := filepath.Join(t.TempDir(), "not-a-directory") - if err := os.WriteFile(blockingPath, []byte("block child creation"), 0o600); err != nil { - t.Fatalf("WriteFile() error = %v", err) - } - service.store = NewConversationFileStore(blockingPath) - if err := service.replaceCheckpointConversation(stream, testConversation([]HistoryEntry{ - testUserMessageEntry(t, 1, stream.RequestID, "hello"), - })); err != nil { - t.Fatalf("replaceCheckpointConversation() error = %v", err) - } - - if err := service.handleCancelIntent(InboundIntent{Kind: "cancel", RequestID: stream.RequestID}); err != nil { - t.Fatalf("handleCancelIntent() error = %v", err) - } - if err := acknowledgePendingCheckpointBlobs(service, stream); err != nil { - t.Fatalf("acknowledge cancellation checkpoint: %v", err) - } - assertCanceledEndEvent(t, service, stream) -} - -func TestCheckpointTerminalActionMergePriority(t *testing.T) { - complete := checkpointCompletionAction(&pendingTurnCompletion{RequestID: "complete"}) - cancel := checkpointCancellationAction("user canceled") - none := checkpointTerminalAction{kind: checkpointTerminalActionNone} - - tests := []struct { - name string - current checkpointTerminalAction - incoming checkpointTerminalAction - wantKind checkpointTerminalActionKind - wantID string - }{ - {name: "none then complete", current: none, incoming: complete, wantKind: checkpointTerminalActionComplete, wantID: "complete"}, - {name: "complete then none", current: complete, incoming: none, wantKind: checkpointTerminalActionComplete, wantID: "complete"}, - {name: "complete then cancel", current: complete, incoming: cancel, wantKind: checkpointTerminalActionCancel}, - {name: "cancel then complete", current: cancel, incoming: complete, wantKind: checkpointTerminalActionCancel}, - {name: "cancel then none", current: cancel, incoming: none, wantKind: checkpointTerminalActionCancel}, - } - for _, test := range tests { - t.Run(test.name, func(t *testing.T) { - merged := mergeCheckpointTerminalAction(test.current, test.incoming) - if merged.kind != test.wantKind { - t.Fatalf("merged kind = %d, want %d", merged.kind, test.wantKind) - } - if test.wantID != "" { - completion := merged.completionValue() - if completion == nil || completion.RequestID != test.wantID { - t.Fatalf("merged completion = %#v, want request_id=%s", completion, test.wantID) - } - } - }) - } -} - -func TestCheckpointTerminalActionIsMutuallyExclusive(t *testing.T) { - completion := &pendingTurnCompletion{RequestID: "request-1"} - action := checkpointCompletionAction(completion) - if action.kind != checkpointTerminalActionComplete || action.completionValue() == nil { - t.Fatalf("completion action = %#v", action) - } - empty := checkpointCompletionAction(nil) - if empty.kind != checkpointTerminalActionNone || empty.completionValue() != nil { - t.Fatalf("empty action = %#v", empty) - } -} - -func assertCanceledEndEvent(t *testing.T, service *Service, stream *ActiveStream) { - t.Helper() - events, err := service.broker.ReadFromCursor(stream.RequestID, 0) - if err != nil { - t.Fatalf("ReadFromCursor() error = %v", err) - } - for _, event := range events { - if event.End && event.TerminalErrorCode == "canceled" { - return - } - } - t.Fatal("cancellation did not publish canceled end event") -} - -func acknowledgePendingCheckpointBlobs(service *Service, stream *ActiveStream) error { - for { - stream.mu.Lock() - requestIDs := make([]uint32, 0, len(stream.PendingCheckpointBlobWrites)) - for requestID := range stream.PendingCheckpointBlobWrites { - requestIDs = append(requestIDs, requestID) - } - stream.mu.Unlock() - if len(requestIDs) == 0 { - return nil - } - for _, requestID := range requestIDs { - if err := service.handleCheckpointBlobResult(stream, &agentv1.KvClientMessage{ - Id: requestID, - Message: &agentv1.KvClientMessage_SetBlobResult{ - SetBlobResult: &agentv1.SetBlobResult{}, - }, - }); err != nil { - return err - } - } - } -} - -func rejectCheckpointBlob(service *Service, stream *ActiveStream, requestID uint32, message string) error { - return service.handleCheckpointBlobResult(stream, &agentv1.KvClientMessage{ - Id: requestID, - Message: &agentv1.KvClientMessage_SetBlobResult{ - SetBlobResult: &agentv1.SetBlobResult{Error: &agentv1.Error{Message: message}}, - }, - }) -} - -func TestImportedTurnIDsRemainCheckpointPrefix(t *testing.T) { - importedID := make([]byte, 32) - for index := range importedID { - importedID[index] = byte(index + 1) - } - conversation := testConversation([]HistoryEntry{ - testUserMessageEntry(t, 2, "request-2", "continued question"), - }) - conversation.ImportedTurnIDs = [][]byte{importedID} - projection, err := NewHistoryProjector().ProjectCheckpointProjection(conversation) - if err != nil { - t.Fatalf("ProjectCheckpointProjection() error = %v", err) - } - if len(projection.State.GetTurns()) != 2 { - t.Fatalf("turns = %d, want imported prefix plus projected turn", len(projection.State.GetTurns())) - } - if string(projection.State.GetTurns()[0]) != string(importedID) { - t.Fatal("imported turn ID was not preserved as the checkpoint prefix") - } -} - -func testCheckpointBlobService(t *testing.T) (*Service, *ActiveStream) { - t.Helper() - broker := NewStreamBroker() - service := &Service{ - projector: NewHistoryProjector(), - broker: broker, - checkpointBlobs: make(map[string]*checkpointBlobCacheEntry), - } - stream, err := broker.OpenStream( - "request-1", - "conversation-1", - 1, - "default", - "default", - agentv1.AgentMode_AGENT_MODE_AGENT, - "hello", - ) - if err != nil { - t.Fatalf("OpenStream() error = %v", err) - } - return service, stream -} diff --git a/internal/backend/forwarder/broker.go b/internal/backend/forwarder/broker.go index 3d704c3..b88442c 100644 --- a/internal/backend/forwarder/broker.go +++ b/internal/backend/forwarder/broker.go @@ -82,30 +82,28 @@ func (broker *StreamBroker) OpenStream(requestID string, conversationID string, } 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), - PendingCheckpointBlobWrites: make(map[uint32]pendingCheckpointBlobWrite), - PendingCheckpointBlobRequests: make(map[string]uint32), - CreatedAt: now, - UpdatedAt: now, + 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 diff --git a/internal/backend/forwarder/events.go b/internal/backend/forwarder/events.go index ef69776..7b02404 100644 --- a/internal/backend/forwarder/events.go +++ b/internal/backend/forwarder/events.go @@ -245,22 +245,6 @@ func buildCheckpointMessage(state *agentv1.ConversationStateStructure) *agentv1. } } -func buildSetCheckpointBlobMessage(id uint32, blob CheckpointBlob) *agentv1.AgentServerMessage { - return &agentv1.AgentServerMessage{ - Message: &agentv1.AgentServerMessage_KvServerMessage{ - KvServerMessage: &agentv1.KvServerMessage{ - Id: id, - Message: &agentv1.KvServerMessage_SetBlobArgs{ - SetBlobArgs: &agentv1.SetBlobArgs{ - BlobId: append([]byte(nil), blob.ID...), - BlobData: append([]byte(nil), blob.Data...), - }, - }, - }, - }, - } -} - // buildExecAbortMessage 构造对客户端执行桥的 abort 控制消息。 func buildExecAbortMessage(pending runtimecore.PendingExec) *agentv1.AgentServerMessage { return &agentv1.AgentServerMessage{ diff --git a/internal/backend/forwarder/file_store.go b/internal/backend/forwarder/file_store.go index fb96d1b..0732614 100644 --- a/internal/backend/forwarder/file_store.go +++ b/internal/backend/forwarder/file_store.go @@ -750,7 +750,6 @@ func mergeConversationMetadata(target *ConversationFile, source *ConversationFil target.CurrentPlanText = source.CurrentPlanText target.CurrentPlans = clonePlanRegistryEntries(source.CurrentPlans) target.CurrentTodos = cloneTodoItems(source.CurrentTodos) - target.ImportedTurnIDs = cloneByteSlices(source.ImportedTurnIDs) target.LatestRequestPrefix = cloneConversationRequestPrefix(source.LatestRequestPrefix) target.LastProviderCall = cloneConversationProviderCall(source.LastProviderCall) if !source.CreatedAt.IsZero() && (target.CreatedAt.IsZero() || source.CreatedAt.Before(target.CreatedAt)) { @@ -883,7 +882,6 @@ func cloneConversationFile(conversation *ConversationFile) *ConversationFile { cloned := *conversation cloned.CurrentPlans = clonePlanRegistryEntries(conversation.CurrentPlans) cloned.CurrentTodos = cloneTodoItems(conversation.CurrentTodos) - cloned.ImportedTurnIDs = cloneByteSlices(conversation.ImportedTurnIDs) cloned.LatestRequestPrefix = cloneConversationRequestPrefix(conversation.LatestRequestPrefix) cloned.LastProviderCall = cloneConversationProviderCall(conversation.LastProviderCall) cloned.Entries = append([]HistoryEntry(nil), conversation.Entries...) diff --git a/internal/backend/forwarder/imported_blobs.go b/internal/backend/forwarder/imported_blobs.go deleted file mode 100644 index fbd17ed..0000000 --- a/internal/backend/forwarder/imported_blobs.go +++ /dev/null @@ -1,167 +0,0 @@ -package forwarder - -import ( - "crypto/sha256" - "fmt" - - "google.golang.org/protobuf/proto" - - "cursor/gen/agentv1" - modeladapter "cursor/internal/backend/agent/model" - promptengine "cursor/internal/backend/agent/prompt" -) - -type importedBlobStore map[string][]byte - -func newImportedBlobStore(items []*agentv1.PreFetchedBlob) (importedBlobStore, error) { - if len(items) == 0 { - return nil, nil - } - store := make(importedBlobStore, len(items)) - for _, item := range items { - if item == nil || len(item.GetId()) == 0 { - continue - } - if len(item.GetId()) != sha256.Size { - return nil, fmt.Errorf("prefetched blob id length %d, want %d", len(item.GetId()), sha256.Size) - } - digest := sha256.Sum256(item.GetValue()) - if string(digest[:]) != string(item.GetId()) { - return nil, fmt.Errorf("prefetched blob %x failed SHA-256 validation", item.GetId()) - } - store[string(item.GetId())] = append([]byte(nil), item.GetValue()...) - } - return store, nil -} - -func (store importedBlobStore) resolve(id []byte) ([]byte, bool) { - if len(id) == 0 || len(store) == 0 { - return nil, false - } - value, ok := store[string(id)] - return append([]byte(nil), value...), ok -} - -func decodeImportedTurn(raw []byte, blobs importedBlobStore) (*agentv1.ConversationTurnStructure, []byte, error) { - if data, ok := blobs.resolve(raw); ok { - turn := &agentv1.ConversationTurnStructure{} - if err := proto.Unmarshal(data, turn); err != nil || turn.GetTurn() == nil { - return nil, nil, fmt.Errorf("decode imported turn blob %x: %w", raw, firstNonNilError(err, fmt.Errorf("turn payload is empty"))) - } - return turn, append([]byte(nil), raw...), nil - } - turn := &agentv1.ConversationTurnStructure{} - if err := proto.Unmarshal(raw, turn); err == nil && turn.GetTurn() != nil { - return turn, nil, nil - } - if len(raw) == sha256.Size { - return nil, append([]byte(nil), raw...), nil - } - return nil, nil, fmt.Errorf("decode imported inline turn") -} - -func decodeImportedUserMessage(raw []byte, blobs importedBlobStore) (*agentv1.UserMessage, error) { - data := raw - if resolved, ok := blobs.resolve(raw); ok { - data = resolved - } else if len(raw) == sha256.Size { - candidate := &agentv1.UserMessage{} - if err := proto.Unmarshal(raw, candidate); err != nil || !hasKnownUserMessageContent(candidate) { - return nil, fmt.Errorf("missing prefetched user message blob %x", raw) - } - return candidate, nil - } - message := &agentv1.UserMessage{} - if err := proto.Unmarshal(data, message); err != nil { - return nil, fmt.Errorf("decode imported turn user_message: %w", err) - } - return message, nil -} - -func decodeImportedStep(raw []byte, blobs importedBlobStore) (*agentv1.ConversationStep, error) { - data := raw - if resolved, ok := blobs.resolve(raw); ok { - data = resolved - } else if len(raw) == sha256.Size { - candidate := &agentv1.ConversationStep{} - if err := proto.Unmarshal(raw, candidate); err != nil || candidate.GetMessage() == nil { - return nil, fmt.Errorf("missing prefetched conversation step blob %x", raw) - } - return candidate, nil - } - step := &agentv1.ConversationStep{} - if err := proto.Unmarshal(data, step); err != nil { - return nil, fmt.Errorf("decode imported turn step: %w", err) - } - if step.GetMessage() == nil { - return nil, fmt.Errorf("decode imported turn step: payload is empty") - } - return step, nil -} - -func importedBlobTurnMessages(turn *agentv1.ConversationTurnStructure, blobs importedBlobStore) ([]modeladapter.Message, error) { - if turn == nil || turn.GetAgentConversationTurn() == nil { - return nil, nil - } - agentTurn := turn.GetAgentConversationTurn() - messages := make([]modeladapter.Message, 0, 1+len(agentTurn.GetSteps())) - if len(agentTurn.GetUserMessage()) > 0 { - userMessage, err := decodeImportedUserMessage(agentTurn.GetUserMessage(), blobs) - if err != nil { - return nil, err - } - if replay, ok := promptengine.BuildUserMessageReplayMessage(userMessage); ok { - messages = append(messages, toModelMessage(replay)) - } - } - for _, rawStep := range agentTurn.GetSteps() { - if len(rawStep) == 0 { - continue - } - step, err := decodeImportedStep(rawStep, blobs) - if err != nil { - return nil, err - } - for _, replay := range promptengine.BuildLegacyMessagesFromConversationStep(step) { - messages = append(messages, toModelMessage(replay)) - } - } - return messages, nil -} - -func importedTurnIDs(turns [][]byte, blobs importedBlobStore) ([][]byte, error) { - ids := make([][]byte, 0, len(turns)) - for _, raw := range turns { - if len(raw) == 0 { - continue - } - _, id, err := decodeImportedTurn(raw, blobs) - if err != nil { - return nil, err - } - if len(id) > 0 { - ids = append(ids, id) - } - } - return ids, nil -} - -func hasKnownUserMessageContent(message *agentv1.UserMessage) bool { - if message == nil { - return false - } - return message.GetText() != "" || - message.GetMessageId() != "" || - message.GetSelectedContext() != nil || - message.GetRichText() != "" || - len(message.GetConversationStateBlobId()) > 0 || - len(message.GetTextBlobId()) > 0 || - len(message.GetRichTextBlobId()) > 0 -} - -func firstNonNilError(err error, fallback error) error { - if err != nil { - return err - } - return fallback -} diff --git a/internal/backend/forwarder/projector.go b/internal/backend/forwarder/projector.go index 3f612d2..5726f8f 100644 --- a/internal/backend/forwarder/projector.go +++ b/internal/backend/forwarder/projector.go @@ -2,7 +2,6 @@ package forwarder import ( - "crypto/sha256" "encoding/json" "fmt" "strings" @@ -20,51 +19,6 @@ const projectedConversationMaxTokens = 130000 type HistoryProjector struct { } -type CheckpointBlob struct { - ID []byte - Data []byte -} - -type CheckpointProjection struct { - State *agentv1.ConversationStateStructure - Blobs []CheckpointBlob -} - -type checkpointBlobGraph struct { - blobs map[[sha256.Size]byte][]byte - order [][sha256.Size]byte -} - -func newCheckpointBlobGraph() *checkpointBlobGraph { - return &checkpointBlobGraph{blobs: make(map[[sha256.Size]byte][]byte)} -} - -func (graph *checkpointBlobGraph) add(data []byte) []byte { - if graph == nil || len(data) == 0 { - return nil - } - id := sha256.Sum256(data) - if _, exists := graph.blobs[id]; !exists { - graph.blobs[id] = append([]byte(nil), data...) - graph.order = append(graph.order, id) - } - return append([]byte(nil), id[:]...) -} - -func (graph *checkpointBlobGraph) list() []CheckpointBlob { - if graph == nil || len(graph.order) == 0 { - return nil - } - blobs := make([]CheckpointBlob, 0, len(graph.order)) - for _, id := range graph.order { - blobs = append(blobs, CheckpointBlob{ - ID: append([]byte(nil), id[:]...), - Data: append([]byte(nil), graph.blobs[id]...), - }) - } - return blobs -} - // NewHistoryProjector 创建 history 投影器。 func NewHistoryProjector() *HistoryProjector { return &HistoryProjector{} @@ -521,16 +475,6 @@ func isHistoricalReplayToolResult(conversation *ConversationFile, entry HistoryE // ProjectLegacyCheckpoint 按需从 JSON history 投影出兼容旧客户端的 checkpoint 结构。 func (projector *HistoryProjector) ProjectLegacyCheckpoint(conversation *ConversationFile) (*agentv1.ConversationStateStructure, error) { - projection, err := projector.ProjectCheckpointProjection(conversation) - if err != nil || projection == nil { - return nil, err - } - return projection.State, nil -} - -// ProjectCheckpointProjection 同时返回 checkpoint 状态及其引用的内容寻址 Blob。 -func (projector *HistoryProjector) ProjectCheckpointProjection(conversation *ConversationFile) (*CheckpointProjection, error) { - blobs := newCheckpointBlobGraph() state := &agentv1.ConversationStateStructure{ TokenDetails: &agentv1.ConversationTokenDetails{ UsedTokens: conversationTokenDetailsUsedTokens(conversation), @@ -544,7 +488,7 @@ func (projector *HistoryProjector) ProjectCheckpointProjection(conversation *Con if conversation == nil { mode := agentv1.AgentMode_AGENT_MODE_AGENT state.Mode = &mode - return &CheckpointProjection{State: state}, nil + return state, nil } mode, err := parseModeAlias(conversation.Mode) if err != nil { @@ -562,11 +506,158 @@ func (projector *HistoryProjector) ProjectCheckpointProjection(conversation *Con if structuredState.HasTodos { state.Todos = encodeConversationTodoBytes(structuredState.Todos) } - turnIDs, err := projectCheckpointTurnBlobs(conversation, blobs) - if err != nil { - return nil, err + grouped := make(map[int64][]HistoryEntry) + order := make([]int64, 0, conversation.NextTurnSeq) + for _, entry := range checkpointProjectionEntries(conversation.Entries) { + if entry.TurnSeq <= 0 { + continue + } + if _, ok := grouped[entry.TurnSeq]; !ok { + order = append(order, entry.TurnSeq) + } + grouped[entry.TurnSeq] = append(grouped[entry.TurnSeq], entry) + } + + for _, turnSeq := range order { + entries := grouped[turnSeq] + var rawUserMessage []byte + var turnRequestID string + steps := make([][]byte, 0, len(entries)) + seenToolCalls := make(map[string]struct{}) + openToolCalls := make(map[string]struct{}) + for _, entry := range entries { + if turnRequestID == "" { + turnRequestID = strings.TrimSpace(entry.RequestID) + } + switch strings.TrimSpace(entry.Kind) { + case "user_message": + userMessage := &agentv1.UserMessage{} + if err := protojson.Unmarshal(entry.Payload, userMessage); err != nil { + return nil, fmt.Errorf("decode checkpoint user_message: %w", err) + } + payload, err := proto.Marshal(userMessage) + if err != nil { + return nil, err + } + rawUserMessage = payload + case "assistant_text": + var payload assistantTextPayload + if err := json.Unmarshal(entry.Payload, &payload); err != nil { + return nil, err + } + if strings.TrimSpace(payload.Text) == "" && strings.TrimSpace(payload.ReasoningContent) != "" && len(openToolCalls) > 0 { + continue + } + if strings.TrimSpace(payload.ReasoningContent) != "" { + stepPayload, err := marshalThinkingStep(payload.ReasoningContent) + if err != nil { + return nil, err + } + steps = append(steps, stepPayload) + } + if strings.TrimSpace(payload.Text) == "" { + continue + } + stepPayload, err := proto.Marshal(&agentv1.ConversationStep{ + Message: &agentv1.ConversationStep_AssistantMessage{ + AssistantMessage: &agentv1.AssistantMessage{Text: strings.TrimSpace(payload.Text)}, + }, + }) + if err != nil { + return nil, err + } + steps = append(steps, stepPayload) + case "tool_call": + var payload toolCallEntryPayload + if err := json.Unmarshal(entry.Payload, &payload); err != nil { + return nil, err + } + if strings.TrimSpace(payload.ReasoningContent) != "" { + stepPayload, err := marshalThinkingStep(payload.ReasoningContent) + if err != nil { + return nil, err + } + steps = append(steps, stepPayload) + } + toolCall := &agentv1.ToolCall{} + if err := protojson.Unmarshal(payload.ToolCall, toolCall); err != nil { + return nil, err + } + if !shouldPersistToolResultName(firstNonEmpty(strings.TrimSpace(payload.ToolName), inferToolName(toolCall))) { + continue + } + stepPayload, err := proto.Marshal(&agentv1.ConversationStep{ + Message: &agentv1.ConversationStep_ToolCall{ + ToolCall: toolCall, + }, + }) + if err != nil { + return nil, err + } + steps = append(steps, stepPayload) + if toolCallID := strings.TrimSpace(payload.ToolCallID); toolCallID != "" { + seenToolCalls[toolCallID] = struct{}{} + openToolCalls[toolCallID] = struct{}{} + } + case "tool_result": + var payload toolResultEntryPayload + if err := json.Unmarshal(entry.Payload, &payload); err != nil { + return nil, err + } + if toolCallID := strings.TrimSpace(payload.ToolCallID); toolCallID != "" { + if _, ok := seenToolCalls[toolCallID]; ok { + delete(openToolCalls, toolCallID) + continue + } + } + if strings.TrimSpace(payload.ReasoningContent) != "" { + stepPayload, err := marshalThinkingStep(payload.ReasoningContent) + if err != nil { + return nil, err + } + steps = append(steps, stepPayload) + } + if len(payload.ToolCall) == 0 { + continue + } + toolCall := &agentv1.ToolCall{} + if err := protojson.Unmarshal(payload.ToolCall, toolCall); err != nil { + return nil, err + } + if !shouldPersistToolResultName(firstNonEmpty(strings.TrimSpace(payload.ToolName), inferToolName(toolCall))) { + continue + } + stepPayload, err := proto.Marshal(&agentv1.ConversationStep{ + Message: &agentv1.ConversationStep_ToolCall{ + ToolCall: toolCall, + }, + }) + if err != nil { + return nil, err + } + steps = append(steps, stepPayload) + } + } + if len(rawUserMessage) == 0 && len(steps) == 0 { + continue + } + agentTurn := &agentv1.AgentConversationTurnStructure{ + UserMessage: rawUserMessage, + Steps: steps, + } + if turnRequestID != "" { + agentTurn.RequestId = &turnRequestID + } + turnPayload, err := proto.Marshal(&agentv1.ConversationTurnStructure{ + Turn: &agentv1.ConversationTurnStructure_AgentConversationTurn{ + AgentConversationTurn: agentTurn, + }, + }) + if err != nil { + return nil, err + } + state.Turns = append(state.Turns, turnPayload) } - state.Turns = append(cloneByteSlices(conversation.ImportedTurnIDs), turnIDs...) replayMessages, err := projector.ProjectPromptReplay(conversation) if err != nil { return nil, err @@ -594,183 +685,15 @@ func (projector *HistoryProjector) ProjectCheckpointProjection(conversation *Con return nil, err } state.RootPromptMessagesJson = rootPromptMessages - return &CheckpointProjection{State: state, Blobs: blobs.list()}, nil + return state, nil } -func projectCheckpointTurnBlobs(conversation *ConversationFile, blobs *checkpointBlobGraph) ([][]byte, error) { - if conversation == nil || blobs == nil { - return nil, nil - } - grouped := make(map[int64][]HistoryEntry) - order := make([]int64, 0, conversation.NextTurnSeq) - for _, entry := range checkpointProjectionEntries(conversation.Entries) { - if entry.TurnSeq <= 0 { - continue - } - if _, ok := grouped[entry.TurnSeq]; !ok { - order = append(order, entry.TurnSeq) - } - grouped[entry.TurnSeq] = append(grouped[entry.TurnSeq], entry) - } - - turnIDs := make([][]byte, 0, len(order)) - for _, turnSeq := range order { - entries := grouped[turnSeq] - var userMessageID []byte - var turnRequestID string - stepIDs := make([][]byte, 0, len(entries)) - seenToolCalls := make(map[string]struct{}) - openToolCalls := make(map[string]struct{}) - for _, entry := range entries { - if turnRequestID == "" { - turnRequestID = strings.TrimSpace(entry.RequestID) - } - switch strings.TrimSpace(entry.Kind) { - case "user_message": - userMessage := &agentv1.UserMessage{} - if err := protojson.Unmarshal(entry.Payload, userMessage); err != nil { - return nil, fmt.Errorf("decode checkpoint user_message: %w", err) - } - payload, err := proto.Marshal(userMessage) - if err != nil { - return nil, err - } - userMessageID = blobs.add(payload) - case "assistant_text": - var payload assistantTextPayload - if err := json.Unmarshal(entry.Payload, &payload); err != nil { - return nil, err - } - if strings.TrimSpace(payload.Text) == "" && strings.TrimSpace(payload.ReasoningContent) != "" && len(openToolCalls) > 0 { - continue - } - if strings.TrimSpace(payload.ReasoningContent) != "" { - stepID, err := addCheckpointStepBlob(blobs, &agentv1.ConversationStep{ - Message: &agentv1.ConversationStep_ThinkingMessage{ - ThinkingMessage: &agentv1.ThinkingMessage{Text: payload.ReasoningContent}, - }, - }) - if err != nil { - return nil, err - } - stepIDs = append(stepIDs, stepID) - } - if strings.TrimSpace(payload.Text) == "" { - continue - } - stepID, err := addCheckpointStepBlob(blobs, &agentv1.ConversationStep{ - Message: &agentv1.ConversationStep_AssistantMessage{ - AssistantMessage: &agentv1.AssistantMessage{Text: strings.TrimSpace(payload.Text)}, - }, - }) - if err != nil { - return nil, err - } - stepIDs = append(stepIDs, stepID) - case "tool_call": - var payload toolCallEntryPayload - if err := json.Unmarshal(entry.Payload, &payload); err != nil { - return nil, err - } - if strings.TrimSpace(payload.ReasoningContent) != "" { - stepID, err := addCheckpointStepBlob(blobs, &agentv1.ConversationStep{ - Message: &agentv1.ConversationStep_ThinkingMessage{ - ThinkingMessage: &agentv1.ThinkingMessage{Text: payload.ReasoningContent}, - }, - }) - if err != nil { - return nil, err - } - stepIDs = append(stepIDs, stepID) - } - toolCall := &agentv1.ToolCall{} - if err := protojson.Unmarshal(payload.ToolCall, toolCall); err != nil { - return nil, err - } - if !shouldPersistToolResultName(firstNonEmpty(strings.TrimSpace(payload.ToolName), inferToolName(toolCall))) { - continue - } - stepID, err := addCheckpointStepBlob(blobs, &agentv1.ConversationStep{ - Message: &agentv1.ConversationStep_ToolCall{ToolCall: toolCall}, - }) - if err != nil { - return nil, err - } - stepIDs = append(stepIDs, stepID) - if toolCallID := strings.TrimSpace(payload.ToolCallID); toolCallID != "" { - seenToolCalls[toolCallID] = struct{}{} - openToolCalls[toolCallID] = struct{}{} - } - case "tool_result": - var payload toolResultEntryPayload - if err := json.Unmarshal(entry.Payload, &payload); err != nil { - return nil, err - } - if toolCallID := strings.TrimSpace(payload.ToolCallID); toolCallID != "" { - if _, ok := seenToolCalls[toolCallID]; ok { - delete(openToolCalls, toolCallID) - continue - } - } - if strings.TrimSpace(payload.ReasoningContent) != "" { - stepID, err := addCheckpointStepBlob(blobs, &agentv1.ConversationStep{ - Message: &agentv1.ConversationStep_ThinkingMessage{ - ThinkingMessage: &agentv1.ThinkingMessage{Text: payload.ReasoningContent}, - }, - }) - if err != nil { - return nil, err - } - stepIDs = append(stepIDs, stepID) - } - if len(payload.ToolCall) == 0 { - continue - } - toolCall := &agentv1.ToolCall{} - if err := protojson.Unmarshal(payload.ToolCall, toolCall); err != nil { - return nil, err - } - if !shouldPersistToolResultName(firstNonEmpty(strings.TrimSpace(payload.ToolName), inferToolName(toolCall))) { - continue - } - stepID, err := addCheckpointStepBlob(blobs, &agentv1.ConversationStep{ - Message: &agentv1.ConversationStep_ToolCall{ToolCall: toolCall}, - }) - if err != nil { - return nil, err - } - stepIDs = append(stepIDs, stepID) - } - } - if len(userMessageID) == 0 && len(stepIDs) == 0 { - continue - } - agentTurn := &agentv1.AgentConversationTurnStructure{ - UserMessage: userMessageID, - Steps: stepIDs, - } - if turnRequestID != "" { - agentTurn.RequestId = &turnRequestID - } - turnPayload, err := proto.Marshal(&agentv1.ConversationTurnStructure{ - Turn: &agentv1.ConversationTurnStructure_AgentConversationTurn{ - AgentConversationTurn: agentTurn, - }, - }) - if err != nil { - return nil, err - } - turnIDs = append(turnIDs, blobs.add(turnPayload)) - } - return turnIDs, nil -} - -func addCheckpointStepBlob(blobs *checkpointBlobGraph, step *agentv1.ConversationStep) ([]byte, error) { - payload, err := proto.Marshal(step) - if err != nil { - return nil, err - } - return blobs.add(payload), nil +func marshalThinkingStep(text string) ([]byte, error) { + return proto.Marshal(&agentv1.ConversationStep{ + Message: &agentv1.ConversationStep_ThinkingMessage{ + ThinkingMessage: &agentv1.ThinkingMessage{Text: text}, + }, + }) } func conversationTokenDetailsUsedTokens(conversation *ConversationFile) uint32 { @@ -1218,6 +1141,60 @@ func shouldPersistToolResultName(toolName string) bool { } } +func filterCheckpointTurns(rawTurns [][]byte) [][]byte { + if len(rawTurns) == 0 { + return nil + } + filtered := make([][]byte, 0, len(rawTurns)) + for _, rawTurn := range rawTurns { + if len(rawTurn) == 0 { + continue + } + turn := &agentv1.ConversationTurnStructure{} + if err := proto.Unmarshal(rawTurn, turn); err != nil { + filtered = append(filtered, append([]byte(nil), rawTurn...)) + continue + } + agentTurn := turn.GetAgentConversationTurn() + if agentTurn == nil { + filtered = append(filtered, append([]byte(nil), rawTurn...)) + continue + } + + nextSteps := make([][]byte, 0, len(agentTurn.GetSteps())) + for _, rawStep := range agentTurn.GetSteps() { + if len(rawStep) == 0 { + continue + } + step := &agentv1.ConversationStep{} + if err := proto.Unmarshal(rawStep, step); err != nil { + continue + } + if toolCall := step.GetToolCall(); toolCall != nil && !shouldPersistToolResultName(inferToolName(toolCall)) { + continue + } + nextSteps = append(nextSteps, append([]byte(nil), rawStep...)) + } + if len(agentTurn.GetUserMessage()) == 0 && len(nextSteps) == 0 { + continue + } + encoded, err := proto.Marshal(&agentv1.ConversationTurnStructure{ + Turn: &agentv1.ConversationTurnStructure_AgentConversationTurn{ + AgentConversationTurn: &agentv1.AgentConversationTurnStructure{ + UserMessage: append([]byte(nil), agentTurn.GetUserMessage()...), + Steps: nextSteps, + }, + }, + }) + if err != nil { + filtered = append(filtered, append([]byte(nil), rawTurn...)) + continue + } + filtered = append(filtered, encoded) + } + return filtered +} + func filterCheckpointPersistentToolReplay(messages []promptengine.Message) []promptengine.Message { if len(messages) == 0 { return nil @@ -1254,7 +1231,7 @@ func filterCheckpointPersistentToolReplay(messages []promptengine.Message) []pro return filtered } -func restoreImportedReplayUserMessages(messages []promptengine.Message, importedTurns [][]byte, blobs importedBlobStore) []promptengine.Message { +func restoreImportedReplayUserMessages(messages []promptengine.Message, importedTurns [][]byte) []promptengine.Message { if len(messages) == 0 || len(importedTurns) == 0 { return messages } @@ -1263,16 +1240,16 @@ func restoreImportedReplayUserMessages(messages []promptengine.Message, imported if len(rawTurn) == 0 { continue } - turn, _, err := decodeImportedTurn(rawTurn, blobs) - if err != nil || turn == nil { + turn := &agentv1.ConversationTurnStructure{} + if err := proto.Unmarshal(rawTurn, turn); err != nil { continue } agentTurn := turn.GetAgentConversationTurn() if agentTurn == nil || len(agentTurn.GetUserMessage()) == 0 { continue } - userMessage, err := decodeImportedUserMessage(agentTurn.GetUserMessage(), blobs) - if err != nil { + userMessage := &agentv1.UserMessage{} + if err := proto.Unmarshal(agentTurn.GetUserMessage(), userMessage); err != nil { continue } replay, ok := promptengine.BuildUserMessageReplayMessage(userMessage) diff --git a/internal/backend/forwarder/projector_test.go b/internal/backend/forwarder/projector_test.go deleted file mode 100644 index 10330ba..0000000 --- a/internal/backend/forwarder/projector_test.go +++ /dev/null @@ -1,433 +0,0 @@ -package forwarder - -import ( - "crypto/sha256" - "encoding/json" - "fmt" - "strings" - "testing" - - "google.golang.org/protobuf/encoding/protojson" - "google.golang.org/protobuf/proto" - - "cursor/gen/agentv1" - modeladapter "cursor/internal/backend/agent/model" - promptengine "cursor/internal/backend/agent/prompt" -) - -func TestProjectCheckpointProjectionBuildsBlobBackedTurns(t *testing.T) { - toolCall := testEditToolCall(t, "file.txt") - - tests := []struct { - name string - entries []HistoryEntry - }{ - { - name: "no tools", - entries: []HistoryEntry{ - testUserMessageEntry(t, 1, "request-1", "hello"), - newAssistantTextEntry(1, "request-1", "hi", "", ""), - }, - }, - { - name: "completed tool call", - entries: []HistoryEntry{ - testUserMessageEntry(t, 1, "request-1", "edit the file"), - newToolCallEntry(1, "request-1", "call-1", "Edit", "", "", toolCall), - newToolResultEntry(1, "request-1", "call-1", "Edit", `{"path":"file.txt"}`, "edited", "", toolCall), - newAssistantTextEntry(1, "request-1", "done", "", ""), - }, - }, - { - name: "unfinished tool call", - entries: []HistoryEntry{ - testUserMessageEntry(t, 1, "request-1", "edit the file"), - newToolCallEntry(1, "request-1", "call-1", "Edit", "", "", toolCall), - }, - }, - { - name: "orphan tool result", - entries: []HistoryEntry{ - testUserMessageEntry(t, 1, "request-1", "edit the file"), - newToolResultEntry(1, "request-1", "call-1", "Edit", `{"path":"file.txt"}`, "edited", "", toolCall), - }, - }, - } - - for _, test := range tests { - t.Run(test.name, func(t *testing.T) { - conversation := testConversation(test.entries) - projection, err := NewHistoryProjector().ProjectCheckpointProjection(conversation) - if err != nil { - t.Fatalf("ProjectCheckpointProjection() error = %v", err) - } - if len(projection.State.GetTurns()) != 1 { - t.Fatalf("ProjectCheckpointProjection() turns = %d, want 1 Blob ID", len(projection.State.GetTurns())) - } - assertCheckpointBlobGraph(t, projection) - messages, err := promptengine.DecodeReplayMessages(projection.State.GetRootPromptMessagesJson()) - if err != nil { - t.Fatalf("DecodeReplayMessages() error = %v", err) - } - if len(messages) == 0 { - t.Fatal("ProjectCheckpointProjection() removed all root prompt replay history") - } - if messages[0].Role != "user" || messages[0].Content == "" { - t.Fatalf("first replay message = %#v, want retained user history", messages[0]) - } - }) - } -} - -func TestProjectLegacyCheckpointLargeModelHistoryUsesRootReplay(t *testing.T) { - entries := make([]HistoryEntry, 0, 400) - for turn := int64(1); turn <= 200; turn++ { - requestID := fmt.Sprintf("request-%d", turn) - entries = append(entries, - testModelMessageEntry(t, turn, requestID, modeladapter.Message{Role: "user", Content: fmt.Sprintf("question %d", turn)}), - testModelMessageEntry(t, turn, requestID, modeladapter.Message{Role: "assistant", Content: fmt.Sprintf("answer %d", turn)}), - ) - } - - state, err := NewHistoryProjector().ProjectLegacyCheckpoint(testConversation(entries)) - if err != nil { - t.Fatalf("ProjectLegacyCheckpoint() error = %v", err) - } - if len(state.GetTurns()) != 0 { - t.Fatalf("ProjectLegacyCheckpoint() model-only turns = %d, want 0", len(state.GetTurns())) - } - messages, err := promptengine.DecodeReplayMessages(state.GetRootPromptMessagesJson()) - if err != nil { - t.Fatalf("DecodeReplayMessages() error = %v", err) - } - if len(messages) != 400 { - t.Fatalf("decoded replay messages = %d, want 400", len(messages)) - } -} - -func TestProjectLegacyCheckpointSnapshotIsIsolatedFromLaterHistory(t *testing.T) { - conversation := testConversation([]HistoryEntry{ - testUserMessageEntry(t, 1, "request-1", "first question"), - newAssistantTextEntry(1, "request-1", "first answer", "", ""), - }) - projector := NewHistoryProjector() - midpointProjection, err := projector.ProjectCheckpointProjection(conversation) - if err != nil { - t.Fatalf("midpoint ProjectCheckpointProjection() error = %v", err) - } - midpoint := midpointProjection.State - - appendEntriesInPlace(conversation, []HistoryEntry{ - testUserMessageEntry(t, 2, "request-2", "second question"), - newAssistantTextEntry(2, "request-2", "second answer", "", ""), - }) - latestProjection, err := projector.ProjectCheckpointProjection(conversation) - if err != nil { - t.Fatalf("latest ProjectCheckpointProjection() error = %v", err) - } - latest := latestProjection.State - - midpointMessages, err := promptengine.DecodeReplayMessages(midpoint.GetRootPromptMessagesJson()) - if err != nil { - t.Fatalf("decode midpoint replay: %v", err) - } - latestMessages, err := promptengine.DecodeReplayMessages(latest.GetRootPromptMessagesJson()) - if err != nil { - t.Fatalf("decode latest replay: %v", err) - } - if len(midpointMessages) != 2 { - t.Fatalf("midpoint replay messages = %d, want 2", len(midpointMessages)) - } - if len(latestMessages) != 4 { - t.Fatalf("latest replay messages = %d, want 4", len(latestMessages)) - } - if len(midpoint.GetTurns()) != 1 || len(latest.GetTurns()) != 2 { - t.Fatalf("checkpoint turn counts = (%d, %d), want (1, 2)", len(midpoint.GetTurns()), len(latest.GetTurns())) - } - assertCheckpointBlobGraph(t, midpointProjection) - assertCheckpointBlobGraph(t, latestProjection) -} - -func TestProjectCheckpointProjectionKeepsVisibleTurnsAcrossCompaction(t *testing.T) { - summaryPayload, err := json.Marshal(compactionSummaryEntryPayload{Summary: "first turn summarized"}) - if err != nil { - t.Fatalf("marshal compaction summary: %v", err) - } - conversation := testConversation([]HistoryEntry{ - testUserMessageEntry(t, 1, "request-1", "first question"), - newAssistantTextEntry(1, "request-1", "first answer", "", ""), - {TurnSeq: 0, Role: "system", Kind: "compaction_summary", Payload: summaryPayload}, - testUserMessageEntry(t, 2, "request-2", "second question"), - newAssistantTextEntry(2, "request-2", "second answer", "", ""), - }) - - projection, err := NewHistoryProjector().ProjectCheckpointProjection(conversation) - if err != nil { - t.Fatalf("ProjectCheckpointProjection() error = %v", err) - } - if len(projection.State.GetTurns()) != 2 { - t.Fatalf("visible turns after compaction = %d, want 2", len(projection.State.GetTurns())) - } - assertCheckpointBlobGraph(t, projection) - - messages, err := promptengine.DecodeReplayMessages(projection.State.GetRootPromptMessagesJson()) - if err != nil { - t.Fatalf("DecodeReplayMessages() error = %v", err) - } - if len(messages) != 3 { - t.Fatalf("compacted root replay messages = %d, want summary plus latest turn", len(messages)) - } - if messages[0].Role != "user" || messages[0].Content != "\nfirst turn summarized\n" { - t.Fatalf("first compacted replay message = %#v", messages[0]) - } -} - -func TestImportedConversationStateRejectsBlobTurnIDsWithoutPrefetchedData(t *testing.T) { - turnID := sha256.Sum256([]byte("imported turn")) - state := &agentv1.ConversationStateStructure{Turns: [][]byte{turnID[:]}} - if _, err := importedConversationStateModelMessages(state, nil); err == nil { - t.Fatal("importedConversationStateModelMessages() accepted unresolved Blob turn") - } - conversation := testConversation(nil) - service := &Service{} - if _, err := service.importConversationState(conversation, state, nil); err == nil { - t.Fatal("importConversationState() accepted unresolved Blob turn") - } -} - -func TestImportedConversationStateRestoresBlobOnlyForkFromPrefetchedBlobs(t *testing.T) { - projection, err := NewHistoryProjector().ProjectCheckpointProjection(testConversation([]HistoryEntry{ - testUserMessageEntry(t, 1, "request-1", "parent question"), - newAssistantTextEntry(1, "request-1", "parent answer", "", ""), - })) - if err != nil { - t.Fatalf("ProjectCheckpointProjection() error = %v", err) - } - prefetched := make([]*agentv1.PreFetchedBlob, 0, len(projection.Blobs)) - for _, blob := range projection.Blobs { - prefetched = append(prefetched, &agentv1.PreFetchedBlob{Id: blob.ID, Value: blob.Data}) - } - state := proto.Clone(projection.State).(*agentv1.ConversationStateStructure) - state.RootPromptMessagesJson = nil - conversation := testConversation(nil) - entries, err := (&Service{}).importConversationState(conversation, state, prefetched) - if err != nil { - t.Fatalf("importConversationState() error = %v", err) - } - if len(conversation.ImportedTurnIDs) != 1 { - t.Fatalf("ImportedTurnIDs = %d, want 1", len(conversation.ImportedTurnIDs)) - } - if len(entries) != 2 { - t.Fatalf("imported model entries = %d, want user and assistant", len(entries)) - } -} - -func TestImportedInlineTurnWithSHA256LengthIsNotMisclassified(t *testing.T) { - var rawTurn []byte - for size := 1; size <= 128; size++ { - rawUser, err := proto.Marshal(&agentv1.UserMessage{Text: strings.Repeat("x", size), MessageId: "inline"}) - if err != nil { - t.Fatalf("marshal user message: %v", err) - } - rawTurn, err = proto.Marshal(&agentv1.ConversationTurnStructure{ - Turn: &agentv1.ConversationTurnStructure_AgentConversationTurn{ - AgentConversationTurn: &agentv1.AgentConversationTurnStructure{UserMessage: rawUser}, - }, - }) - if err != nil { - t.Fatalf("marshal turn: %v", err) - } - if len(rawTurn) == sha256.Size { - break - } - } - if len(rawTurn) != sha256.Size { - t.Fatal("test could not construct a 32-byte inline turn") - } - ids, err := importedTurnIDs([][]byte{rawTurn}, nil) - if err != nil { - t.Fatalf("importedTurnIDs() error = %v", err) - } - if len(ids) != 0 { - t.Fatal("32-byte inline turn was misclassified as a Blob ID") - } - messages, err := importedConversationStateModelMessages(&agentv1.ConversationStateStructure{Turns: [][]byte{rawTurn}}, nil) - if err != nil { - t.Fatalf("importedConversationStateModelMessages() error = %v", err) - } - if len(messages) != 1 || messages[0].Role != "user" { - t.Fatalf("inline turn messages = %#v, want one user message", messages) - } -} - -func TestImportedTurnIDsPersistThroughConversationStore(t *testing.T) { - store := NewConversationFileStore(t.TempDir()) - turnID := sha256.Sum256([]byte("parent turn")) - conversation := testConversation(nil) - conversation.ImportedTurnIDs = [][]byte{turnID[:]} - persisted, err := store.SaveConversationWithEntries(conversation.ConversationID, conversation, []HistoryEntry{ - testUserMessageEntry(t, 2, "request-2", "fork question"), - }) - if err != nil { - t.Fatalf("SaveConversationWithEntries() error = %v", err) - } - if len(persisted.ImportedTurnIDs) != 1 || string(persisted.ImportedTurnIDs[0]) != string(turnID[:]) { - t.Fatalf("persisted ImportedTurnIDs = %x, want %x", persisted.ImportedTurnIDs, turnID) - } - loaded, err := store.LoadConversation(conversation.ConversationID) - if err != nil { - t.Fatalf("LoadConversation() error = %v", err) - } - if len(loaded.ImportedTurnIDs) != 1 || string(loaded.ImportedTurnIDs[0]) != string(turnID[:]) { - t.Fatalf("loaded ImportedTurnIDs = %x, want %x", loaded.ImportedTurnIDs, turnID) - } -} - -func TestRewindImportedTurnPrefixUsesClientForkPoint(t *testing.T) { - ids := make([][]byte, 4) - for index := range ids { - digest := sha256.Sum256([]byte(fmt.Sprintf("turn-%d", index+1))) - ids[index] = digest[:] - } - trimmed := rewindImportedTurnPrefix(ids, runRewindDecision{ - TargetTurnSeq: 4, - HasClientTurnCount: true, - ClientTurnCount: 2, - }) - if len(trimmed) != 2 || string(trimmed[0]) != string(ids[0]) || string(trimmed[1]) != string(ids[1]) { - t.Fatalf("rewindImportedTurnPrefix() = %x, want first two IDs", trimmed) - } -} - -func TestRewindImportedTurnPrefixClearsAllIDsAtClientTurnZero(t *testing.T) { - ids := make([][]byte, 2) - for index := range ids { - digest := sha256.Sum256([]byte(fmt.Sprintf("turn-%d", index+1))) - ids[index] = digest[:] - } - trimmed := rewindImportedTurnPrefix(ids, runRewindDecision{ - TargetTurnSeq: 3, - HasClientTurnCount: true, - ClientTurnCount: 0, - }) - if trimmed != nil { - t.Fatalf("rewindImportedTurnPrefix() = %x, want nil at client turn zero", trimmed) - } -} - -func TestRewindImportedTurnPrefixUsesTargetWithoutClientCount(t *testing.T) { - ids := make([][]byte, 4) - for index := range ids { - digest := sha256.Sum256([]byte(fmt.Sprintf("turn-%d", index+1))) - ids[index] = digest[:] - } - trimmed := rewindImportedTurnPrefix(ids, runRewindDecision{TargetTurnSeq: 3}) - if len(trimmed) != 2 || string(trimmed[0]) != string(ids[0]) || string(trimmed[1]) != string(ids[1]) { - t.Fatalf("rewindImportedTurnPrefix() = %x, want target-derived first two IDs", trimmed) - } -} - -func assertCheckpointBlobGraph(t *testing.T, projection *CheckpointProjection) { - t.Helper() - if projection == nil || projection.State == nil { - t.Fatal("checkpoint projection is nil") - } - blobByID := make(map[string][]byte, len(projection.Blobs)) - for _, blob := range projection.Blobs { - if len(blob.ID) != sha256.Size { - t.Fatalf("blob id length = %d, want %d", len(blob.ID), sha256.Size) - } - digest := sha256.Sum256(blob.Data) - if string(blob.ID) != string(digest[:]) { - t.Fatal("blob id does not match SHA-256(data)") - } - blobByID[string(blob.ID)] = blob.Data - } - for _, turnID := range projection.State.GetTurns() { - turnData, ok := blobByID[string(turnID)] - if !ok { - t.Fatal("turn references missing blob") - } - turn := &agentv1.ConversationTurnStructure{} - if err := proto.Unmarshal(turnData, turn); err != nil { - t.Fatalf("decode turn blob: %v", err) - } - agentTurn := turn.GetAgentConversationTurn() - if agentTurn == nil { - continue - } - if userID := agentTurn.GetUserMessage(); len(userID) > 0 { - userData, exists := blobByID[string(userID)] - if !exists { - t.Fatal("turn references missing user message blob") - } - if err := proto.Unmarshal(userData, &agentv1.UserMessage{}); err != nil { - t.Fatalf("decode user message blob: %v", err) - } - } - for _, stepID := range agentTurn.GetSteps() { - stepData, exists := blobByID[string(stepID)] - if !exists { - t.Fatal("turn references missing step blob") - } - if err := proto.Unmarshal(stepData, &agentv1.ConversationStep{}); err != nil { - t.Fatalf("decode conversation step blob: %v", err) - } - } - } -} - -func testConversation(entries []HistoryEntry) *ConversationFile { - conversation := &ConversationFile{ - ConversationID: "conversation-1", - RootConversationID: "conversation-1", - Mode: "agent", - NextTurnSeq: 1, - NextEntrySeq: 1, - Entries: make([]HistoryEntry, 0, len(entries)), - } - appendEntriesInPlace(conversation, entries) - return conversation -} - -func testUserMessageEntry(t *testing.T, turnSeq int64, requestID string, text string) HistoryEntry { - t.Helper() - payload, err := protojson.Marshal(&agentv1.UserMessage{Text: text, MessageId: fmt.Sprintf("message-%d", turnSeq)}) - if err != nil { - t.Fatalf("marshal user message: %v", err) - } - return HistoryEntry{ - TurnSeq: turnSeq, - RequestID: requestID, - Role: "user", - Kind: "user_message", - Payload: payload, - } -} - -func testModelMessageEntry(t *testing.T, turnSeq int64, requestID string, message modeladapter.Message) HistoryEntry { - t.Helper() - entry, ok, err := newModelMessageEntry(turnSeq, requestID, message) - if err != nil { - t.Fatalf("newModelMessageEntry() error = %v", err) - } - if !ok { - t.Fatal("newModelMessageEntry() rejected test message") - } - return entry -} - -func testEditToolCall(t *testing.T, path string) []byte { - t.Helper() - payload, err := protojson.Marshal(&agentv1.ToolCall{ - Tool: &agentv1.ToolCall_EditToolCall{ - EditToolCall: &agentv1.EditToolCall{ - Args: &agentv1.EditArgs{Path: path}, - }, - }, - }) - if err != nil { - t.Fatalf("marshal edit tool call: %v", err) - } - return payload -} diff --git a/internal/backend/forwarder/rewind.go b/internal/backend/forwarder/rewind.go index 1f11ee0..ba7c851 100644 --- a/internal/backend/forwarder/rewind.go +++ b/internal/backend/forwarder/rewind.go @@ -36,7 +36,7 @@ type runRewindMatch struct { } func (service *Service) decideRunRewind(intent InboundIntent, conversation *ConversationFile) runRewindDecision { - decision := runRewindDecision{} + decision := runRewindDecision{ClientTurnCount: -1} if !shouldEvaluateRunRewind(intent) { return decision } @@ -121,7 +121,7 @@ func selectRunRewindMatch(matches []runRewindMatch, clientTurnCount int, hasClie if len(matches) == 0 { return runRewindMatch{}, "no_match" } - if hasClientTurnCount { + if hasClientTurnCount && clientTurnCount >= 0 { targetTurnSeq := int64(clientTurnCount) + 1 for _, match := range matches { if match.Entry.TurnSeq == targetTurnSeq { @@ -224,7 +224,6 @@ func (service *Service) applyRunRewindToConversation(conversation *ConversationF conversation.Entries = nil conversation.NextEntrySeq = 1 conversation.NextTurnSeq = 1 - conversation.ImportedTurnIDs = rewindImportedTurnPrefix(conversation.ImportedTurnIDs, decision) appendEntriesInPlace(conversation, appendReplacementRunEntries(decision.PrefixEntries, entries)) applyRunRewindConversationState(conversation, intent, turnSeq) deriveConversationLoopState(conversation) @@ -270,30 +269,10 @@ func applyRunRewindMetadata(conversation *ConversationFile, source *Conversation if source.TokenDetailsMaxTokens > 0 { conversation.TokenDetailsMaxTokens = source.TokenDetailsMaxTokens } - decision := runRewindDecision{TargetTurnSeq: turnSeq} - if intent.ConversationState != nil { - decision.HasClientTurnCount = true - decision.ClientTurnCount = len(intent.ConversationState.GetTurns()) - } - conversation.ImportedTurnIDs = rewindImportedTurnPrefix(source.ImportedTurnIDs, decision) } applyRunRewindConversationState(conversation, intent, turnSeq) } -func rewindImportedTurnPrefix(importedTurnIDs [][]byte, decision runRewindDecision) [][]byte { - keep := decision.TargetTurnSeq - 1 - if decision.HasClientTurnCount { - keep = int64(decision.ClientTurnCount) - } - if keep <= 0 || len(importedTurnIDs) == 0 { - return nil - } - if keep > int64(len(importedTurnIDs)) { - keep = int64(len(importedTurnIDs)) - } - return cloneByteSlices(importedTurnIDs[:keep]) -} - func (service *Service) logRunRewindDecision(requestID string, conversationID string, eventName string, decision runRewindDecision) { if service == nil || !decision.Evaluated { return diff --git a/internal/backend/forwarder/runtime_summary.go b/internal/backend/forwarder/runtime_summary.go index b302c8a..9abbba3 100644 --- a/internal/backend/forwarder/runtime_summary.go +++ b/internal/backend/forwarder/runtime_summary.go @@ -50,7 +50,7 @@ func (service *Service) bootstrapRuntimeConversation(intent InboundIntent) (*Con } importedEntries := []HistoryEntry(nil) if len(conversation.Entries) == 0 && intent.ConversationState != nil { - importedEntries, err = service.importConversationState(conversation, intent.ConversationState, intent.PreFetchedBlobs) + importedEntries, err = service.importConversationState(conversation, intent.ConversationState) if err != nil { return nil, agentv1.AgentMode_AGENT_MODE_AGENT, 0, nil, err } @@ -138,7 +138,6 @@ func (service *Service) syncConversationRecord(conversationID string, conversati item.AutoCompactionReserveTokens = conversation.AutoCompactionReserveTokens item.AutoCompactionTriggeredAt = conversation.AutoCompactionTriggeredAt item.AutoCompactionSourceModelCallID = conversation.AutoCompactionSourceModelCallID - item.ImportedTurnIDs = cloneByteSlices(conversation.ImportedTurnIDs) item.LatestRequestPrefix = cloneConversationRequestPrefix(conversation.LatestRequestPrefix) item.LastProviderCall = cloneConversationProviderCall(conversation.LastProviderCall) item.CreatedAt = conversation.CreatedAt diff --git a/internal/backend/forwarder/service.go b/internal/backend/forwarder/service.go index fb4c848..5bb04ca 100644 --- a/internal/backend/forwarder/service.go +++ b/internal/backend/forwarder/service.go @@ -9,7 +9,6 @@ import ( "log" "sort" "strings" - "sync" "time" "connectrpc.com/connect" @@ -262,8 +261,6 @@ type Service struct { execBridge execbridge.ExecBridge interactionBridge interactionbridge.InteractionBridge appendSeq *appendSequenceTracker - checkpointBlobMu sync.Mutex - checkpointBlobs map[string]*checkpointBlobCacheEntry } type agentModelMemory interface { @@ -303,7 +300,6 @@ func NewService(historyRoot string, resolver modeladapter.ChannelResolver) *Serv execBridge: execbridge.NewBridge(), interactionBridge: interactionbridge.NewBridge(), appendSeq: newAppendSequenceTracker(), - checkpointBlobs: make(map[string]*checkpointBlobCacheEntry), } service.startHistoryMaintenance() store.SyncAllCursorTranscriptsBestEffort() @@ -332,7 +328,6 @@ func newServiceWithDependencies(store *ConversationFileStore, projector *History execBridge: execbridge.NewBridge(), interactionBridge: interactionbridge.NewBridge(), appendSeq: newAppendSequenceTracker(), - checkpointBlobs: make(map[string]*checkpointBlobCacheEntry), } } @@ -562,7 +557,6 @@ func (service *Service) decodeInboundIntent(requestID string, message *agentv1.A } intent.ConversationID = conversationID intent.ConversationState = runRequest.GetConversationState() - intent.PreFetchedBlobs = runRequest.GetPreFetchedBlobs() intent.UserMessage = extractUserMessage(message) intent.RequestContext = extractRequestContext(message) if service.shouldIgnoreEmptyResumeRunRequest(requestID, runRequest, intent.UserMessage, intent.RequestContext) { @@ -610,7 +604,6 @@ func (service *Service) decodeInboundIntent(requestID string, message *agentv1.A intent.ConversationID = conversationID intent.SubagentTypeName = strings.TrimSpace(prewarmRequest.GetSubagentTypeName()) intent.ConversationState = prewarmRequest.GetConversationState() - intent.PreFetchedBlobs = prewarmRequest.GetPreFetchedBlobs() intent.Mode, intent.ModeSource, intent.HasExplicitMode, err = extractPrewarmMode(prewarmRequest) if err != nil { return InboundIntent{}, err @@ -826,14 +819,11 @@ func (service *Service) snapshotVisibleTurns(conversation *ConversationFile) ([] if service == nil || service.projector == nil || conversation == nil { return nil, nil } - projection, err := service.projector.ProjectCheckpointProjection(conversation) + state, err := service.projector.ProjectLegacyCheckpoint(conversation) if err != nil { return nil, err } - if projection == nil || projection.State == nil { - return nil, fmt.Errorf("checkpoint projection is empty") - } - return cloneByteSlices(projection.State.GetTurns()), nil + return cloneByteSlices(state.GetTurns()), nil } // handleCancelIntent 处理取消请求,并向客户端发送执行桥 abort。 @@ -843,60 +833,42 @@ func (service *Service) handleCancelIntent(intent InboundIntent) error { return fmt.Errorf("request is not active: %s", intent.RequestID) } hasCheckpoint := checkpointConversationInitialized(stream) + if hasCheckpoint { + cancelReason := firstNonEmpty(intent.CancelReason, "user aborted") + _, err := service.appendConversationEntries(stream, stream.ConversationID, []HistoryEntry{ + newMetadataEntry(stream.TurnSeq, intent.RequestID, "control", map[string]any{ + "status": "canceled", + "reason": cancelReason, + "replay_policy": cancelReplayPolicyForReason(cancelReason), + }), + }) + if err != nil { + return err + } + } stream.mu.Lock() pendingExecs := make([]runtimecore.PendingExec, 0, len(stream.PendingExecs)) for _, pending := range stream.PendingExecs { pendingExecs = append(pendingExecs, pending) } - if stream.ProviderCancel != nil { - stream.ProviderCancel() - stream.ProviderCancel = nil - } - stream.ProviderActive = false - stream.CurrentProviderToken++ - stream.CurrentCompactionToken++ - stream.PendingProviderAction = providerActionNone - stream.PendingCompaction = nil - stream.UpdatedAt = time.Now().UTC() stream.mu.Unlock() - if hasCheckpoint { - cancelReason := firstNonEmpty(intent.CancelReason, "user aborted") - cancelEntry := newMetadataEntry(stream.TurnSeq, intent.RequestID, "control", map[string]any{ - "status": "canceled", - "reason": cancelReason, - "replay_policy": cancelReplayPolicyForReason(cancelReason), - }) - if _, err := service.appendConversationEntries(stream, stream.ConversationID, []HistoryEntry{cancelEntry}); err != nil { - log.Printf("forwarder cancellation metadata persistence failed request_id=%s conversation_id=%s err=%v", stream.RequestID, stream.ConversationID, err) - if memoryErr := service.appendCheckpointEntries(stream, []HistoryEntry{cancelEntry}); memoryErr != nil { - return memoryErr - } - } - } for _, pending := range pendingExecs { _ = service.broker.Publish(intent.RequestID, StreamEvent{ Message: buildExecAbortMessage(pending), }) } + if hasCheckpoint { + if err := service.publishCheckpoint(stream.RequestID, stream.ConversationID); err != nil { + return err + } + } clearPendingProviderCompletion(stream) - terminalMessage := firstNonEmpty(intent.CancelReason, "[canceled] User aborted request") stream.mu.Lock() - stream.PendingExecs = make(map[string]runtimecore.PendingExec) - stream.PendingInteractions = make(map[string]runtimecore.PendingInteraction) + stream.PendingProviderAction = providerActionNone stream.UpdatedAt = time.Now().UTC() stream.mu.Unlock() - service.discardPendingCheckpoint(stream, fmt.Errorf("checkpoint superseded by cancellation")) - if hasCheckpoint { - if err := service.publishCheckpointWithTerminalAction( - stream.RequestID, - stream.ConversationID, - checkpointCancellationAction(terminalMessage), - ); err != nil { - return service.failTerminalCheckpointSync(stream, err) - } - return nil - } - return service.finishCanceledTurnAfterCheckpoint(stream, terminalMessage) + service.setTurnPhase(stream, TurnPhaseCanceled) + return service.broker.Cancel(intent.RequestID, firstNonEmpty(intent.CancelReason, "[canceled] User aborted request")) } // handleExecResult 处理客户端返回的执行桥结果,并在终态时把 tool_result 写回 history。 @@ -2150,18 +2122,9 @@ func (service *Service) completeSuccessfulTurn(stream *ActiveStream, completion err, ) } - if err := service.publishCheckpointWithCompletion(requestID, conversationID, &completion); err != nil { + if err := service.publishCheckpoint(requestID, conversationID); err != nil { return err } - return nil -} - -func (service *Service) finishSuccessfulTurnAfterCheckpoint(stream *ActiveStream, completion pendingTurnCompletion) error { - if stream == nil { - return nil - } - requestID := firstNonEmpty(strings.TrimSpace(completion.RequestID), strings.TrimSpace(stream.RequestID)) - usage := completion.Usage if err := service.broker.Publish(requestID, StreamEvent{ Message: buildTurnEndedMessage(usage.InputTokens, usage.OutputTokens, usage.CacheReadTokens, usage.CacheWriteTokens), }); err != nil { @@ -2188,15 +2151,7 @@ func (service *Service) failStreamIfNonTerminal(stream *ActiveStream, terminalCo } // publishCheckpoint 按当前内存会话镜像投影出 checkpoint,并广播给所有 RunSSE 订阅者。 -func (service *Service) publishCheckpoint(requestID string, conversationID string) error { - return service.publishCheckpointWithCompletion(requestID, conversationID, nil) -} - -func (service *Service) publishCheckpointWithCompletion(requestID string, conversationID string, completion *pendingTurnCompletion) error { - return service.publishCheckpointWithTerminalAction(requestID, conversationID, checkpointCompletionAction(completion)) -} - -func (service *Service) publishCheckpointWithTerminalAction(requestID string, conversationID string, terminalAction checkpointTerminalAction) error { +func (service *Service) publishCheckpoint(requestID string, _ string) error { stream, ok := service.broker.Get(requestID) if !ok || stream == nil { return fmt.Errorf("request is not active: %s", requestID) @@ -2205,16 +2160,15 @@ func (service *Service) publishCheckpointWithTerminalAction(requestID string, co if err != nil { return err } - projection, err := service.projector.ProjectCheckpointProjection(conversation) + state, err := service.projector.ProjectLegacyCheckpoint(conversation) if err != nil { return err } - if projection == nil || projection.State == nil { - return fmt.Errorf("checkpoint projection is empty") - } - projection.State.PendingToolCalls = buildPendingToolCalls(pendingExecs, pendingInteractions) - service.rewriteCheckpointTokenDetailsForClient(stream, conversation, projection.State) - return service.queueCheckpointProjection(stream, projection, terminalAction) + state.PendingToolCalls = buildPendingToolCalls(pendingExecs, pendingInteractions) + service.rewriteCheckpointTokenDetailsForClient(stream, conversation, state) + return service.broker.Publish(requestID, StreamEvent{ + Message: buildCheckpointMessage(state), + }) } func (service *Service) rewriteCheckpointTokenDetailsForClient(stream *ActiveStream, conversation *ConversationFile, state *agentv1.ConversationStateStructure) { diff --git a/internal/backend/forwarder/token_usage.go b/internal/backend/forwarder/token_usage.go index 20ad0fe..ea3e894 100644 --- a/internal/backend/forwarder/token_usage.go +++ b/internal/backend/forwarder/token_usage.go @@ -45,25 +45,13 @@ func (snapshot turnUsageSnapshot) requestTokensTotal() int64 { return snapshot.promptTokensTotal() + nonNegativeInt64(snapshot.OutputTokens) } -func (service *Service) importConversationState(item *ConversationFile, state *agentv1.ConversationStateStructure, prefetchedBlobs []*agentv1.PreFetchedBlob) ([]HistoryEntry, error) { +func (service *Service) importConversationState(item *ConversationFile, state *agentv1.ConversationStateStructure) ([]HistoryEntry, error) { if item == nil || state == nil { return nil, nil } - blobs, err := newImportedBlobStore(prefetchedBlobs) - if err != nil { - return nil, err - } - importedIDs, err := importedTurnIDs(state.GetTurns(), blobs) - if err != nil { - return nil, err - } item.TokenDetailsUsedTokens = state.GetTokenDetails().GetUsedTokens() - item.ImportedTurnIDs = importedIDs - if minimumNextTurnSeq := int64(len(item.ImportedTurnIDs)) + 1; item.NextTurnSeq < minimumNextTurnSeq { - item.NextTurnSeq = minimumNextTurnSeq - } entries := make([]HistoryEntry, 0, 2) - if messages, err := importedConversationStateModelMessages(state, blobs); err != nil { + if messages, err := importedConversationStateModelMessages(state); err != nil { return nil, err } else { for _, message := range messages { @@ -116,7 +104,7 @@ func (service *Service) importConversationState(item *ConversationFile, state *a return entries, nil } -func importedConversationStateModelMessages(state *agentv1.ConversationStateStructure, blobs importedBlobStore) ([]modeladapter.Message, error) { +func importedConversationStateModelMessages(state *agentv1.ConversationStateStructure) ([]modeladapter.Message, error) { if state == nil { return nil, nil } @@ -125,7 +113,7 @@ func importedConversationStateModelMessages(state *agentv1.ConversationStateStru if err != nil { return nil, fmt.Errorf("decode imported replay messages: %w", err) } - decoded = restoreImportedReplayUserMessages(decoded, state.GetTurns(), blobs) + decoded = restoreImportedReplayUserMessages(decoded, state.GetTurns()) decoded = filterLegacyPlainWriteReplay(decoded) decoded = filterInternalPromptContextReplay(decoded) messages := make([]modeladapter.Message, 0, len(decoded)) @@ -145,18 +133,35 @@ func importedConversationStateModelMessages(state *agentv1.ConversationStateStru if len(rawTurn) == 0 { continue } - turn, turnID, err := decodeImportedTurn(rawTurn, blobs) - if err != nil { - return nil, err + turn := &agentv1.ConversationTurnStructure{} + if err := proto.Unmarshal(rawTurn, turn); err != nil { + return nil, fmt.Errorf("decode imported turn: %w", err) } - if turn == nil && len(turnID) > 0 { - return nil, fmt.Errorf("missing prefetched turn blob %x", turnID) + agentTurn := turn.GetAgentConversationTurn() + if agentTurn == nil { + continue } - turnMessages, err := importedBlobTurnMessages(turn, blobs) - if err != nil { - return nil, err + if rawUser := agentTurn.GetUserMessage(); len(rawUser) > 0 { + userMessage := &agentv1.UserMessage{} + if err := proto.Unmarshal(rawUser, userMessage); err != nil { + return nil, fmt.Errorf("decode imported turn user_message: %w", err) + } + if replay, ok := promptengine.BuildUserMessageReplayMessage(userMessage); ok { + messages = append(messages, toModelMessage(replay)) + } + } + for _, rawStep := range agentTurn.GetSteps() { + if len(rawStep) == 0 { + continue + } + step := &agentv1.ConversationStep{} + if err := proto.Unmarshal(rawStep, step); err != nil { + return nil, fmt.Errorf("decode imported turn step: %w", err) + } + for _, replay := range promptengine.BuildLegacyMessagesFromConversationStep(step) { + messages = append(messages, toModelMessage(replay)) + } } - messages = append(messages, turnMessages...) } return normalizeReplayMessageSequence(messages), nil } diff --git a/internal/backend/forwarder/types.go b/internal/backend/forwarder/types.go index 4b5a6ef..db78eee 100644 --- a/internal/backend/forwarder/types.go +++ b/internal/backend/forwarder/types.go @@ -39,7 +39,6 @@ type ConversationFile struct { CurrentPlanText string `json:"current_plan_text,omitempty"` CurrentPlans map[string]*agentv1.PlanRegistryEntry `json:"current_plans,omitempty"` CurrentTodos []*agentv1.TodoItem `json:"current_todos,omitempty"` - ImportedTurnIDs [][]byte `json:"imported_turn_ids,omitempty"` LatestRequestPrefix *ConversationRequestPrefix `json:"latest_request_prefix,omitempty"` LastProviderCall *ConversationProviderCall `json:"last_provider_call,omitempty"` CreatedAt time.Time `json:"created_at"` @@ -164,11 +163,6 @@ type ActiveStream struct { ProviderUsage turnUsageSnapshot ProviderTerminalToolInvocation bool PendingCompaction *PendingCompaction - PendingCheckpointBlobWrites map[uint32]pendingCheckpointBlobWrite - PendingCheckpointBlobRequests map[string]uint32 - NextCheckpointBlobRequestID uint32 - NextCheckpointRevision uint64 - PendingCheckpoint *pendingCheckpointPublish Backlog []StreamEvent Subscribers map[string]*StreamSubscriber @@ -225,32 +219,6 @@ type pendingTurnCompletion struct { Disposition pendingCompletionDisposition } -type pendingCheckpointBlobWrite struct { - Key string - Revision uint64 -} - -type checkpointTerminalActionKind uint8 - -const ( - checkpointTerminalActionNone checkpointTerminalActionKind = iota - checkpointTerminalActionComplete - checkpointTerminalActionCancel -) - -type checkpointTerminalAction struct { - kind checkpointTerminalActionKind - completion pendingTurnCompletion - cancelMessage string -} - -type pendingCheckpointPublish struct { - Revision uint64 - State *agentv1.ConversationStateStructure - Required map[string]struct{} - TerminalAction checkpointTerminalAction -} - type PendingCompaction struct { Trigger string ContextTokens int64 @@ -450,7 +418,6 @@ type InboundIntent struct { SubagentTypeName string SubagentModelOverrides map[string]runtimecore.SubagentModelOverrideSelection ConversationState *agentv1.ConversationStateStructure - PreFetchedBlobs []*agentv1.PreFetchedBlob UserMessage *agentv1.UserMessage RequestContext *agentv1.RequestContext ClientMessage *agentv1.AgentClientMessage From 06ab0d8daecfe47a60374a0e72665d3271157f1a Mon Sep 17 00:00:00 2001 From: leookun Date: Wed, 5 Aug 2026 01:25:39 +0800 Subject: [PATCH 11/13] chore(release): update version to 0.0.45 and fix conversation disappearance issue - Updated version number to 0.0.45 across all relevant files. - Fixed an issue that could cause conversations to disappear. --- build/config.yml | 2 +- build/darwin/Info.dev.plist | 4 ++-- build/darwin/Info.plist | 4 ++-- build/linux/nfpm/nfpm.yaml | 2 +- build/windows/info.json | 4 ++-- build/windows/nsis/wails_tools.nsh | 2 +- build/windows/wails.exe.manifest | 2 +- release-notes.md | 1 + 8 files changed, 11 insertions(+), 10 deletions(-) diff --git a/build/config.yml b/build/config.yml index c72bd69..16b98d9 100644 --- a/build/config.yml +++ b/build/config.yml @@ -8,7 +8,7 @@ info: description: "Cursor助手" copyright: "© 2026, Cursor助手" comments: "Cursor助手" - version: "0.0.44" + version: "0.0.45" dev_mode: root_path: . diff --git a/build/darwin/Info.dev.plist b/build/darwin/Info.dev.plist index 4ca9f2e..831b455 100644 --- a/build/darwin/Info.dev.plist +++ b/build/darwin/Info.dev.plist @@ -17,9 +17,9 @@ CFBundlePackageType APPL CFBundleShortVersionString - 0.0.44 + 0.0.45 CFBundleVersion - 0.0.44 + 0.0.45 LSMinimumSystemVersion 12.0.0 LSUIElement diff --git a/build/darwin/Info.plist b/build/darwin/Info.plist index 3591e80..56b50c2 100644 --- a/build/darwin/Info.plist +++ b/build/darwin/Info.plist @@ -17,9 +17,9 @@ CFBundlePackageType APPL CFBundleShortVersionString - 0.0.44 + 0.0.45 CFBundleVersion - 0.0.44 + 0.0.45 LSMinimumSystemVersion 12.0.0 LSUIElement diff --git a/build/linux/nfpm/nfpm.yaml b/build/linux/nfpm/nfpm.yaml index 7389f1f..af16591 100644 --- a/build/linux/nfpm/nfpm.yaml +++ b/build/linux/nfpm/nfpm.yaml @@ -6,7 +6,7 @@ name: "Cursor助手" arch: ${GOARCH} platform: "linux" -version: "0.0.44" +version: "0.0.45" section: "default" priority: "extra" maintainer: ${GIT_COMMITTER_NAME} <${GIT_COMMITTER_EMAIL}> diff --git a/build/windows/info.json b/build/windows/info.json index d8c9a4a..115ac3e 100644 --- a/build/windows/info.json +++ b/build/windows/info.json @@ -1,10 +1,10 @@ { "fixed": { - "file_version": "0.0.44" + "file_version": "0.0.45" }, "info": { "0000": { - "ProductVersion": "0.0.44", + "ProductVersion": "0.0.45", "CompanyName": "Cursor助手", "FileDescription": "Cursor助手", "LegalCopyright": "© 2026, Cursor助手", diff --git a/build/windows/nsis/wails_tools.nsh b/build/windows/nsis/wails_tools.nsh index f011a54..eccd696 100644 --- a/build/windows/nsis/wails_tools.nsh +++ b/build/windows/nsis/wails_tools.nsh @@ -14,7 +14,7 @@ !define INFO_PRODUCTNAME "Cursor助手" !endif !ifndef INFO_PRODUCTVERSION - !define INFO_PRODUCTVERSION "0.0.44" + !define INFO_PRODUCTVERSION "0.0.45" !endif !ifndef INFO_COPYRIGHT !define INFO_COPYRIGHT "© 2026, Cursor助手" diff --git a/build/windows/wails.exe.manifest b/build/windows/wails.exe.manifest index 3ed917e..7b30836 100644 --- a/build/windows/wails.exe.manifest +++ b/build/windows/wails.exe.manifest @@ -1,6 +1,6 @@ - + diff --git a/release-notes.md b/release-notes.md index 1c11759..690b827 100644 --- a/release-notes.md +++ b/release-notes.md @@ -9,4 +9,5 @@ Tg群组: https://t.me/cursor_byok - 支持cursor-cli +- 修复对话中错误可能导致的消失问题 From 4e9335d82f0f7a655e02d459ce4716d19b1d0e46 Mon Sep 17 00:00:00 2001 From: leookun Date: Wed, 5 Aug 2026 01:46:47 +0800 Subject: [PATCH 12/13] Implement tests for ProjectLegacyCheckpoint to ensure proper handling of conversation state and message imports. This includes verifying that no dangling inline blobs are present and that the correct user and assistant messages are imported. --- internal/backend/forwarder/projector.go | 164 +----------------- .../backend/forwarder/projector_fork_test.go | 54 ++++++ 2 files changed, 58 insertions(+), 160 deletions(-) create mode 100644 internal/backend/forwarder/projector_fork_test.go diff --git a/internal/backend/forwarder/projector.go b/internal/backend/forwarder/projector.go index 5726f8f..ed6dcb9 100644 --- a/internal/backend/forwarder/projector.go +++ b/internal/backend/forwarder/projector.go @@ -506,158 +506,10 @@ func (projector *HistoryProjector) ProjectLegacyCheckpoint(conversation *Convers if structuredState.HasTodos { state.Todos = encodeConversationTodoBytes(structuredState.Todos) } - grouped := make(map[int64][]HistoryEntry) - order := make([]int64, 0, conversation.NextTurnSeq) - for _, entry := range checkpointProjectionEntries(conversation.Entries) { - if entry.TurnSeq <= 0 { - continue - } - if _, ok := grouped[entry.TurnSeq]; !ok { - order = append(order, entry.TurnSeq) - } - grouped[entry.TurnSeq] = append(grouped[entry.TurnSeq], entry) - } - - for _, turnSeq := range order { - entries := grouped[turnSeq] - var rawUserMessage []byte - var turnRequestID string - steps := make([][]byte, 0, len(entries)) - seenToolCalls := make(map[string]struct{}) - openToolCalls := make(map[string]struct{}) - for _, entry := range entries { - if turnRequestID == "" { - turnRequestID = strings.TrimSpace(entry.RequestID) - } - switch strings.TrimSpace(entry.Kind) { - case "user_message": - userMessage := &agentv1.UserMessage{} - if err := protojson.Unmarshal(entry.Payload, userMessage); err != nil { - return nil, fmt.Errorf("decode checkpoint user_message: %w", err) - } - payload, err := proto.Marshal(userMessage) - if err != nil { - return nil, err - } - rawUserMessage = payload - case "assistant_text": - var payload assistantTextPayload - if err := json.Unmarshal(entry.Payload, &payload); err != nil { - return nil, err - } - if strings.TrimSpace(payload.Text) == "" && strings.TrimSpace(payload.ReasoningContent) != "" && len(openToolCalls) > 0 { - continue - } - if strings.TrimSpace(payload.ReasoningContent) != "" { - stepPayload, err := marshalThinkingStep(payload.ReasoningContent) - if err != nil { - return nil, err - } - steps = append(steps, stepPayload) - } - if strings.TrimSpace(payload.Text) == "" { - continue - } - stepPayload, err := proto.Marshal(&agentv1.ConversationStep{ - Message: &agentv1.ConversationStep_AssistantMessage{ - AssistantMessage: &agentv1.AssistantMessage{Text: strings.TrimSpace(payload.Text)}, - }, - }) - if err != nil { - return nil, err - } - steps = append(steps, stepPayload) - case "tool_call": - var payload toolCallEntryPayload - if err := json.Unmarshal(entry.Payload, &payload); err != nil { - return nil, err - } - if strings.TrimSpace(payload.ReasoningContent) != "" { - stepPayload, err := marshalThinkingStep(payload.ReasoningContent) - if err != nil { - return nil, err - } - steps = append(steps, stepPayload) - } - toolCall := &agentv1.ToolCall{} - if err := protojson.Unmarshal(payload.ToolCall, toolCall); err != nil { - return nil, err - } - if !shouldPersistToolResultName(firstNonEmpty(strings.TrimSpace(payload.ToolName), inferToolName(toolCall))) { - continue - } - stepPayload, err := proto.Marshal(&agentv1.ConversationStep{ - Message: &agentv1.ConversationStep_ToolCall{ - ToolCall: toolCall, - }, - }) - if err != nil { - return nil, err - } - steps = append(steps, stepPayload) - if toolCallID := strings.TrimSpace(payload.ToolCallID); toolCallID != "" { - seenToolCalls[toolCallID] = struct{}{} - openToolCalls[toolCallID] = struct{}{} - } - case "tool_result": - var payload toolResultEntryPayload - if err := json.Unmarshal(entry.Payload, &payload); err != nil { - return nil, err - } - if toolCallID := strings.TrimSpace(payload.ToolCallID); toolCallID != "" { - if _, ok := seenToolCalls[toolCallID]; ok { - delete(openToolCalls, toolCallID) - continue - } - } - if strings.TrimSpace(payload.ReasoningContent) != "" { - stepPayload, err := marshalThinkingStep(payload.ReasoningContent) - if err != nil { - return nil, err - } - steps = append(steps, stepPayload) - } - if len(payload.ToolCall) == 0 { - continue - } - toolCall := &agentv1.ToolCall{} - if err := protojson.Unmarshal(payload.ToolCall, toolCall); err != nil { - return nil, err - } - if !shouldPersistToolResultName(firstNonEmpty(strings.TrimSpace(payload.ToolName), inferToolName(toolCall))) { - continue - } - stepPayload, err := proto.Marshal(&agentv1.ConversationStep{ - Message: &agentv1.ConversationStep_ToolCall{ - ToolCall: toolCall, - }, - }) - if err != nil { - return nil, err - } - steps = append(steps, stepPayload) - } - } - if len(rawUserMessage) == 0 && len(steps) == 0 { - continue - } - agentTurn := &agentv1.AgentConversationTurnStructure{ - UserMessage: rawUserMessage, - Steps: steps, - } - if turnRequestID != "" { - agentTurn.RequestId = &turnRequestID - } - turnPayload, err := proto.Marshal(&agentv1.ConversationTurnStructure{ - Turn: &agentv1.ConversationTurnStructure_AgentConversationTurn{ - AgentConversationTurn: agentTurn, - }, - }) - if err != nil { - return nil, err - } - state.Turns = append(state.Turns, turnPayload) - } + // Cursor 3.14 treats ConversationStateStructure.turns as content-addressed + // blob IDs. Inline protobuf turns make Fork Chat look up the serialized turn + // itself as an ID and fail with "Missing turn blob". Keep model-visible + // history in root_prompt_messages_json until checkpoint blob sync is safe. replayMessages, err := projector.ProjectPromptReplay(conversation) if err != nil { return nil, err @@ -688,14 +540,6 @@ func (projector *HistoryProjector) ProjectLegacyCheckpoint(conversation *Convers return state, nil } -func marshalThinkingStep(text string) ([]byte, error) { - return proto.Marshal(&agentv1.ConversationStep{ - Message: &agentv1.ConversationStep_ThinkingMessage{ - ThinkingMessage: &agentv1.ThinkingMessage{Text: text}, - }, - }) -} - func conversationTokenDetailsUsedTokens(conversation *ConversationFile) uint32 { if conversation == nil { return 0 diff --git a/internal/backend/forwarder/projector_fork_test.go b/internal/backend/forwarder/projector_fork_test.go new file mode 100644 index 0000000..dcc3ed4 --- /dev/null +++ b/internal/backend/forwarder/projector_fork_test.go @@ -0,0 +1,54 @@ +package forwarder + +import ( + "strings" + "testing" + + "google.golang.org/protobuf/encoding/protojson" + + "cursor/gen/agentv1" +) + +func TestProjectLegacyCheckpointKeepsForkContextWithoutDanglingTurnBlobs(t *testing.T) { + userPayload, err := protojson.Marshal(&agentv1.UserMessage{ + Text: "parent question", + MessageId: "message-1", + }) + if err != nil { + t.Fatalf("marshal user message: %v", err) + } + conversation := &ConversationFile{ + ConversationID: "conversation-1", + RootConversationID: "conversation-1", + Mode: "agent", + NextTurnSeq: 2, + NextEntrySeq: 3, + TokenDetailsMaxTokens: projectedConversationMaxTokens, + Entries: []HistoryEntry{ + {Seq: 1, TurnSeq: 1, RequestID: "request-1", Role: "user", Kind: "user_message", Payload: userPayload}, + newAssistantTextEntry(1, "request-1", "parent answer", "", ""), + }, + } + + state, err := NewHistoryProjector().ProjectLegacyCheckpoint(conversation) + if err != nil { + t.Fatalf("ProjectLegacyCheckpoint() error = %v", err) + } + if len(state.GetTurns()) != 0 { + t.Fatalf("ProjectLegacyCheckpoint() turns = %d, want 0 dangling inline blobs", len(state.GetTurns())) + } + + messages, err := importedConversationStateModelMessages(state) + if err != nil { + t.Fatalf("importedConversationStateModelMessages() error = %v", err) + } + if len(messages) != 2 { + t.Fatalf("imported messages = %d, want parent user and assistant context", len(messages)) + } + if messages[0].Role != "user" || !strings.Contains(messages[0].Content, "parent question") { + t.Fatalf("first imported message = %#v", messages[0]) + } + if messages[1].Role != "assistant" || messages[1].Content != "parent answer" { + t.Fatalf("second imported message = %#v", messages[1]) + } +} From d3adfffbd8c0a93a86aa633903b5d7bc7d51d031 Mon Sep 17 00:00:00 2001 From: leookun Date: Wed, 5 Aug 2026 02:09:45 +0800 Subject: [PATCH 13/13] Implement checkpoint blob handling in forwarder service - Added support for checkpointing phases and blob management in the forwarder. - Introduced new types and methods for handling checkpoint blobs, including queuing and publishing checkpoints. - Enhanced the projector to build checkpoint projections with content-addressed blobs. - Implemented tests to ensure proper checkpoint blob synchronization and handling of cancellation scenarios. --- internal/backend/forwarder/actor.go | 7 + internal/backend/forwarder/broker.go | 8 + .../backend/forwarder/checkpoint_blobs.go | 238 +++++++++++++++++ .../forwarder/checkpoint_blobs_test.go | 208 +++++++++++++++ internal/backend/forwarder/events.go | 16 ++ internal/backend/forwarder/projector.go | 245 +++++++++++++++++- .../backend/forwarder/projector_fork_test.go | 96 ++++++- internal/backend/forwarder/service.go | 33 ++- internal/backend/forwarder/types.go | 10 + 9 files changed, 838 insertions(+), 23 deletions(-) create mode 100644 internal/backend/forwarder/checkpoint_blobs.go create mode 100644 internal/backend/forwarder/checkpoint_blobs_test.go diff --git a/internal/backend/forwarder/actor.go b/internal/backend/forwarder/actor.go index 2614cc1..a712976 100644 --- a/internal/backend/forwarder/actor.go +++ b/internal/backend/forwarder/actor.go @@ -23,6 +23,7 @@ const ( TurnPhaseWaitingExternal TurnPhase = "waiting_external" TurnPhaseAwaitingUser TurnPhase = "awaiting_user" TurnPhaseCompacting TurnPhase = "compacting" + TurnPhaseCheckpointing TurnPhase = "checkpointing" TurnPhaseCompleted TurnPhase = "completed" TurnPhaseFailed TurnPhase = "failed" TurnPhaseCanceled TurnPhase = "canceled" @@ -66,6 +67,7 @@ const ( streamTimerNonStreamingRecovery streamTimerKind = "non_streaming_recovery" streamTimerShellForeground streamTimerKind = "shell_foreground" streamTimerShellTransportClose streamTimerKind = "shell_transport_close" + streamTimerCheckpointBlobs streamTimerKind = "checkpoint_blobs" streamTimerOrphanCancel streamTimerKind = "orphan_cancel" ) @@ -318,6 +320,9 @@ func (service *Service) handleStreamCommand(stream *ActiveStream, command stream case streamCommandCancel: return service.handleCancelIntent(command.Intent) case streamCommandMetadata: + if strings.TrimSpace(command.Intent.Kind) == "kv_result" { + return service.handleCheckpointBlobResult(stream, command.Intent.KVClientMessage) + } return service.handleMetadataIntent(command.Intent) case streamCommandExecResult: return service.handleExecResult(command.Intent) @@ -1003,6 +1008,8 @@ func (service *Service) handleTimerEvent(stream *ActiveStream, payload *streamTi return nil } return service.recoverShellWithoutTerminal(stream, current, shellRecoveryReasonTransportClosed) + case streamTimerCheckpointBlobs: + return service.handleCheckpointBlobTimeout(stream) case streamTimerOrphanCancel: stream.mu.Lock() subscriberCount := len(stream.Subscribers) diff --git a/internal/backend/forwarder/broker.go b/internal/backend/forwarder/broker.go index b88442c..7a0f1ce 100644 --- a/internal/backend/forwarder/broker.go +++ b/internal/backend/forwarder/broker.go @@ -76,6 +76,12 @@ func (broker *StreamBroker) OpenStream(requestID string, conversationID string, if existing.BackgroundShellActions == nil { existing.BackgroundShellActions = make(map[string]time.Time) } + if existing.PendingCheckpointBlobWrites == nil { + existing.PendingCheckpointBlobWrites = make(map[uint32]string) + } + if existing.ConfirmedCheckpointBlobs == nil { + existing.ConfirmedCheckpointBlobs = make(map[string]struct{}) + } existing.UpdatedAt = time.Now().UTC() existing.mu.Unlock() return existing, nil @@ -102,6 +108,8 @@ func (broker *StreamBroker) OpenStream(requestID string, conversationID string, BackgroundShellsByMessageID: make(map[uint32]string), BackgroundShellsByExecID: make(map[string]string), BackgroundShellActions: make(map[string]time.Time), + PendingCheckpointBlobWrites: make(map[uint32]string), + ConfirmedCheckpointBlobs: make(map[string]struct{}), CreatedAt: now, UpdatedAt: now, } diff --git a/internal/backend/forwarder/checkpoint_blobs.go b/internal/backend/forwarder/checkpoint_blobs.go new file mode 100644 index 0000000..c07c07f --- /dev/null +++ b/internal/backend/forwarder/checkpoint_blobs.go @@ -0,0 +1,238 @@ +package forwarder + +import ( + "encoding/hex" + "fmt" + "log" + "strings" + "time" + + "google.golang.org/protobuf/proto" + + "cursor/gen/agentv1" +) + +const checkpointBlobWriteTimeout = 5 * time.Second + +type pendingCheckpointBlobWrite struct { + requestID uint32 + blob CheckpointBlob +} + +func clonePendingTurnCompletion(completion *pendingTurnCompletion) *pendingTurnCompletion { + if completion == nil { + return nil + } + cloned := *completion + return &cloned +} + +func (service *Service) queueCheckpointProjection(stream *ActiveStream, projection *CheckpointProjection, completion *pendingTurnCompletion) error { + if service == nil || stream == nil || projection == nil || projection.State == nil { + return nil + } + state, ok := proto.Clone(projection.State).(*agentv1.ConversationStateStructure) + if !ok || state == nil { + return fmt.Errorf("clone checkpoint state") + } + + stream.mu.Lock() + if stream.PendingCheckpointBlobWrites == nil { + stream.PendingCheckpointBlobWrites = make(map[uint32]string) + } + if stream.ConfirmedCheckpointBlobs == nil { + stream.ConfirmedCheckpointBlobs = make(map[string]struct{}) + } + if completion == nil && stream.PendingCheckpoint != nil { + completion = stream.PendingCheckpoint.Completion + } + required := make(map[string]struct{}, len(projection.Blobs)) + pendingKeys := make(map[string]struct{}, len(stream.PendingCheckpointBlobWrites)) + for _, key := range stream.PendingCheckpointBlobWrites { + pendingKeys[key] = struct{}{} + } + toWrite := make([]pendingCheckpointBlobWrite, 0, len(projection.Blobs)) + for _, blob := range projection.Blobs { + key := string(blob.ID) + if key == "" { + continue + } + required[key] = struct{}{} + if _, confirmed := stream.ConfirmedCheckpointBlobs[key]; confirmed { + continue + } + if _, pending := pendingKeys[key]; pending { + continue + } + stream.NextCheckpointBlobRequestID++ + if stream.NextCheckpointBlobRequestID == 0 { + stream.NextCheckpointBlobRequestID++ + } + requestID := stream.NextCheckpointBlobRequestID + stream.PendingCheckpointBlobWrites[requestID] = key + pendingKeys[key] = struct{}{} + toWrite = append(toWrite, pendingCheckpointBlobWrite{requestID: requestID, blob: blob}) + } + stream.PendingCheckpoint = &pendingCheckpointPublish{ + State: state, + Required: required, + Completion: clonePendingTurnCompletion(completion), + } + if completion != nil { + stream.Phase = TurnPhaseCheckpointing + } + stream.UpdatedAt = time.Now().UTC() + stream.mu.Unlock() + + for _, write := range toWrite { + if err := service.broker.Publish(stream.RequestID, StreamEvent{ + Message: buildSetCheckpointBlobMessage(write.requestID, write.blob), + }); err != nil { + return service.finishAfterCheckpointSyncFailure(stream, fmt.Errorf("publish checkpoint blob: %w", err)) + } + } + if service.checkpointProjectionReady(stream) { + return service.publishReadyCheckpoint(stream) + } + service.scheduleStreamTimer( + stream, + providerTimerKey(streamTimerCheckpointBlobs, ""), + checkpointBlobWriteTimeout, + streamTimerCheckpointBlobs, + "", + 0, + "checkpoint blob write timeout", + ) + return nil +} + +func (service *Service) checkpointProjectionReady(stream *ActiveStream) bool { + if stream == nil { + return false + } + stream.mu.Lock() + defer stream.mu.Unlock() + if stream.PendingCheckpoint == nil { + return false + } + for key := range stream.PendingCheckpoint.Required { + if _, confirmed := stream.ConfirmedCheckpointBlobs[key]; !confirmed { + return false + } + } + return true +} + +func (service *Service) handleCheckpointBlobResult(stream *ActiveStream, message *agentv1.KvClientMessage) error { + if service == nil || stream == nil || message == nil || message.GetSetBlobResult() == nil { + return nil + } + stream.mu.Lock() + key, ok := stream.PendingCheckpointBlobWrites[message.GetId()] + if ok { + delete(stream.PendingCheckpointBlobWrites, message.GetId()) + } + required := false + if ok && stream.PendingCheckpoint != nil { + _, required = stream.PendingCheckpoint.Required[key] + } + if ok && message.GetSetBlobResult().GetError() == nil { + stream.ConfirmedCheckpointBlobs[key] = struct{}{} + } + stream.UpdatedAt = time.Now().UTC() + stream.mu.Unlock() + if !ok { + return nil + } + if blobErr := message.GetSetBlobResult().GetError(); blobErr != nil && required { + return service.finishAfterCheckpointSyncFailure(stream, fmt.Errorf( + "client rejected checkpoint blob %s: %s", + hex.EncodeToString([]byte(key)), + firstNonEmpty(strings.TrimSpace(blobErr.GetMessage()), "unknown error"), + )) + } + if service.checkpointProjectionReady(stream) { + return service.publishReadyCheckpoint(stream) + } + return nil +} + +func (service *Service) publishReadyCheckpoint(stream *ActiveStream) error { + if service == nil || stream == nil { + return nil + } + stream.mu.Lock() + pending := stream.PendingCheckpoint + if pending == nil { + stream.mu.Unlock() + return nil + } + for key := range pending.Required { + if _, confirmed := stream.ConfirmedCheckpointBlobs[key]; !confirmed { + stream.mu.Unlock() + return nil + } + } + stream.PendingCheckpoint = nil + state := pending.State + completion := clonePendingTurnCompletion(pending.Completion) + stream.UpdatedAt = time.Now().UTC() + stream.mu.Unlock() + clearStreamTimer(stream, providerTimerKey(streamTimerCheckpointBlobs, "")) + if err := service.broker.Publish(stream.RequestID, StreamEvent{Message: buildCheckpointMessage(state)}); err != nil { + if completion != nil { + log.Printf("forwarder checkpoint publish skipped before successful terminal request_id=%s err=%v", stream.RequestID, err) + return service.finishSuccessfulTurnAfterCheckpoint(stream, *completion) + } + return err + } + if completion != nil { + return service.finishSuccessfulTurnAfterCheckpoint(stream, *completion) + } + return nil +} + +func (service *Service) handleCheckpointBlobTimeout(stream *ActiveStream) error { + if stream == nil { + return nil + } + stream.mu.Lock() + pendingCount := len(stream.PendingCheckpointBlobWrites) + stream.mu.Unlock() + return service.finishAfterCheckpointSyncFailure(stream, fmt.Errorf("%d checkpoint blob writes timed out", pendingCount)) +} + +func (service *Service) finishAfterCheckpointSyncFailure(stream *ActiveStream, cause error) error { + if stream == nil { + return nil + } + stream.mu.Lock() + pending := stream.PendingCheckpoint + stream.PendingCheckpoint = nil + stream.PendingCheckpointBlobWrites = make(map[uint32]string) + stream.UpdatedAt = time.Now().UTC() + stream.mu.Unlock() + clearStreamTimer(stream, providerTimerKey(streamTimerCheckpointBlobs, "")) + if cause != nil { + log.Printf("forwarder checkpoint blob sync skipped request_id=%s conversation_id=%s err=%v", stream.RequestID, stream.ConversationID, cause) + } + if pending != nil && pending.Completion != nil { + return service.finishSuccessfulTurnAfterCheckpoint(stream, *pending.Completion) + } + return nil +} + +func (service *Service) discardPendingCheckpoint(stream *ActiveStream, reason string) { + if stream == nil { + return + } + stream.mu.Lock() + stream.PendingCheckpoint = nil + stream.PendingCheckpointBlobWrites = make(map[uint32]string) + stream.UpdatedAt = time.Now().UTC() + stream.mu.Unlock() + clearStreamTimer(stream, providerTimerKey(streamTimerCheckpointBlobs, "")) + if strings.TrimSpace(reason) != "" { + log.Printf("forwarder pending checkpoint discarded request_id=%s conversation_id=%s reason=%s", stream.RequestID, stream.ConversationID, strings.TrimSpace(reason)) + } +} diff --git a/internal/backend/forwarder/checkpoint_blobs_test.go b/internal/backend/forwarder/checkpoint_blobs_test.go new file mode 100644 index 0000000..9ada5d3 --- /dev/null +++ b/internal/backend/forwarder/checkpoint_blobs_test.go @@ -0,0 +1,208 @@ +package forwarder + +import ( + "testing" + + "google.golang.org/protobuf/encoding/protojson" + + "cursor/gen/agentv1" +) + +func TestCheckpointBlobSyncPublishesCheckpointAfterAcknowledgements(t *testing.T) { + service, stream, projection := testCheckpointBlobProjection(t) + if err := service.queueCheckpointProjection(stream, projection, nil); err != nil { + t.Fatalf("queueCheckpointProjection() error = %v", err) + } + events := readCheckpointTestEvents(t, service, stream) + if len(events) != len(projection.Blobs) { + t.Fatalf("events before ACK = %d, want %d Blob writes", len(events), len(projection.Blobs)) + } + for _, event := range events { + if event.Message.GetKvServerMessage().GetSetBlobArgs() == nil { + t.Fatalf("event before ACK = %#v, want set_blob_args", event.Message) + } + } + + acknowledgeCheckpointBlobs(t, service, stream) + events = readCheckpointTestEvents(t, service, stream) + checkpoint := events[len(events)-1].Message.GetConversationCheckpointUpdate() + if checkpoint == nil || len(checkpoint.GetTurns()) != 1 { + t.Fatalf("last event checkpoint = %#v, want one Blob-backed turn", checkpoint) + } +} + +func TestCheckpointBlobSyncPublishesCheckpointBeforeSuccessfulTerminal(t *testing.T) { + service, stream, projection := testCheckpointBlobProjection(t) + completion := &pendingTurnCompletion{ + RequestID: stream.RequestID, + Usage: turnUsageSnapshot{InputTokens: 11, OutputTokens: 7}, + } + if err := service.queueCheckpointProjection(stream, projection, completion); err != nil { + t.Fatalf("queueCheckpointProjection() error = %v", err) + } + acknowledgeCheckpointBlobs(t, service, stream) + + events := readCheckpointTestEvents(t, service, stream) + checkpointIndex, turnEndedIndex, endIndex := -1, -1, -1 + for index, event := range events { + switch { + case event.Message.GetConversationCheckpointUpdate() != nil: + checkpointIndex = index + case event.Message.GetInteractionUpdate().GetTurnEnded() != nil: + turnEndedIndex = index + case event.End: + endIndex = index + } + } + if checkpointIndex < 0 || turnEndedIndex <= checkpointIndex || endIndex <= turnEndedIndex { + t.Fatalf("terminal order checkpoint=%d turn_ended=%d end=%d", checkpointIndex, turnEndedIndex, endIndex) + } +} + +func TestCheckpointBlobTimeoutDoesNotFailSuccessfulTurn(t *testing.T) { + service, stream, projection := testCheckpointBlobProjection(t) + completion := &pendingTurnCompletion{ + RequestID: stream.RequestID, + Usage: turnUsageSnapshot{InputTokens: 11, OutputTokens: 7}, + } + if err := service.queueCheckpointProjection(stream, projection, completion); err != nil { + t.Fatalf("queueCheckpointProjection() error = %v", err) + } + if err := service.handleCheckpointBlobTimeout(stream); err != nil { + t.Fatalf("handleCheckpointBlobTimeout() error = %v", err) + } + + events := readCheckpointTestEvents(t, service, stream) + var checkpoint, turnEnded, successfulEnd bool + for _, event := range events { + checkpoint = checkpoint || event.Message.GetConversationCheckpointUpdate() != nil + turnEnded = turnEnded || event.Message.GetInteractionUpdate().GetTurnEnded() != nil + successfulEnd = successfulEnd || event.End && event.TerminalErrorCode == "" + } + if checkpoint || !turnEnded || !successfulEnd { + t.Fatalf("timeout events checkpoint=%v turn_ended=%v successful_end=%v", checkpoint, turnEnded, successfulEnd) + } +} + +func TestCancellationDiscardsPendingCheckpointAndIgnoresLateAcknowledgements(t *testing.T) { + service, stream, projection := testCheckpointBlobProjection(t) + if err := service.queueCheckpointProjection(stream, projection, nil); err != nil { + t.Fatalf("queueCheckpointProjection() error = %v", err) + } + stream.mu.Lock() + requestIDs := make([]uint32, 0, len(stream.PendingCheckpointBlobWrites)) + for requestID := range stream.PendingCheckpointBlobWrites { + requestIDs = append(requestIDs, requestID) + } + stream.mu.Unlock() + if err := service.handleCancelIntent(InboundIntent{ + Kind: "cancel", + RequestID: stream.RequestID, + CancelReason: "user stopped", + }); err != nil { + t.Fatalf("handleCancelIntent() error = %v", err) + } + for _, requestID := range requestIDs { + if err := service.handleCheckpointBlobResult(stream, &agentv1.KvClientMessage{ + Id: requestID, + Message: &agentv1.KvClientMessage_SetBlobResult{ + SetBlobResult: &agentv1.SetBlobResult{}, + }, + }); err != nil { + t.Fatalf("late ACK %d error = %v", requestID, err) + } + } + + events := readCheckpointTestEvents(t, service, stream) + var checkpoint, canceledEnd bool + for _, event := range events { + checkpoint = checkpoint || event.Message.GetConversationCheckpointUpdate() != nil + canceledEnd = canceledEnd || event.End && event.TerminalErrorCode == "canceled" + } + stream.mu.Lock() + pending := stream.PendingCheckpoint + stream.mu.Unlock() + if checkpoint || !canceledEnd || pending != nil { + t.Fatalf("cancel events checkpoint=%v canceled_end=%v pending=%v", checkpoint, canceledEnd, pending != nil) + } +} + +func testCheckpointBlobProjection(t *testing.T) (*Service, *ActiveStream, *CheckpointProjection) { + t.Helper() + broker := NewStreamBroker() + service := &Service{ + store: NewConversationFileStore(t.TempDir()), + projector: NewHistoryProjector(), + broker: broker, + } + stream, err := broker.OpenStream( + "request-1", "conversation-1", 1, "default", "default", + agentv1.AgentMode_AGENT_MODE_AGENT, "hello", + ) + if err != nil { + t.Fatalf("OpenStream() error = %v", err) + } + conversation := &ConversationFile{ + ConversationID: "conversation-1", + RootConversationID: "conversation-1", + Mode: "agent", + NextTurnSeq: 2, + NextEntrySeq: 3, + TokenDetailsMaxTokens: projectedConversationMaxTokens, + Entries: []HistoryEntry{ + testCheckpointUserEntry(t), + newAssistantTextEntry(1, "request-1", "hi", "", ""), + }, + } + projection, err := service.projector.ProjectCheckpointProjection(conversation) + if err != nil { + t.Fatalf("ProjectCheckpointProjection() error = %v", err) + } + if err := service.replaceCheckpointConversation(stream, conversation); err != nil { + t.Fatalf("replaceCheckpointConversation() error = %v", err) + } + return service, stream, projection +} + +func testCheckpointUserEntry(t *testing.T) HistoryEntry { + t.Helper() + payload, err := protojson.Marshal(&agentv1.UserMessage{Text: "hello", MessageId: "message-1"}) + if err != nil { + t.Fatalf("marshal user message: %v", err) + } + return HistoryEntry{Seq: 1, TurnSeq: 1, RequestID: "request-1", Role: "user", Kind: "user_message", Payload: payload} +} + +func acknowledgeCheckpointBlobs(t *testing.T, service *Service, stream *ActiveStream) { + t.Helper() + for { + stream.mu.Lock() + requestIDs := make([]uint32, 0, len(stream.PendingCheckpointBlobWrites)) + for requestID := range stream.PendingCheckpointBlobWrites { + requestIDs = append(requestIDs, requestID) + } + stream.mu.Unlock() + if len(requestIDs) == 0 { + return + } + for _, requestID := range requestIDs { + if err := service.handleCheckpointBlobResult(stream, &agentv1.KvClientMessage{ + Id: requestID, + Message: &agentv1.KvClientMessage_SetBlobResult{ + SetBlobResult: &agentv1.SetBlobResult{}, + }, + }); err != nil { + t.Fatalf("handleCheckpointBlobResult(%d) error = %v", requestID, err) + } + } + } +} + +func readCheckpointTestEvents(t *testing.T, service *Service, stream *ActiveStream) []StreamEvent { + t.Helper() + events, err := service.broker.ReadFromCursor(stream.RequestID, 0) + if err != nil { + t.Fatalf("ReadFromCursor() error = %v", err) + } + return events +} diff --git a/internal/backend/forwarder/events.go b/internal/backend/forwarder/events.go index 7b02404..ef69776 100644 --- a/internal/backend/forwarder/events.go +++ b/internal/backend/forwarder/events.go @@ -245,6 +245,22 @@ func buildCheckpointMessage(state *agentv1.ConversationStateStructure) *agentv1. } } +func buildSetCheckpointBlobMessage(id uint32, blob CheckpointBlob) *agentv1.AgentServerMessage { + return &agentv1.AgentServerMessage{ + Message: &agentv1.AgentServerMessage_KvServerMessage{ + KvServerMessage: &agentv1.KvServerMessage{ + Id: id, + Message: &agentv1.KvServerMessage_SetBlobArgs{ + SetBlobArgs: &agentv1.SetBlobArgs{ + BlobId: append([]byte(nil), blob.ID...), + BlobData: append([]byte(nil), blob.Data...), + }, + }, + }, + }, + } +} + // buildExecAbortMessage 构造对客户端执行桥的 abort 控制消息。 func buildExecAbortMessage(pending runtimecore.PendingExec) *agentv1.AgentServerMessage { return &agentv1.AgentServerMessage{ diff --git a/internal/backend/forwarder/projector.go b/internal/backend/forwarder/projector.go index ed6dcb9..7095860 100644 --- a/internal/backend/forwarder/projector.go +++ b/internal/backend/forwarder/projector.go @@ -2,6 +2,7 @@ package forwarder import ( + "crypto/sha256" "encoding/json" "fmt" "strings" @@ -19,6 +20,51 @@ const projectedConversationMaxTokens = 130000 type HistoryProjector struct { } +type CheckpointBlob struct { + ID []byte + Data []byte +} + +type CheckpointProjection struct { + State *agentv1.ConversationStateStructure + Blobs []CheckpointBlob +} + +type checkpointBlobGraph struct { + blobs map[[sha256.Size]byte][]byte + order [][sha256.Size]byte +} + +func newCheckpointBlobGraph() *checkpointBlobGraph { + return &checkpointBlobGraph{blobs: make(map[[sha256.Size]byte][]byte)} +} + +func (graph *checkpointBlobGraph) add(data []byte) []byte { + if graph == nil || len(data) == 0 { + return nil + } + id := sha256.Sum256(data) + if _, exists := graph.blobs[id]; !exists { + graph.blobs[id] = append([]byte(nil), data...) + graph.order = append(graph.order, id) + } + return append([]byte(nil), id[:]...) +} + +func (graph *checkpointBlobGraph) list() []CheckpointBlob { + if graph == nil || len(graph.order) == 0 { + return nil + } + blobs := make([]CheckpointBlob, 0, len(graph.order)) + for _, id := range graph.order { + blobs = append(blobs, CheckpointBlob{ + ID: append([]byte(nil), id[:]...), + Data: append([]byte(nil), graph.blobs[id]...), + }) + } + return blobs +} + // NewHistoryProjector 创建 history 投影器。 func NewHistoryProjector() *HistoryProjector { return &HistoryProjector{} @@ -475,6 +521,16 @@ func isHistoricalReplayToolResult(conversation *ConversationFile, entry HistoryE // ProjectLegacyCheckpoint 按需从 JSON history 投影出兼容旧客户端的 checkpoint 结构。 func (projector *HistoryProjector) ProjectLegacyCheckpoint(conversation *ConversationFile) (*agentv1.ConversationStateStructure, error) { + projection, err := projector.ProjectCheckpointProjection(conversation) + if err != nil || projection == nil { + return nil, err + } + return projection.State, nil +} + +// ProjectCheckpointProjection 同时返回 checkpoint 状态及其引用的内容寻址 Blob。 +func (projector *HistoryProjector) ProjectCheckpointProjection(conversation *ConversationFile) (*CheckpointProjection, error) { + blobs := newCheckpointBlobGraph() state := &agentv1.ConversationStateStructure{ TokenDetails: &agentv1.ConversationTokenDetails{ UsedTokens: conversationTokenDetailsUsedTokens(conversation), @@ -488,7 +544,7 @@ func (projector *HistoryProjector) ProjectLegacyCheckpoint(conversation *Convers if conversation == nil { mode := agentv1.AgentMode_AGENT_MODE_AGENT state.Mode = &mode - return state, nil + return &CheckpointProjection{State: state}, nil } mode, err := parseModeAlias(conversation.Mode) if err != nil { @@ -506,10 +562,11 @@ func (projector *HistoryProjector) ProjectLegacyCheckpoint(conversation *Convers if structuredState.HasTodos { state.Todos = encodeConversationTodoBytes(structuredState.Todos) } - // Cursor 3.14 treats ConversationStateStructure.turns as content-addressed - // blob IDs. Inline protobuf turns make Fork Chat look up the serialized turn - // itself as an ID and fail with "Missing turn blob". Keep model-visible - // history in root_prompt_messages_json until checkpoint blob sync is safe. + turnIDs, err := projectCheckpointTurnBlobs(conversation, blobs) + if err != nil { + return nil, err + } + state.Turns = turnIDs replayMessages, err := projector.ProjectPromptReplay(conversation) if err != nil { return nil, err @@ -537,7 +594,183 @@ func (projector *HistoryProjector) ProjectLegacyCheckpoint(conversation *Convers return nil, err } state.RootPromptMessagesJson = rootPromptMessages - return state, nil + return &CheckpointProjection{State: state, Blobs: blobs.list()}, nil +} + +func projectCheckpointTurnBlobs(conversation *ConversationFile, blobs *checkpointBlobGraph) ([][]byte, error) { + if conversation == nil || blobs == nil { + return nil, nil + } + grouped := make(map[int64][]HistoryEntry) + order := make([]int64, 0, conversation.NextTurnSeq) + for _, entry := range checkpointProjectionEntries(conversation.Entries) { + if entry.TurnSeq <= 0 { + continue + } + if _, ok := grouped[entry.TurnSeq]; !ok { + order = append(order, entry.TurnSeq) + } + grouped[entry.TurnSeq] = append(grouped[entry.TurnSeq], entry) + } + + turnIDs := make([][]byte, 0, len(order)) + for _, turnSeq := range order { + entries := grouped[turnSeq] + var userMessageID []byte + var turnRequestID string + stepIDs := make([][]byte, 0, len(entries)) + seenToolCalls := make(map[string]struct{}) + openToolCalls := make(map[string]struct{}) + for _, entry := range entries { + if turnRequestID == "" { + turnRequestID = strings.TrimSpace(entry.RequestID) + } + switch strings.TrimSpace(entry.Kind) { + case "user_message": + userMessage := &agentv1.UserMessage{} + if err := protojson.Unmarshal(entry.Payload, userMessage); err != nil { + return nil, fmt.Errorf("decode checkpoint user_message: %w", err) + } + payload, err := proto.Marshal(userMessage) + if err != nil { + return nil, err + } + userMessageID = blobs.add(payload) + case "assistant_text": + var payload assistantTextPayload + if err := json.Unmarshal(entry.Payload, &payload); err != nil { + return nil, err + } + if strings.TrimSpace(payload.Text) == "" && strings.TrimSpace(payload.ReasoningContent) != "" && len(openToolCalls) > 0 { + continue + } + if strings.TrimSpace(payload.ReasoningContent) != "" { + stepID, err := addCheckpointStepBlob(blobs, &agentv1.ConversationStep{ + Message: &agentv1.ConversationStep_ThinkingMessage{ + ThinkingMessage: &agentv1.ThinkingMessage{Text: payload.ReasoningContent}, + }, + }) + if err != nil { + return nil, err + } + stepIDs = append(stepIDs, stepID) + } + if strings.TrimSpace(payload.Text) == "" { + continue + } + stepID, err := addCheckpointStepBlob(blobs, &agentv1.ConversationStep{ + Message: &agentv1.ConversationStep_AssistantMessage{ + AssistantMessage: &agentv1.AssistantMessage{Text: strings.TrimSpace(payload.Text)}, + }, + }) + if err != nil { + return nil, err + } + stepIDs = append(stepIDs, stepID) + case "tool_call": + var payload toolCallEntryPayload + if err := json.Unmarshal(entry.Payload, &payload); err != nil { + return nil, err + } + if strings.TrimSpace(payload.ReasoningContent) != "" { + stepID, err := addCheckpointStepBlob(blobs, &agentv1.ConversationStep{ + Message: &agentv1.ConversationStep_ThinkingMessage{ + ThinkingMessage: &agentv1.ThinkingMessage{Text: payload.ReasoningContent}, + }, + }) + if err != nil { + return nil, err + } + stepIDs = append(stepIDs, stepID) + } + toolCall := &agentv1.ToolCall{} + if err := protojson.Unmarshal(payload.ToolCall, toolCall); err != nil { + return nil, err + } + if !shouldPersistToolResultName(firstNonEmpty(strings.TrimSpace(payload.ToolName), inferToolName(toolCall))) { + continue + } + stepID, err := addCheckpointStepBlob(blobs, &agentv1.ConversationStep{ + Message: &agentv1.ConversationStep_ToolCall{ToolCall: toolCall}, + }) + if err != nil { + return nil, err + } + stepIDs = append(stepIDs, stepID) + if toolCallID := strings.TrimSpace(payload.ToolCallID); toolCallID != "" { + seenToolCalls[toolCallID] = struct{}{} + openToolCalls[toolCallID] = struct{}{} + } + case "tool_result": + var payload toolResultEntryPayload + if err := json.Unmarshal(entry.Payload, &payload); err != nil { + return nil, err + } + if toolCallID := strings.TrimSpace(payload.ToolCallID); toolCallID != "" { + if _, ok := seenToolCalls[toolCallID]; ok { + delete(openToolCalls, toolCallID) + continue + } + } + if strings.TrimSpace(payload.ReasoningContent) != "" { + stepID, err := addCheckpointStepBlob(blobs, &agentv1.ConversationStep{ + Message: &agentv1.ConversationStep_ThinkingMessage{ + ThinkingMessage: &agentv1.ThinkingMessage{Text: payload.ReasoningContent}, + }, + }) + if err != nil { + return nil, err + } + stepIDs = append(stepIDs, stepID) + } + if len(payload.ToolCall) == 0 { + continue + } + toolCall := &agentv1.ToolCall{} + if err := protojson.Unmarshal(payload.ToolCall, toolCall); err != nil { + return nil, err + } + if !shouldPersistToolResultName(firstNonEmpty(strings.TrimSpace(payload.ToolName), inferToolName(toolCall))) { + continue + } + stepID, err := addCheckpointStepBlob(blobs, &agentv1.ConversationStep{ + Message: &agentv1.ConversationStep_ToolCall{ToolCall: toolCall}, + }) + if err != nil { + return nil, err + } + stepIDs = append(stepIDs, stepID) + } + } + if len(userMessageID) == 0 && len(stepIDs) == 0 { + continue + } + agentTurn := &agentv1.AgentConversationTurnStructure{ + UserMessage: userMessageID, + Steps: stepIDs, + } + if turnRequestID != "" { + agentTurn.RequestId = &turnRequestID + } + turnPayload, err := proto.Marshal(&agentv1.ConversationTurnStructure{ + Turn: &agentv1.ConversationTurnStructure_AgentConversationTurn{ + AgentConversationTurn: agentTurn, + }, + }) + if err != nil { + return nil, err + } + turnIDs = append(turnIDs, blobs.add(turnPayload)) + } + return turnIDs, nil +} + +func addCheckpointStepBlob(blobs *checkpointBlobGraph, step *agentv1.ConversationStep) ([]byte, error) { + payload, err := proto.Marshal(step) + if err != nil { + return nil, err + } + return blobs.add(payload), nil } func conversationTokenDetailsUsedTokens(conversation *ConversationFile) uint32 { diff --git a/internal/backend/forwarder/projector_fork_test.go b/internal/backend/forwarder/projector_fork_test.go index dcc3ed4..56fc0af 100644 --- a/internal/backend/forwarder/projector_fork_test.go +++ b/internal/backend/forwarder/projector_fork_test.go @@ -1,15 +1,17 @@ package forwarder import ( + "crypto/sha256" "strings" "testing" "google.golang.org/protobuf/encoding/protojson" + "google.golang.org/protobuf/proto" "cursor/gen/agentv1" ) -func TestProjectLegacyCheckpointKeepsForkContextWithoutDanglingTurnBlobs(t *testing.T) { +func TestProjectCheckpointProjectionBuildsResolvableForkState(t *testing.T) { userPayload, err := protojson.Marshal(&agentv1.UserMessage{ Text: "parent question", MessageId: "message-1", @@ -30,12 +32,41 @@ func TestProjectLegacyCheckpointKeepsForkContextWithoutDanglingTurnBlobs(t *test }, } - state, err := NewHistoryProjector().ProjectLegacyCheckpoint(conversation) + projection, err := NewHistoryProjector().ProjectCheckpointProjection(conversation) if err != nil { - t.Fatalf("ProjectLegacyCheckpoint() error = %v", err) + t.Fatalf("ProjectCheckpointProjection() error = %v", err) } - if len(state.GetTurns()) != 0 { - t.Fatalf("ProjectLegacyCheckpoint() turns = %d, want 0 dangling inline blobs", len(state.GetTurns())) + state := projection.State + if len(state.GetTurns()) != 1 { + t.Fatalf("ProjectCheckpointProjection() turns = %d, want 1 Blob-backed turn", len(state.GetTurns())) + } + blobs := make(map[string][]byte, len(projection.Blobs)) + for _, blob := range projection.Blobs { + digest := sha256.Sum256(blob.Data) + if len(blob.ID) != sha256.Size || string(blob.ID) != string(digest[:]) { + t.Fatalf("invalid content-addressed Blob id=%x", blob.ID) + } + blobs[string(blob.ID)] = blob.Data + } + turnPayload, ok := blobs[string(state.GetTurns()[0])] + if !ok { + t.Fatal("turn references a missing Blob") + } + turn := &agentv1.ConversationTurnStructure{} + if err := proto.Unmarshal(turnPayload, turn); err != nil { + t.Fatalf("decode turn Blob: %v", err) + } + agentTurn := turn.GetAgentConversationTurn() + if agentTurn == nil { + t.Fatal("turn Blob does not contain an agent turn") + } + if _, ok := blobs[string(agentTurn.GetUserMessage())]; !ok { + t.Fatal("turn references a missing user message Blob") + } + for _, stepID := range agentTurn.GetSteps() { + if _, ok := blobs[string(stepID)]; !ok { + t.Fatal("turn references a missing step Blob") + } } messages, err := importedConversationStateModelMessages(state) @@ -52,3 +83,58 @@ func TestProjectLegacyCheckpointKeepsForkContextWithoutDanglingTurnBlobs(t *test t.Fatalf("second imported message = %#v", messages[1]) } } + +func TestProjectCheckpointProjectionKeepsForkPointIsolatedFromLaterHistory(t *testing.T) { + firstUser, err := protojson.Marshal(&agentv1.UserMessage{Text: "first question", MessageId: "message-1"}) + if err != nil { + t.Fatalf("marshal first user message: %v", err) + } + conversation := &ConversationFile{ + ConversationID: "conversation-1", + RootConversationID: "conversation-1", + Mode: "agent", + NextTurnSeq: 2, + NextEntrySeq: 3, + TokenDetailsMaxTokens: projectedConversationMaxTokens, + Entries: []HistoryEntry{ + {Seq: 1, TurnSeq: 1, RequestID: "request-1", Role: "user", Kind: "user_message", Payload: firstUser}, + newAssistantTextEntry(1, "request-1", "first answer", "", ""), + }, + } + projector := NewHistoryProjector() + midpoint, err := projector.ProjectCheckpointProjection(conversation) + if err != nil { + t.Fatalf("midpoint projection: %v", err) + } + + secondUser, err := protojson.Marshal(&agentv1.UserMessage{Text: "second question", MessageId: "message-2"}) + if err != nil { + t.Fatalf("marshal second user message: %v", err) + } + appendEntriesInPlace(conversation, []HistoryEntry{ + {TurnSeq: 2, RequestID: "request-2", Role: "user", Kind: "user_message", Payload: secondUser}, + newAssistantTextEntry(2, "request-2", "second answer", "", ""), + }) + latest, err := projector.ProjectCheckpointProjection(conversation) + if err != nil { + t.Fatalf("latest projection: %v", err) + } + + midpointMessages, err := importedConversationStateModelMessages(midpoint.State) + if err != nil { + t.Fatalf("import midpoint messages: %v", err) + } + latestMessages, err := importedConversationStateModelMessages(latest.State) + if err != nil { + t.Fatalf("import latest messages: %v", err) + } + if len(midpoint.State.GetTurns()) != 1 || len(midpointMessages) != 2 { + t.Fatalf("midpoint turns=%d messages=%d, want 1 turn and 2 messages", len(midpoint.State.GetTurns()), len(midpointMessages)) + } + if len(latest.State.GetTurns()) != 2 || len(latestMessages) != 4 { + t.Fatalf("latest turns=%d messages=%d, want 2 turns and 4 messages", len(latest.State.GetTurns()), len(latestMessages)) + } + if midpointMessages[1].Content != "first answer" || latestMessages[3].Content != "second answer" { + t.Fatalf("fork snapshots are not isolated: midpoint=%#v latest=%#v", midpointMessages, latestMessages) + } +} diff --git a/internal/backend/forwarder/service.go b/internal/backend/forwarder/service.go index 5bb04ca..4cf5c5d 100644 --- a/internal/backend/forwarder/service.go +++ b/internal/backend/forwarder/service.go @@ -858,9 +858,7 @@ func (service *Service) handleCancelIntent(intent InboundIntent) error { }) } if hasCheckpoint { - if err := service.publishCheckpoint(stream.RequestID, stream.ConversationID); err != nil { - return err - } + service.discardPendingCheckpoint(stream, "checkpoint superseded by cancellation") } clearPendingProviderCompletion(stream) stream.mu.Lock() @@ -2122,9 +2120,15 @@ func (service *Service) completeSuccessfulTurn(stream *ActiveStream, completion err, ) } - if err := service.publishCheckpoint(requestID, conversationID); err != nil { - return err + return service.publishCheckpointWithCompletion(requestID, conversationID, &completion) +} + +func (service *Service) finishSuccessfulTurnAfterCheckpoint(stream *ActiveStream, completion pendingTurnCompletion) error { + if stream == nil { + return nil } + requestID := firstNonEmpty(strings.TrimSpace(completion.RequestID), strings.TrimSpace(stream.RequestID)) + usage := completion.Usage if err := service.broker.Publish(requestID, StreamEvent{ Message: buildTurnEndedMessage(usage.InputTokens, usage.OutputTokens, usage.CacheReadTokens, usage.CacheWriteTokens), }); err != nil { @@ -2151,7 +2155,11 @@ func (service *Service) failStreamIfNonTerminal(stream *ActiveStream, terminalCo } // publishCheckpoint 按当前内存会话镜像投影出 checkpoint,并广播给所有 RunSSE 订阅者。 -func (service *Service) publishCheckpoint(requestID string, _ string) error { +func (service *Service) publishCheckpoint(requestID string, conversationID string) error { + return service.publishCheckpointWithCompletion(requestID, conversationID, nil) +} + +func (service *Service) publishCheckpointWithCompletion(requestID string, _ string, completion *pendingTurnCompletion) error { stream, ok := service.broker.Get(requestID) if !ok || stream == nil { return fmt.Errorf("request is not active: %s", requestID) @@ -2160,15 +2168,16 @@ func (service *Service) publishCheckpoint(requestID string, _ string) error { if err != nil { return err } - state, err := service.projector.ProjectLegacyCheckpoint(conversation) + projection, err := service.projector.ProjectCheckpointProjection(conversation) if err != nil { return err } - state.PendingToolCalls = buildPendingToolCalls(pendingExecs, pendingInteractions) - service.rewriteCheckpointTokenDetailsForClient(stream, conversation, state) - return service.broker.Publish(requestID, StreamEvent{ - Message: buildCheckpointMessage(state), - }) + if projection == nil || projection.State == nil { + return fmt.Errorf("checkpoint projection is empty") + } + projection.State.PendingToolCalls = buildPendingToolCalls(pendingExecs, pendingInteractions) + service.rewriteCheckpointTokenDetailsForClient(stream, conversation, projection.State) + return service.queueCheckpointProjection(stream, projection, completion) } func (service *Service) rewriteCheckpointTokenDetailsForClient(stream *ActiveStream, conversation *ConversationFile, state *agentv1.ConversationStateStructure) { diff --git a/internal/backend/forwarder/types.go b/internal/backend/forwarder/types.go index db78eee..7b33d20 100644 --- a/internal/backend/forwarder/types.go +++ b/internal/backend/forwarder/types.go @@ -163,6 +163,10 @@ type ActiveStream struct { ProviderUsage turnUsageSnapshot ProviderTerminalToolInvocation bool PendingCompaction *PendingCompaction + PendingCheckpointBlobWrites map[uint32]string + ConfirmedCheckpointBlobs map[string]struct{} + NextCheckpointBlobRequestID uint32 + PendingCheckpoint *pendingCheckpointPublish Backlog []StreamEvent Subscribers map[string]*StreamSubscriber @@ -219,6 +223,12 @@ type pendingTurnCompletion struct { Disposition pendingCompletionDisposition } +type pendingCheckpointPublish struct { + State *agentv1.ConversationStateStructure + Required map[string]struct{} + Completion *pendingTurnCompletion +} + type PendingCompaction struct { Trigger string ContextTokens int64