mirror of
https://wget.la/https://github.com/leookun/cursor-byok
synced 2026-08-18 03:57:06 +08:00
Merge branch 'main' into feat/cursor-local-history-transcripts
This commit is contained in:
@@ -1,6 +1,6 @@
|
||||
# Backend 架构说明
|
||||
|
||||
`internal/backend` 当前支持本地助手模式与直连上游模式。
|
||||
`internal/backend` 只支持本地助手模式。
|
||||
|
||||
关于 backend agent「最小事实集合」的第一阶段研究文档,见 [`../../docs/backend-agent-minimum-facts-phase1.md`](../../docs/backend-agent-minimum-facts-phase1.md)。
|
||||
|
||||
@@ -34,7 +34,6 @@ internal/backend/
|
||||
errors.go
|
||||
local.go
|
||||
middleware.go
|
||||
policy.go
|
||||
route.go
|
||||
url.go
|
||||
|
||||
@@ -134,7 +133,7 @@ history/
|
||||
## 请求流
|
||||
|
||||
1. 请求进入 backend 根路由。
|
||||
2. `PolicyMiddleware` 根据 `routing.mode` 与 `X-Server-Upstream-URL` 选择本地或上游分支。
|
||||
2. `ServerContext` 解析 MITM 带入的原始目标地址,路由始终执行本地 action。
|
||||
3. `BidiAppend` / `RunSSE` 进入 `forwarder`。
|
||||
4. `forwarder` 先把当前 loop 状态写入 `state.json`,再把已发生语义事件追加到 `context.json`。
|
||||
5. 发给 LLM 的 prompt 只由 `context.json` 投射生成;`state.json` 不保存可投射历史。
|
||||
|
||||
@@ -524,7 +524,7 @@ func (service *Service) applyProviderModelEvent(stream *ActiveStream, event mode
|
||||
stream.mu.Unlock()
|
||||
}
|
||||
if shouldEmitSyntheticThinking {
|
||||
if err := service.broker.Publish(requestID, StreamEvent{Message: buildThinkingDeltaMessage("Thinking is encrypted. Please wait a moment.", event.ThinkingStyle)}); err != nil {
|
||||
if err := service.broker.Publish(requestID, StreamEvent{Message: buildThinkingDeltaMessage("The reasoning process is encrypted. Please wait a moment. (This message does not affect any functionality; it only indicates the current reasoning status.)", event.ThinkingStyle)}); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
+121
-178
@@ -27,10 +27,11 @@ const healthPath = "/healthz"
|
||||
const tabServerBaseURL = "https://tab.leokun.cn"
|
||||
|
||||
type Host struct {
|
||||
store *serverconfig.Store
|
||||
listenAddr string
|
||||
configs *serverconfig.Manager
|
||||
healthHTTP *http.Client
|
||||
store *serverconfig.Store
|
||||
listenAddr string
|
||||
configs *serverconfig.Manager
|
||||
healthHTTP *http.Client
|
||||
controlPlaneAuth upstream.AuthorizationProvider
|
||||
|
||||
runMu sync.RWMutex
|
||||
httpServer *http.Server
|
||||
@@ -40,7 +41,7 @@ type Host struct {
|
||||
mux http.Handler
|
||||
}
|
||||
|
||||
func NewHost(store *serverconfig.Store) (*Host, error) {
|
||||
func NewHost(store *serverconfig.Store, controlPlaneAuth upstream.AuthorizationProvider) (*Host, error) {
|
||||
if store == nil {
|
||||
return nil, fmt.Errorf("backend config store is required")
|
||||
}
|
||||
@@ -50,10 +51,11 @@ func NewHost(store *serverconfig.Store) (*Host, error) {
|
||||
}
|
||||
cfg := configs.Current()
|
||||
host := &Host{
|
||||
store: store,
|
||||
listenAddr: cfg.BackendListenAddr,
|
||||
configs: configs,
|
||||
healthHTTP: newLoopbackHTTPClient(),
|
||||
store: store,
|
||||
listenAddr: cfg.BackendListenAddr,
|
||||
configs: configs,
|
||||
healthHTTP: newLoopbackHTTPClient(),
|
||||
controlPlaneAuth: controlPlaneAuth,
|
||||
}
|
||||
if err := host.rebuild(cfg); err != nil {
|
||||
return nil, err
|
||||
@@ -279,7 +281,6 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
server.Use(
|
||||
server.Recover(),
|
||||
server.ServerContext(),
|
||||
server.PolicyMiddleware(host.configs),
|
||||
server.ErrorEncoder(),
|
||||
),
|
||||
server.Mount(ads.RoutePrefix, ads.NewHTTPHandler(appdata.AdsRootPath())),
|
||||
@@ -292,17 +293,11 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
server.Name("bidi_append"),
|
||||
server.ConnectUnary(),
|
||||
server.Local(server.HTTPHandlerAction(agentModule.LocalBidiHandler)),
|
||||
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
|
||||
Name: "bidi_append",
|
||||
})),
|
||||
),
|
||||
server.POST(legacyRunSSEProcedure,
|
||||
server.Name("run_sse"),
|
||||
server.ConnectStream(),
|
||||
server.Local(server.HTTPHandlerAction(agentModule.LocalRunSSE)),
|
||||
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
|
||||
Name: "run_sse",
|
||||
})),
|
||||
),
|
||||
server.POST("/aiserver.v1.AiService/ServerTime",
|
||||
server.Name("server_time"),
|
||||
@@ -313,9 +308,6 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
MockProtoType: "aiserver.v1.ServerTimeResponse",
|
||||
MockBuilder: upstream.ServerTimeMockBuilder,
|
||||
})),
|
||||
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
|
||||
Name: "server_time",
|
||||
})),
|
||||
),
|
||||
server.POST("/aiserver.v1.AiService/GetServerConfig",
|
||||
server.Name("server_config"),
|
||||
@@ -326,9 +318,6 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
MockProtoType: "aiserver.v1.GetServerConfigResponse",
|
||||
MockBuilder: upstream.ServerConfigMockBuilder,
|
||||
})),
|
||||
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
|
||||
Name: "server_config",
|
||||
})),
|
||||
),
|
||||
server.POST("/aiserver.v1.AiService/AvailableModels",
|
||||
server.Name("available_models"),
|
||||
@@ -339,9 +328,6 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
MockProtoType: "aiserver.v1.AvailableModelsResponse",
|
||||
MockBuilder: upstream.AvailableModelsMockBuilder,
|
||||
})),
|
||||
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
|
||||
Name: "available_models",
|
||||
})),
|
||||
),
|
||||
server.POST("/aiserver.v1.AiService/GetDefaultModelNudgeData",
|
||||
server.Name("default_model_nudge"),
|
||||
@@ -352,9 +338,6 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
MockProtoType: "aiserver.v1.GetDefaultModelNudgeDataResponse",
|
||||
MockBuilder: upstream.DefaultModelNudgeMockBuilder,
|
||||
})),
|
||||
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
|
||||
Name: "default_model_nudge",
|
||||
})),
|
||||
),
|
||||
server.POST("/aiserver.v1.AnalyticsService/BootstrapStatsig",
|
||||
server.Name("bootstrap_statsig"),
|
||||
@@ -365,9 +348,6 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
MockProtoType: "aiserver.v1.BootstrapStatsigResponse",
|
||||
MockBuilder: upstream.BootstrapStatsigMockBuilder,
|
||||
})),
|
||||
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
|
||||
Name: "bootstrap_statsig",
|
||||
})),
|
||||
),
|
||||
server.POST("/aiserver.v1.AnalyticsService/GetFirstWindowStatsigDecision",
|
||||
server.Name("first_window_statsig_decision"),
|
||||
@@ -378,9 +358,6 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
MockProtoType: "aiserver.v1.GetFirstWindowStatsigDecisionResponse",
|
||||
MockBuilder: upstream.FirstWindowStatsigDecisionMockBuilder,
|
||||
})),
|
||||
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
|
||||
Name: "first_window_statsig_decision",
|
||||
})),
|
||||
),
|
||||
server.POST("/oauth/token",
|
||||
server.Name("oauth_token"),
|
||||
@@ -389,9 +366,6 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
Name: "oauth_token",
|
||||
StatusCode: http.StatusOK,
|
||||
})),
|
||||
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
|
||||
Name: "oauth_token",
|
||||
})),
|
||||
),
|
||||
server.POST("/aiserver.v1.AuthService/GetEmail",
|
||||
server.Name("auth_service_get_email"),
|
||||
@@ -400,50 +374,44 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
Name: "auth_service_get_email",
|
||||
StatusCode: http.StatusOK,
|
||||
})),
|
||||
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
|
||||
Name: "auth_service_get_email",
|
||||
})),
|
||||
),
|
||||
tabServerUpstreamProcedure("/aiserver.v1.AiService/StreamCpp", "ai_stream_cpp", server.ConnectStream(), routeDeps),
|
||||
tabServerUpstreamProcedure("/aiserver.v1.AiService/StreamNextCursorPrediction", "ai_stream_next_cursor_prediction", server.ConnectStream(), routeDeps),
|
||||
tabServerUpstreamProcedure("/aiserver.v1.AiService/GetCppEditClassification", "ai_get_cpp_edit_classification", server.ConnectUnary(), routeDeps),
|
||||
tabServerUpstreamProcedure("/aiserver.v1.AiService/RefreshTabContext", "ai_refresh_tab_context", server.ConnectUnary(), routeDeps),
|
||||
tabServerUpstreamProcedure("/aiserver.v1.AiService/CppConfig", "ai_cpp_config", server.ConnectUnary(), routeDeps),
|
||||
tabServerUpstreamProcedure("/aiserver.v1.AiService/CppEditHistoryStatus", "ai_cpp_edit_history_status", server.ConnectUnary(), routeDeps),
|
||||
tabServerUpstreamProcedure("/aiserver.v1.AiService/CppAppend", "ai_cpp_append", server.ConnectUnary(), routeDeps),
|
||||
tabServerUpstreamProcedure("/aiserver.v1.AiService/CppEditHistoryAppend", "ai_cpp_edit_history_append", server.ConnectUnary(), routeDeps),
|
||||
tabServerUpstreamProcedure("/aiserver.v1.AiService/ReportAiCodeChangeMetrics", "ai_report_ai_code_change_metrics", server.ConnectUnary(), routeDeps),
|
||||
tabServerUpstreamProcedure("/aiserver.v1.AiService/WriteGitCommitMessage", "ai_write_git_commit_message", server.ConnectUnary(), routeDeps),
|
||||
tabServerUpstreamProcedure("/aiserver.v1.AiService/WriteGitBranchName", "ai_write_git_branch_name", server.ConnectUnary(), routeDeps),
|
||||
repositoryServiceProcedure(forwarder.RepositoryServiceFastRepoInitHandshakeV2Procedure, "repository_fast_repo_init_handshake_v2", server.ConnectUnary(), agentModule, routeDeps),
|
||||
repositoryServiceProcedure(forwarder.RepositoryServiceFastRepoInitHandshakeProcedure, "repository_fast_repo_init_handshake", server.ConnectUnary(), agentModule, routeDeps),
|
||||
repositoryServiceProcedure(forwarder.RepositoryServiceFastRepoSyncCompleteProcedure, "repository_fast_repo_sync_complete", server.ConnectUnary(), agentModule, routeDeps),
|
||||
repositoryServiceProcedure(forwarder.RepositoryServiceSyncMerkleSubtreeV2Procedure, "repository_sync_merkle_subtree_v2", server.ConnectUnary(), agentModule, routeDeps),
|
||||
repositoryServiceProcedure(forwarder.RepositoryServiceSyncMerkleSubtreeProcedure, "repository_sync_merkle_subtree", server.ConnectUnary(), agentModule, routeDeps),
|
||||
repositoryServiceProcedure(forwarder.RepositoryServiceFastUpdateFileV2Procedure, "repository_fast_update_file_v2", server.ConnectUnary(), agentModule, routeDeps),
|
||||
repositoryServiceProcedure(forwarder.RepositoryServiceFastUpdateFileProcedure, "repository_fast_update_file", server.ConnectUnary(), agentModule, routeDeps),
|
||||
repositoryServiceProcedure(forwarder.RepositoryServiceEnsureIndexCreatedProcedure, "repository_ensure_index_created", server.ConnectUnary(), agentModule, routeDeps),
|
||||
repositoryServiceProcedure(forwarder.RepositoryServiceGetCopyStatusProcedure, "repository_get_copy_status", server.ConnectUnary(), agentModule, routeDeps),
|
||||
repositoryServiceProcedure(forwarder.RepositoryServiceGetUploadLimitsProcedure, "repository_get_upload_limits", server.ConnectUnary(), agentModule, routeDeps),
|
||||
repositoryServiceProcedure(forwarder.RepositoryServiceGetNumFilesToSendProcedure, "repository_get_num_files_to_send", server.ConnectUnary(), agentModule, routeDeps),
|
||||
repositoryServiceProcedure(forwarder.RepositoryServiceGetAvailableChunkingStrategiesProcedure, "repository_get_available_chunking_strategies", server.ConnectUnary(), agentModule, routeDeps),
|
||||
repositoryServiceProcedure(forwarder.RepositoryServiceGetHighLevelFolderDescriptionProcedure, "repository_get_high_level_folder_description", server.ConnectUnary(), agentModule, routeDeps),
|
||||
repositoryServiceProcedure(forwarder.RepositoryServiceRepositoryStatusProcedure, "repository_status", server.ConnectUnary(), agentModule, routeDeps),
|
||||
repositoryServiceProcedure(forwarder.RepositoryServiceBatchRepositoryStatusProcedure, "repository_batch_status", server.ConnectUnary(), agentModule, routeDeps),
|
||||
uploadServiceProcedure(forwarder.UploadServiceUploadDocumentationProcedure, "upload_documentation", server.ConnectUnary(), agentModule, routeDeps),
|
||||
uploadServiceProcedure(forwarder.UploadServiceGetDocProcedure, "upload_get_doc", server.ConnectUnary(), agentModule, routeDeps),
|
||||
uploadServiceProcedure(forwarder.UploadServiceGetPagesProcedure, "upload_get_pages", server.ConnectUnary(), agentModule, routeDeps),
|
||||
uploadServiceProcedure(forwarder.UploadServiceUploadedStatusProcedure, "upload_uploaded_status", server.ConnectUnary(), agentModule, routeDeps),
|
||||
tabServerProcedure("/aiserver.v1.AiService/StreamCpp", "ai_stream_cpp", server.ConnectStream(), routeDeps),
|
||||
tabServerProcedure("/aiserver.v1.AiService/StreamNextCursorPrediction", "ai_stream_next_cursor_prediction", server.ConnectStream(), routeDeps),
|
||||
tabServerProcedure("/aiserver.v1.AiService/GetCppEditClassification", "ai_get_cpp_edit_classification", server.ConnectUnary(), routeDeps),
|
||||
tabServerProcedure("/aiserver.v1.AiService/RefreshTabContext", "ai_refresh_tab_context", server.ConnectUnary(), routeDeps),
|
||||
tabServerProcedure("/aiserver.v1.AiService/CppConfig", "ai_cpp_config", server.ConnectUnary(), routeDeps),
|
||||
tabServerProcedure("/aiserver.v1.AiService/CppEditHistoryStatus", "ai_cpp_edit_history_status", server.ConnectUnary(), routeDeps),
|
||||
tabServerProcedure("/aiserver.v1.AiService/CppAppend", "ai_cpp_append", server.ConnectUnary(), routeDeps),
|
||||
tabServerProcedure("/aiserver.v1.AiService/CppEditHistoryAppend", "ai_cpp_edit_history_append", server.ConnectUnary(), routeDeps),
|
||||
tabServerProcedure("/aiserver.v1.AiService/ReportAiCodeChangeMetrics", "ai_report_ai_code_change_metrics", server.ConnectUnary(), routeDeps),
|
||||
tabServerProcedure("/aiserver.v1.AiService/WriteGitCommitMessage", "ai_write_git_commit_message", server.ConnectUnary(), routeDeps),
|
||||
tabServerProcedure("/aiserver.v1.AiService/WriteGitBranchName", "ai_write_git_branch_name", server.ConnectUnary(), routeDeps),
|
||||
repositoryServiceProcedure(forwarder.RepositoryServiceFastRepoInitHandshakeV2Procedure, "repository_fast_repo_init_handshake_v2", server.ConnectUnary(), agentModule),
|
||||
repositoryServiceProcedure(forwarder.RepositoryServiceFastRepoInitHandshakeProcedure, "repository_fast_repo_init_handshake", server.ConnectUnary(), agentModule),
|
||||
repositoryServiceProcedure(forwarder.RepositoryServiceFastRepoSyncCompleteProcedure, "repository_fast_repo_sync_complete", server.ConnectUnary(), agentModule),
|
||||
repositoryServiceProcedure(forwarder.RepositoryServiceSyncMerkleSubtreeV2Procedure, "repository_sync_merkle_subtree_v2", server.ConnectUnary(), agentModule),
|
||||
repositoryServiceProcedure(forwarder.RepositoryServiceSyncMerkleSubtreeProcedure, "repository_sync_merkle_subtree", server.ConnectUnary(), agentModule),
|
||||
repositoryServiceProcedure(forwarder.RepositoryServiceFastUpdateFileV2Procedure, "repository_fast_update_file_v2", server.ConnectUnary(), agentModule),
|
||||
repositoryServiceProcedure(forwarder.RepositoryServiceFastUpdateFileProcedure, "repository_fast_update_file", server.ConnectUnary(), agentModule),
|
||||
repositoryServiceProcedure(forwarder.RepositoryServiceEnsureIndexCreatedProcedure, "repository_ensure_index_created", server.ConnectUnary(), agentModule),
|
||||
repositoryServiceProcedure(forwarder.RepositoryServiceGetCopyStatusProcedure, "repository_get_copy_status", server.ConnectUnary(), agentModule),
|
||||
repositoryServiceProcedure(forwarder.RepositoryServiceGetUploadLimitsProcedure, "repository_get_upload_limits", server.ConnectUnary(), agentModule),
|
||||
repositoryServiceProcedure(forwarder.RepositoryServiceGetNumFilesToSendProcedure, "repository_get_num_files_to_send", server.ConnectUnary(), agentModule),
|
||||
repositoryServiceProcedure(forwarder.RepositoryServiceGetAvailableChunkingStrategiesProcedure, "repository_get_available_chunking_strategies", server.ConnectUnary(), agentModule),
|
||||
repositoryServiceProcedure(forwarder.RepositoryServiceGetHighLevelFolderDescriptionProcedure, "repository_get_high_level_folder_description", server.ConnectUnary(), agentModule),
|
||||
repositoryServiceProcedure(forwarder.RepositoryServiceRepositoryStatusProcedure, "repository_status", server.ConnectUnary(), agentModule),
|
||||
repositoryServiceProcedure(forwarder.RepositoryServiceBatchRepositoryStatusProcedure, "repository_batch_status", server.ConnectUnary(), agentModule),
|
||||
uploadServiceProcedure(forwarder.UploadServiceUploadDocumentationProcedure, "upload_documentation", server.ConnectUnary(), agentModule),
|
||||
uploadServiceProcedure(forwarder.UploadServiceGetDocProcedure, "upload_get_doc", server.ConnectUnary(), agentModule),
|
||||
uploadServiceProcedure(forwarder.UploadServiceGetPagesProcedure, "upload_get_pages", server.ConnectUnary(), agentModule),
|
||||
uploadServiceProcedure(forwarder.UploadServiceUploadedStatusProcedure, "upload_uploaded_status", server.ConnectUnary(), agentModule),
|
||||
server.Any("/aiserver.v1.AiService/*",
|
||||
server.Name("ai_service"),
|
||||
server.HTTP(),
|
||||
server.Local(server.HTTPHandlerAction(agentModule.AiHandler)),
|
||||
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
|
||||
Name: "ai_service",
|
||||
})),
|
||||
),
|
||||
tabServerUpstreamProcedure("/aiserver.v1.CppService/AvailableModels", "cpp_available_models", server.ConnectUnary(), routeDeps),
|
||||
tabServerUpstreamProcedure("/aiserver.v1.CppService/RecordCppFate", "cpp_record_cpp_fate", server.ConnectUnary(), routeDeps),
|
||||
tabServerProcedure("/aiserver.v1.CppService/AvailableModels", "cpp_available_models", server.ConnectUnary(), routeDeps),
|
||||
tabServerProcedure("/aiserver.v1.CppService/RecordCppFate", "cpp_record_cpp_fate", server.ConnectUnary(), routeDeps),
|
||||
server.Any("/aiserver.v1.CppService/*",
|
||||
server.Name("cpp_service"),
|
||||
server.HTTP(),
|
||||
@@ -451,14 +419,11 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
http.NotFound(ctx.Writer, ctx.Request)
|
||||
return nil
|
||||
}),
|
||||
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
|
||||
Name: "cpp_service",
|
||||
})),
|
||||
),
|
||||
tabServerUpstreamProcedure("/aiserver.v1.FileSyncService/FSSyncFile", "file_sync_sync_file", server.ConnectUnary(), routeDeps),
|
||||
tabServerUpstreamProcedure("/aiserver.v1.FileSyncService/FSIsEnabledForUser", "file_sync_is_enabled_for_user", server.ConnectUnary(), routeDeps),
|
||||
tabServerUpstreamProcedure("/aiserver.v1.FileSyncService/FSConfig", "file_sync_config", server.ConnectUnary(), routeDeps),
|
||||
tabServerUpstreamProcedure("/aiserver.v1.FileSyncService/FSUploadFile", "file_sync_upload_file", server.ConnectUnary(), routeDeps),
|
||||
tabServerProcedure("/aiserver.v1.FileSyncService/FSSyncFile", "file_sync_sync_file", server.ConnectUnary(), routeDeps),
|
||||
tabServerProcedure("/aiserver.v1.FileSyncService/FSIsEnabledForUser", "file_sync_is_enabled_for_user", server.ConnectUnary(), routeDeps),
|
||||
tabServerProcedure("/aiserver.v1.FileSyncService/FSConfig", "file_sync_config", server.ConnectUnary(), routeDeps),
|
||||
tabServerProcedure("/aiserver.v1.FileSyncService/FSUploadFile", "file_sync_upload_file", server.ConnectUnary(), routeDeps),
|
||||
server.Any("/aiserver.v1.FileSyncService/*",
|
||||
server.Name("file_sync"),
|
||||
server.HTTP(),
|
||||
@@ -466,25 +431,16 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
http.NotFound(ctx.Writer, ctx.Request)
|
||||
return nil
|
||||
}),
|
||||
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
|
||||
Name: "file_sync",
|
||||
})),
|
||||
),
|
||||
server.POST("/aiserver.v1.DashboardService/GetTokenUsage",
|
||||
server.Name("dashboard_token_usage"),
|
||||
server.HTTP(),
|
||||
server.Local(server.HTTPHandlerAction(agentModule.AiHandler)),
|
||||
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
|
||||
Name: "dashboard_token_usage",
|
||||
})),
|
||||
),
|
||||
server.POST("/aiserver.v1.DashboardService/GetGlassEarlyPreviewEnrollment",
|
||||
server.Name("dashboard_glass_early_preview_enrollment"),
|
||||
server.ConnectUnary(),
|
||||
server.Local(server.HTTPHandlerAction(agentModule.AiHandler)),
|
||||
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
|
||||
Name: "dashboard_glass_early_preview_enrollment",
|
||||
})),
|
||||
),
|
||||
server.POST("/aiserver.v1.DashboardService/GetCurrentPeriodUsage",
|
||||
server.Name("dashboard_current_period_usage"),
|
||||
@@ -495,9 +451,6 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
MockProtoType: "aiserver.v1.GetCurrentPeriodUsageResponse",
|
||||
MockBuilder: upstream.DashboardCurrentPeriodUsageMockBuilder,
|
||||
})),
|
||||
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
|
||||
Name: "dashboard_current_period_usage",
|
||||
})),
|
||||
),
|
||||
server.POST("/aiserver.v1.DashboardService/GetTeams",
|
||||
server.Name("dashboard_get_teams"),
|
||||
@@ -508,22 +461,21 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
MockProtoType: "aiserver.v1.GetTeamsResponse",
|
||||
MockBuilder: upstream.DashboardTeamsMockBuilder,
|
||||
})),
|
||||
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
|
||||
Name: "dashboard_get_teams",
|
||||
})),
|
||||
),
|
||||
server.POST("/aiserver.v1.DashboardService/GetManagedSkills",
|
||||
server.Name("dashboard_get_managed_skills"),
|
||||
server.ConnectUnary(),
|
||||
server.Local(upstream.MockProtoAction(routeDeps, upstream.CompatRouteConfig{
|
||||
Name: "dashboard_get_managed_skills",
|
||||
StatusCode: http.StatusOK,
|
||||
MockProtoType: "aiserver.v1.GetManagedSkillsResponse",
|
||||
MockBuilder: upstream.DashboardManagedSkillsMockBuilder,
|
||||
})),
|
||||
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
|
||||
Name: "dashboard_get_managed_skills",
|
||||
})),
|
||||
server.Local(cursorControlPlaneAction(
|
||||
host.controlPlaneAuth,
|
||||
routeDeps,
|
||||
"dashboard_get_managed_skills",
|
||||
upstream.MockProtoAction(routeDeps, upstream.CompatRouteConfig{
|
||||
Name: "dashboard_get_managed_skills",
|
||||
StatusCode: http.StatusOK,
|
||||
MockProtoType: "aiserver.v1.GetManagedSkillsResponse",
|
||||
MockBuilder: upstream.DashboardManagedSkillsMockBuilder,
|
||||
}),
|
||||
)),
|
||||
),
|
||||
server.POST("/aiserver.v1.DashboardService/GetMe",
|
||||
server.Name("dashboard_get_me"),
|
||||
@@ -534,9 +486,6 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
MockProtoType: "aiserver.v1.GetMeResponse",
|
||||
MockBuilder: upstream.DashboardGetMeMockBuilder,
|
||||
})),
|
||||
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
|
||||
Name: "dashboard_get_me",
|
||||
})),
|
||||
),
|
||||
server.POST("/aiserver.v1.DashboardService/GetUserPrivacyMode",
|
||||
server.Name("dashboard_user_privacy_mode"),
|
||||
@@ -547,9 +496,6 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
MockProtoType: "aiserver.v1.GetUserPrivacyModeResponse",
|
||||
MockBuilder: upstream.DashboardUserPrivacyModeMockBuilder,
|
||||
})),
|
||||
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
|
||||
Name: "dashboard_user_privacy_mode",
|
||||
})),
|
||||
),
|
||||
server.POST("/aiserver.v1.DashboardService/GetPlanInfo",
|
||||
server.Name("dashboard_plan_info"),
|
||||
@@ -560,9 +506,6 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
MockProtoType: "aiserver.v1.GetPlanInfoResponse",
|
||||
MockBuilder: upstream.DashboardPlanInfoMockBuilder,
|
||||
})),
|
||||
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
|
||||
Name: "dashboard_plan_info",
|
||||
})),
|
||||
),
|
||||
server.POST("/aiserver.v1.DashboardService/GetUsageLimitStatusAndActiveGrants",
|
||||
server.Name("dashboard_usage_limit_status"),
|
||||
@@ -573,9 +516,6 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
MockProtoType: "aiserver.v1.GetUsageLimitStatusAndActiveGrantsResponse",
|
||||
MockBuilder: upstream.DashboardUsageLimitStatusAndActiveGrantsMockBuilder,
|
||||
})),
|
||||
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
|
||||
Name: "dashboard_usage_limit_status",
|
||||
})),
|
||||
),
|
||||
server.POST("/aiserver.v1.DashboardService/IsOnNewPricing",
|
||||
server.Name("dashboard_is_on_new_pricing"),
|
||||
@@ -586,11 +526,25 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
MockProtoType: "aiserver.v1.IsOnNewPricingResponse",
|
||||
MockBuilder: upstream.DashboardIsOnNewPricingMockBuilder,
|
||||
})),
|
||||
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
|
||||
Name: "dashboard_is_on_new_pricing",
|
||||
})),
|
||||
),
|
||||
// tabServerUpstreamProcedure("/aiserver.v1.DashboardService/GetEffectiveUserPlugins", "dashboard_get_effective_user_plugins", server.ConnectUnary(), routeDeps),
|
||||
cursorControlPlaneProcedure("/aiserver.v1.DashboardService/AddMarketplace", "dashboard_add_marketplace", server.ConnectUnary(), host.controlPlaneAuth, routeDeps),
|
||||
cursorControlPlaneProcedure("/aiserver.v1.DashboardService/AddMcpServersFromPlugin", "dashboard_add_mcp_servers_from_plugin", server.ConnectUnary(), host.controlPlaneAuth, routeDeps),
|
||||
cursorControlPlaneProcedure("/aiserver.v1.DashboardService/BatchGetPluginMcpConfig", "dashboard_batch_get_plugin_mcp_config", server.ConnectUnary(), host.controlPlaneAuth, routeDeps),
|
||||
cursorControlPlaneProcedure("/aiserver.v1.DashboardService/GetAvailableMcpServers", "dashboard_get_available_mcp_servers", server.ConnectUnary(), host.controlPlaneAuth, routeDeps),
|
||||
cursorControlPlaneProcedure("/aiserver.v1.DashboardService/GetEffectiveUserPlugins", "dashboard_get_effective_user_plugins", server.ConnectUnary(), host.controlPlaneAuth, routeDeps),
|
||||
cursorControlPlaneProcedure("/aiserver.v1.DashboardService/GetPlugin", "dashboard_get_plugin", server.ConnectUnary(), host.controlPlaneAuth, routeDeps),
|
||||
cursorControlPlaneProcedure("/aiserver.v1.DashboardService/GetPluginMcpConfig", "dashboard_get_plugin_mcp_config", server.ConnectUnary(), host.controlPlaneAuth, routeDeps),
|
||||
cursorControlPlaneProcedure("/aiserver.v1.DashboardService/InstallUserPlugin", "dashboard_install_user_plugin", server.ConnectUnary(), host.controlPlaneAuth, routeDeps),
|
||||
cursorControlPlaneProcedure("/aiserver.v1.DashboardService/ListMarketplacePlugins", "dashboard_list_marketplace_plugins", server.ConnectUnary(), host.controlPlaneAuth, routeDeps),
|
||||
cursorControlPlaneProcedure("/aiserver.v1.DashboardService/ListMarketplaces", "dashboard_list_marketplaces", server.ConnectUnary(), host.controlPlaneAuth, routeDeps),
|
||||
cursorControlPlaneProcedure("/aiserver.v1.DashboardService/ListUserPluginInstalls", "dashboard_list_user_plugin_installs", server.ConnectUnary(), host.controlPlaneAuth, routeDeps),
|
||||
cursorControlPlaneProcedure("/aiserver.v1.DashboardService/RefreshMarketplace", "dashboard_refresh_marketplace", server.ConnectUnary(), host.controlPlaneAuth, routeDeps),
|
||||
cursorControlPlaneProcedure("/aiserver.v1.DashboardService/RegisterMarketplaceAndPlugins", "dashboard_register_marketplace_and_plugins", server.ConnectUnary(), host.controlPlaneAuth, routeDeps),
|
||||
cursorControlPlaneProcedure("/aiserver.v1.DashboardService/RemoveMarketplace", "dashboard_remove_marketplace", server.ConnectUnary(), host.controlPlaneAuth, routeDeps),
|
||||
cursorControlPlaneProcedure("/aiserver.v1.DashboardService/ResolvePluginsByRef", "dashboard_resolve_plugins_by_ref", server.ConnectUnary(), host.controlPlaneAuth, routeDeps),
|
||||
cursorControlPlaneProcedure("/aiserver.v1.DashboardService/UninstallUserPlugin", "dashboard_uninstall_user_plugin", server.ConnectUnary(), host.controlPlaneAuth, routeDeps),
|
||||
cursorControlPlaneProcedure("/aiserver.v1.DashboardService/UpdateUserPluginInstall", "dashboard_update_user_plugin_install", server.ConnectUnary(), host.controlPlaneAuth, routeDeps),
|
||||
cursorControlPlaneProcedure("/aiserver.v1.MCPRegistryService/GetKnownServers", "mcp_registry_get_known_servers", server.ConnectUnary(), host.controlPlaneAuth, routeDeps),
|
||||
server.Any("/aiserver.v1.DashboardService/*",
|
||||
server.Name("dashboard"),
|
||||
server.HTTP(),
|
||||
@@ -598,9 +552,6 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
http.NotFound(ctx.Writer, ctx.Request)
|
||||
return nil
|
||||
}),
|
||||
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
|
||||
Name: "dashboard",
|
||||
})),
|
||||
),
|
||||
server.Any("/aiserver.v1.NetworkService/*",
|
||||
server.Name("network_service"),
|
||||
@@ -609,9 +560,6 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
http.NotFound(ctx.Writer, ctx.Request)
|
||||
return nil
|
||||
}),
|
||||
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
|
||||
Name: "network_service",
|
||||
})),
|
||||
),
|
||||
server.Any("/aiserver.v1.InAppAdService/*",
|
||||
server.Name("in_app_ad"),
|
||||
@@ -620,9 +568,6 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
http.NotFound(ctx.Writer, ctx.Request)
|
||||
return nil
|
||||
}),
|
||||
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
|
||||
Name: "in_app_ad",
|
||||
})),
|
||||
),
|
||||
server.GET("/auth/full_stripe_profile",
|
||||
server.Name("auth_full_stripe_profile"),
|
||||
@@ -631,9 +576,6 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
Name: "auth_full_stripe_profile",
|
||||
StatusCode: http.StatusOK,
|
||||
})),
|
||||
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
|
||||
Name: "auth_full_stripe_profile",
|
||||
})),
|
||||
),
|
||||
server.GET("/auth/stripe_profile",
|
||||
server.Name("auth_stripe_profile"),
|
||||
@@ -642,9 +584,6 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
Name: "auth_stripe_profile",
|
||||
StatusCode: http.StatusOK,
|
||||
})),
|
||||
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
|
||||
Name: "auth_stripe_profile",
|
||||
})),
|
||||
),
|
||||
server.GET("/auth/has_valid_payment_method",
|
||||
server.Name("auth_has_valid_payment_method"),
|
||||
@@ -656,9 +595,6 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
"hasValidPaymentMethod": true,
|
||||
},
|
||||
})),
|
||||
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
|
||||
Name: "auth_has_valid_payment_method",
|
||||
})),
|
||||
),
|
||||
server.Any("/auth/poll",
|
||||
server.Name("auth_poll"),
|
||||
@@ -667,9 +603,6 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
Name: "auth_poll",
|
||||
StatusCode: http.StatusOK,
|
||||
})),
|
||||
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
|
||||
Name: "auth_poll",
|
||||
})),
|
||||
),
|
||||
server.POST("/auth/logout",
|
||||
server.Name("auth_logout"),
|
||||
@@ -678,9 +611,6 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
Name: "auth_logout",
|
||||
StatusCode: http.StatusNoContent,
|
||||
})),
|
||||
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
|
||||
Name: "auth_logout",
|
||||
})),
|
||||
),
|
||||
server.Any("/auth/*",
|
||||
server.Name("auth_proxy"),
|
||||
@@ -689,58 +619,32 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error {
|
||||
http.NotFound(ctx.Writer, ctx.Request)
|
||||
return nil
|
||||
}),
|
||||
server.Upstream(upstream.DirectAction(routeDeps, upstream.CompatRouteConfig{
|
||||
Name: "auth_proxy",
|
||||
})),
|
||||
),
|
||||
)
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func directUpstreamProcedure(pattern string, name string, protocol server.RouteOption, deps upstream.Dependencies) server.Option {
|
||||
direct := upstream.DirectAction(deps, upstream.CompatRouteConfig{Name: name})
|
||||
action := func(ctx *server.Context) error {
|
||||
if ctx != nil && ctx.UpstreamURL == nil && ctx.Request != nil && ctx.Request.URL != nil {
|
||||
targetURL := *ctx.Request.URL
|
||||
targetURL.Scheme = "https"
|
||||
targetURL.Host = "api2.cursor.sh:443"
|
||||
ctx.UpstreamURL = &targetURL
|
||||
}
|
||||
return direct(ctx)
|
||||
}
|
||||
return server.POST(pattern,
|
||||
server.Name(name),
|
||||
protocol,
|
||||
server.Local(action),
|
||||
server.Upstream(action),
|
||||
)
|
||||
}
|
||||
|
||||
func repositoryServiceProcedure(pattern string, name string, protocol server.RouteOption, module *forwarder.Module, deps upstream.Dependencies) server.Option {
|
||||
func repositoryServiceProcedure(pattern string, name string, protocol server.RouteOption, module *forwarder.Module) server.Option {
|
||||
localAction := server.HTTPHandlerAction(module.RepositoryServiceHandler)
|
||||
upstreamAction := upstream.DirectAction(deps, upstream.CompatRouteConfig{Name: name})
|
||||
return server.POST(pattern,
|
||||
server.Name(name),
|
||||
protocol,
|
||||
server.Local(localAction),
|
||||
server.Upstream(upstreamAction),
|
||||
)
|
||||
}
|
||||
|
||||
func uploadServiceProcedure(pattern string, name string, protocol server.RouteOption, module *forwarder.Module, deps upstream.Dependencies) server.Option {
|
||||
func uploadServiceProcedure(pattern string, name string, protocol server.RouteOption, module *forwarder.Module) server.Option {
|
||||
localAction := server.HTTPHandlerAction(module.UploadServiceHandler)
|
||||
upstreamAction := upstream.DirectAction(deps, upstream.CompatRouteConfig{Name: name})
|
||||
return server.POST(pattern,
|
||||
server.Name(name),
|
||||
protocol,
|
||||
server.Local(localAction),
|
||||
server.Upstream(upstreamAction),
|
||||
)
|
||||
}
|
||||
|
||||
func tabServerUpstreamProcedure(pattern string, name string, protocol server.RouteOption, deps upstream.Dependencies) server.Option {
|
||||
direct := upstream.DirectAction(deps, upstream.CompatRouteConfig{Name: name})
|
||||
func tabServerProcedure(pattern string, name string, protocol server.RouteOption, deps upstream.Dependencies) server.Option {
|
||||
forward := upstream.ForwardAction(deps, upstream.CompatRouteConfig{Name: name})
|
||||
action := func(ctx *server.Context) error {
|
||||
if ctx != nil && ctx.Request != nil && ctx.Request.URL != nil {
|
||||
baseURL, err := url.Parse(tabServerBaseURL)
|
||||
@@ -752,16 +656,55 @@ func tabServerUpstreamProcedure(pattern string, name string, protocol server.Rou
|
||||
targetURL.Host = baseURL.Host
|
||||
ctx.UpstreamURL = &targetURL
|
||||
}
|
||||
return direct(ctx)
|
||||
return forward(ctx)
|
||||
}
|
||||
return server.POST(pattern,
|
||||
server.Name(name),
|
||||
protocol,
|
||||
server.Local(action),
|
||||
server.Upstream(action),
|
||||
)
|
||||
}
|
||||
|
||||
func cursorControlPlaneProcedure(
|
||||
pattern string,
|
||||
name string,
|
||||
protocol server.RouteOption,
|
||||
authorizationProvider upstream.AuthorizationProvider,
|
||||
deps upstream.Dependencies,
|
||||
) server.Option {
|
||||
notFound := func(ctx *server.Context) error {
|
||||
http.NotFound(ctx.Writer, ctx.Request)
|
||||
return nil
|
||||
}
|
||||
return server.POST(pattern,
|
||||
server.Name(name),
|
||||
protocol,
|
||||
server.Local(cursorControlPlaneAction(authorizationProvider, deps, name, notFound)),
|
||||
)
|
||||
}
|
||||
|
||||
func cursorControlPlaneAction(
|
||||
authorizationProvider upstream.AuthorizationProvider,
|
||||
deps upstream.Dependencies,
|
||||
name string,
|
||||
fallback server.HandlerFunc,
|
||||
) server.HandlerFunc {
|
||||
forward := upstream.AuthenticatedForwardAction(deps, upstream.CompatRouteConfig{Name: name}, authorizationProvider)
|
||||
return func(ctx *server.Context) error {
|
||||
if authorizationProvider == nil || !authorizationProvider.SignedIn() {
|
||||
return fallback(ctx)
|
||||
}
|
||||
if ctx == nil || ctx.Request == nil || ctx.Request.URL == nil {
|
||||
return fmt.Errorf("Cursor 控制面请求上下文无效")
|
||||
}
|
||||
targetURL := *ctx.Request.URL
|
||||
targetURL.Scheme = "https"
|
||||
targetURL.Host = "api2.cursor.sh:443"
|
||||
ctx.UpstreamURL = &targetURL
|
||||
return forward(ctx)
|
||||
}
|
||||
}
|
||||
|
||||
type serverSystemSettings struct {
|
||||
configs *serverconfig.Manager
|
||||
}
|
||||
|
||||
@@ -171,20 +171,6 @@ func (manager *Manager) LegacyRuntimeSnapshot(_ context.Context) (legacyruntime.
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (manager *Manager) RouteMode(hasUpstreamURL bool) string {
|
||||
if !hasUpstreamURL {
|
||||
return DefaultRoutingMode
|
||||
}
|
||||
if manager == nil {
|
||||
return DefaultRoutingMode
|
||||
}
|
||||
mode := normalizeRoutingMode(manager.Current().Routing.Mode)
|
||||
if mode == "" {
|
||||
return DefaultRoutingMode
|
||||
}
|
||||
return mode
|
||||
}
|
||||
|
||||
func (manager *Manager) setCurrent(cfg Config) {
|
||||
next := cfg
|
||||
manager.current.Store(&next)
|
||||
|
||||
@@ -136,6 +136,9 @@ func (store *Store) saveLocked(normalized Config) error {
|
||||
}
|
||||
|
||||
func shouldPersistNormalizedConfig(raw []byte, current Config, normalized Config) bool {
|
||||
if yamlHasKey(raw, "routing") {
|
||||
return true
|
||||
}
|
||||
if !yamlHasKey(raw, "backendListenAddr") || !yamlHasKey(raw, "proxyListenAddr") {
|
||||
return true
|
||||
}
|
||||
|
||||
@@ -15,7 +15,6 @@ const (
|
||||
DefaultBackendListenAddr = "127.0.0.1:18090"
|
||||
DefaultProxyListenAddr = "127.0.0.1:18080"
|
||||
DefaultFrontendBaseURL = "http://127.0.0.1"
|
||||
DefaultRoutingMode = "local"
|
||||
DefaultProviderStreamIdleTimeoutSeconds = 240
|
||||
MinProviderStreamIdleTimeoutSeconds = 30
|
||||
)
|
||||
@@ -43,10 +42,6 @@ type ModelAdapterConfig struct {
|
||||
ThinkingBudgetTokens int `json:"thinkingBudgetTokens" yaml:"thinkingBudgetTokens"`
|
||||
}
|
||||
|
||||
type RoutingConfig struct {
|
||||
Mode string `json:"mode" yaml:"mode"`
|
||||
}
|
||||
|
||||
type HomeMetricsConfig struct {
|
||||
IncludeCacheWriteInHitRate bool `json:"includeCacheWriteInHitRate" yaml:"includeCacheWriteInHitRate"`
|
||||
}
|
||||
@@ -57,7 +52,6 @@ type Config struct {
|
||||
BackendListenAddr string `json:"backendListenAddr" yaml:"backendListenAddr"`
|
||||
ProxyListenAddr string `json:"proxyListenAddr" yaml:"proxyListenAddr"`
|
||||
ModelAdapters []ModelAdapterConfig `json:"modelAdapters" yaml:"modelAdapters"`
|
||||
Routing RoutingConfig `json:"routing" yaml:"routing"`
|
||||
HomeMetrics HomeMetricsConfig `json:"homeMetrics" yaml:"homeMetrics"`
|
||||
LastAgentModelHash string `json:"lastAgentModelHash" yaml:"lastAgentModelHash"`
|
||||
}
|
||||
@@ -69,9 +63,6 @@ func DefaultConfig() Config {
|
||||
BackendListenAddr: DefaultBackendListenAddr,
|
||||
ProxyListenAddr: DefaultProxyListenAddr,
|
||||
ModelAdapters: []ModelAdapterConfig{},
|
||||
Routing: RoutingConfig{
|
||||
Mode: DefaultRoutingMode,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
@@ -91,10 +82,6 @@ func NormalizeConfig(input Config) (Config, error) {
|
||||
output.ProxyListenAddr = proxyListenAddr
|
||||
output.HomeMetrics.IncludeCacheWriteInHitRate = input.HomeMetrics.IncludeCacheWriteInHitRate
|
||||
output.LastAgentModelHash = strings.TrimSpace(input.LastAgentModelHash)
|
||||
output.Routing.Mode = normalizeRoutingMode(input.Routing.Mode)
|
||||
if output.Routing.Mode == "" {
|
||||
output.Routing.Mode = DefaultRoutingMode
|
||||
}
|
||||
adapters, err := NormalizeModelAdapterConfigs(input.ModelAdapters)
|
||||
if err != nil {
|
||||
return Config{}, err
|
||||
@@ -280,14 +267,3 @@ func normalizeModelAdapterType(value string) string {
|
||||
return ""
|
||||
}
|
||||
}
|
||||
|
||||
func normalizeRoutingMode(value string) string {
|
||||
switch strings.ToLower(strings.TrimSpace(value)) {
|
||||
case "", "local":
|
||||
return "local"
|
||||
case "upstream":
|
||||
return "upstream"
|
||||
default:
|
||||
return ""
|
||||
}
|
||||
}
|
||||
|
||||
@@ -34,7 +34,6 @@ type Context struct {
|
||||
StartedAt time.Time
|
||||
|
||||
UpstreamURL *url.URL
|
||||
Mode ExecutionMode
|
||||
LastError error
|
||||
|
||||
Logger *slog.Logger
|
||||
@@ -48,7 +47,6 @@ func newContext(writer http.ResponseWriter, request *http.Request, route Route)
|
||||
Protocol: route.Protocol,
|
||||
StartedAt: time.Now(),
|
||||
Logger: slog.Default(),
|
||||
Mode: ModeLocal,
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -1,14 +1,12 @@
|
||||
package server
|
||||
|
||||
import (
|
||||
"cursor/internal/logger"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"runtime/debug"
|
||||
"strings"
|
||||
|
||||
serverconfig "cursor/internal/backend/server/config"
|
||||
legacyruntime "cursor/internal/runtime"
|
||||
)
|
||||
|
||||
@@ -39,16 +37,6 @@ func ServerContext() Middleware {
|
||||
}
|
||||
}
|
||||
|
||||
func PolicyMiddleware(configs *serverconfig.Manager) Middleware {
|
||||
return func(next HandlerFunc) HandlerFunc {
|
||||
return func(ctx *Context) error {
|
||||
ctx.Mode = parseExecutionMode(configs.RouteMode(ctx.UpstreamURL != nil))
|
||||
logger.Infof("ctx.Mode=%s upstream=%t", ctx.Mode, ctx.UpstreamURL != nil)
|
||||
return next(ctx)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func ErrorEncoder() Middleware {
|
||||
return func(next HandlerFunc) HandlerFunc {
|
||||
return func(ctx *Context) error {
|
||||
|
||||
@@ -1,19 +0,0 @@
|
||||
package server
|
||||
|
||||
type ExecutionMode string
|
||||
|
||||
const (
|
||||
// ModeLocal 表示本地模式,适用于直接处理请求的情况。
|
||||
ModeLocal ExecutionMode = "local"
|
||||
// ModeUpstream 表示直连上游模式,适用于将请求转发到原始地址。
|
||||
ModeUpstream ExecutionMode = "upstream"
|
||||
)
|
||||
|
||||
func parseExecutionMode(value string) ExecutionMode {
|
||||
switch value {
|
||||
case string(ModeUpstream):
|
||||
return ModeUpstream
|
||||
default:
|
||||
return ModeLocal
|
||||
}
|
||||
}
|
||||
@@ -17,7 +17,6 @@ type Route struct {
|
||||
Protocol ProtocolClass
|
||||
Middleware []Middleware
|
||||
Local HandlerFunc
|
||||
Upstream HandlerFunc
|
||||
}
|
||||
|
||||
type App struct {
|
||||
@@ -129,12 +128,6 @@ func Local(action HandlerFunc) RouteOption {
|
||||
}
|
||||
}
|
||||
|
||||
func Upstream(action HandlerFunc) RouteOption {
|
||||
return func(route *Route) {
|
||||
route.Upstream = action
|
||||
}
|
||||
}
|
||||
|
||||
func (app *App) registerRoute(route Route) {
|
||||
handler := app.buildRouteHandler(route)
|
||||
if route.Method == "" {
|
||||
@@ -148,12 +141,6 @@ func (app *App) buildRouteHandler(route Route) http.HandlerFunc {
|
||||
chain := append([]Middleware{}, app.globalMiddlewares...)
|
||||
chain = append(chain, route.Middleware...)
|
||||
final := Chain(chain...)(func(ctx *Context) error {
|
||||
if shouldUseUpstreamAction(ctx, route) && route.Upstream != nil {
|
||||
return route.Upstream(ctx)
|
||||
}
|
||||
if shouldUseUpstreamAction(ctx, route) && ctx.UpstreamURL != nil {
|
||||
return fmt.Errorf("route %s is missing upstream action while request targets upstream %s", route.Name, ctx.UpstreamURL.String())
|
||||
}
|
||||
if route.Local != nil {
|
||||
return route.Local(ctx)
|
||||
}
|
||||
@@ -168,14 +155,6 @@ func (app *App) buildRouteHandler(route Route) http.HandlerFunc {
|
||||
}
|
||||
}
|
||||
|
||||
func shouldUseUpstreamAction(ctx *Context, route Route) bool {
|
||||
_ = route
|
||||
if ctx == nil {
|
||||
return false
|
||||
}
|
||||
return ctx.Mode == ModeUpstream
|
||||
}
|
||||
|
||||
func Chain(middlewares ...Middleware) Middleware {
|
||||
return func(final HandlerFunc) HandlerFunc {
|
||||
wrapped := final
|
||||
|
||||
@@ -2,6 +2,7 @@ package upstream
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"strings"
|
||||
@@ -19,7 +20,7 @@ type CompatRouteConfig struct {
|
||||
ConsoleLog bool
|
||||
}
|
||||
|
||||
func DirectAction(deps Dependencies, cfg CompatRouteConfig) server.HandlerFunc {
|
||||
func ForwardAction(deps Dependencies, cfg CompatRouteConfig) server.HandlerFunc {
|
||||
return func(ctx *server.Context) error {
|
||||
reqCtx, route, err := newCompatRouteObjects(ctx, deps, cfg)
|
||||
if err != nil {
|
||||
@@ -29,6 +30,34 @@ func DirectAction(deps Dependencies, cfg CompatRouteConfig) server.HandlerFunc {
|
||||
}
|
||||
}
|
||||
|
||||
// AuthenticatedForwardAction forwards a Cursor control-plane request with the
|
||||
// independent desktop account after the local-mode identity rewrite has run.
|
||||
func AuthenticatedForwardAction(deps Dependencies, cfg CompatRouteConfig, authorizationProvider AuthorizationProvider) server.HandlerFunc {
|
||||
return func(ctx *server.Context) error {
|
||||
reqCtx, _, err := newCompatRouteObjects(ctx, deps, cfg)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if reqCtx == nil || reqCtx.Request == nil {
|
||||
return fmt.Errorf("Cursor 控制面请求上下文无效")
|
||||
}
|
||||
if authorizationProvider == nil {
|
||||
return fmt.Errorf("Cursor 账号服务未初始化")
|
||||
}
|
||||
authorization, err := authorizationProvider.Authorization(reqCtx.Request.Context())
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
_, err = ForwardToUpstream(reqCtx, ForwardOptions{
|
||||
PatchHeaders: func(headers http.Header) {
|
||||
headers.Set("Authorization", authorization)
|
||||
headers.Set("x-cursor-checksum", BuildCursorChecksum(authorization))
|
||||
},
|
||||
})
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
func FixedStatusAction(deps Dependencies, cfg CompatRouteConfig) server.HandlerFunc {
|
||||
return func(ctx *server.Context) error {
|
||||
reqCtx, route, err := newCompatRouteObjects(ctx, deps, cfg)
|
||||
@@ -133,7 +162,6 @@ func newCompatRouteObjects(ctx *server.Context, deps Dependencies, cfg CompatRou
|
||||
Headers: ctx.Request.Header.Clone(),
|
||||
ContentType: strings.TrimSpace(ctx.Request.Header.Get("content-type")),
|
||||
RequestBody: body,
|
||||
Mode: ctx.Mode,
|
||||
Deps: &deps,
|
||||
HTTPRequestID: resolveHTTPRequestID(ctx.Request),
|
||||
}
|
||||
|
||||
@@ -15,7 +15,6 @@ import (
|
||||
"time"
|
||||
|
||||
"cursor/gen/aiserverv1"
|
||||
"cursor/internal/backend/server"
|
||||
"cursor/internal/logger"
|
||||
"cursor/internal/netproxy"
|
||||
legacyruntime "cursor/internal/runtime"
|
||||
@@ -87,7 +86,7 @@ func buildUpstreamRequest(reqCtx *RequestContext, body []byte, options ForwardOp
|
||||
}
|
||||
upstreamRequest.Host = reqCtx.TargetURL.Host
|
||||
|
||||
if reqCtx.Mode == server.ModeLocal && shouldRewriteHost(reqCtx.TargetURL.Hostname()) {
|
||||
if shouldRewriteHost(reqCtx.TargetURL.Hostname()) {
|
||||
auth := formatBearerAuthorization(legacyruntime.LocalRelayToken)
|
||||
if auth == "" {
|
||||
return nil, nil, legacyruntime.ErrInvalidSystemSetting
|
||||
|
||||
@@ -20,6 +20,13 @@ type SystemSettingService interface {
|
||||
ResolveModelAdapters(context.Context) ([]legacyruntime.ModelAdapterConfig, error)
|
||||
}
|
||||
|
||||
// AuthorizationProvider supplies the independent Cursor account used only by
|
||||
// official control-plane requests such as Plugins, Skills, and MCP registry.
|
||||
type AuthorizationProvider interface {
|
||||
Authorization(context.Context) (string, error)
|
||||
SignedIn() bool
|
||||
}
|
||||
|
||||
type HTTPClient interface {
|
||||
Do(req *http.Request) (*http.Response, error)
|
||||
}
|
||||
@@ -41,7 +48,6 @@ type RequestContext struct {
|
||||
Headers http.Header
|
||||
ContentType string
|
||||
RequestBody []byte
|
||||
Mode server.ExecutionMode
|
||||
Deps *Dependencies
|
||||
HTTPRequestID string
|
||||
}
|
||||
|
||||
@@ -24,6 +24,9 @@ type ModelAdapterTestResult = client.ModelAdapterTestResult
|
||||
// ModelAdapterTestResultsPayload 定义测速结果事件载荷。
|
||||
type ModelAdapterTestResultsPayload = client.ModelAdapterTestResultsPayload
|
||||
|
||||
// CursorAccountStatus 是可安全展示给桌面前端的独立 Cursor 账号状态。
|
||||
type CursorAccountStatus = client.CursorAccountStatus
|
||||
|
||||
// LicenseActionRequest 定义了当前模块中的 LicenseActionRequest 类型。
|
||||
type LicenseActionRequest = client.LicenseActionRequest
|
||||
|
||||
@@ -91,6 +94,21 @@ func (s *ProxyService) SaveUserConfig(cfg UserConfig) error {
|
||||
return s.core.SaveUserConfig(cfg)
|
||||
}
|
||||
|
||||
// GetCursorAccountStatus 返回 cursor-byok 独立 Cursor 账号的脱敏状态。
|
||||
func (s *ProxyService) GetCursorAccountStatus() CursorAccountStatus {
|
||||
return s.core.GetCursorAccountStatus()
|
||||
}
|
||||
|
||||
// StartCursorAccountLogin 打开官方浏览器登录并异步等待结果。
|
||||
func (s *ProxyService) StartCursorAccountLogin() (CursorAccountStatus, error) {
|
||||
return s.core.StartCursorAccountLogin()
|
||||
}
|
||||
|
||||
// DisconnectCursorAccount 只断开 cursor-byok 自己的账号。
|
||||
func (s *ProxyService) DisconnectCursorAccount() (CursorAccountStatus, error) {
|
||||
return s.core.DisconnectCursorAccount()
|
||||
}
|
||||
|
||||
// TestModelAdapter 用于处理与 TestModelAdapter 相关的逻辑。
|
||||
func (s *ProxyService) TestModelAdapter(adapter ModelAdapterConfig) (ModelAdapterTestResult, error) {
|
||||
return s.core.TestModelAdapter(adapter)
|
||||
|
||||
@@ -0,0 +1,35 @@
|
||||
package client
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"cursor/internal/cursoraccount"
|
||||
)
|
||||
|
||||
type CursorAccountStatus = cursoraccount.Status
|
||||
|
||||
func (s *ProxyService) GetCursorAccountStatus() CursorAccountStatus {
|
||||
if s == nil || s.cursorAccount == nil {
|
||||
return CursorAccountStatus{State: cursoraccount.StateSignedOut}
|
||||
}
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
defer cancel()
|
||||
s.cursorAccount.EnsureEmail(ctx)
|
||||
return s.cursorAccount.Status()
|
||||
}
|
||||
|
||||
func (s *ProxyService) StartCursorAccountLogin() (CursorAccountStatus, error) {
|
||||
if s == nil || s.cursorAccount == nil {
|
||||
return CursorAccountStatus{State: cursoraccount.StateError}, fmt.Errorf("Cursor 账号服务未初始化")
|
||||
}
|
||||
return s.cursorAccount.StartLogin()
|
||||
}
|
||||
|
||||
func (s *ProxyService) DisconnectCursorAccount() (CursorAccountStatus, error) {
|
||||
if s == nil || s.cursorAccount == nil {
|
||||
return CursorAccountStatus{State: cursoraccount.StateSignedOut}, nil
|
||||
}
|
||||
return s.cursorAccount.Disconnect()
|
||||
}
|
||||
@@ -265,6 +265,9 @@ func (s *ProxyService) ShutdownForQuit() {
|
||||
finalErr = errors.Join(finalErr, err)
|
||||
}
|
||||
}
|
||||
if s.cursorAccount != nil {
|
||||
s.cursorAccount.Shutdown()
|
||||
}
|
||||
if finalErr != nil {
|
||||
s.setLastError(finalErr)
|
||||
}
|
||||
|
||||
@@ -4,6 +4,7 @@ import (
|
||||
"context"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"path/filepath"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
@@ -11,6 +12,7 @@ import (
|
||||
backend "cursor/internal/backend"
|
||||
serverconfig "cursor/internal/backend/server/config"
|
||||
"cursor/internal/certs"
|
||||
"cursor/internal/cursoraccount"
|
||||
"cursor/internal/logger"
|
||||
"cursor/internal/mitm"
|
||||
"cursor/internal/netproxy"
|
||||
@@ -35,6 +37,8 @@ type ProxyService struct {
|
||||
certManager *certs.Manager
|
||||
// backendHost 表示当前嵌入式 backend 服务。
|
||||
backendHost *backend.Host
|
||||
// cursorAccount 持有仅供插件、Skills 和 MCP 控制面使用的真实 Cursor 身份。
|
||||
cursorAccount *cursoraccount.Manager
|
||||
|
||||
// mu 表示当前声明中的 mu。
|
||||
mu sync.RWMutex
|
||||
@@ -84,8 +88,12 @@ func NewProxyService(proxy *mitm.ProxyServer, certManager *certs.Manager, caCert
|
||||
publicClient: netproxy.NewHTTPClient(publicAPITimeout),
|
||||
modelTestResults: make(map[string]ModelAdapterTestResult),
|
||||
}
|
||||
service.cursorAccount = cursoraccount.NewManager(
|
||||
filepath.Join(appdata.DataRootPath(), "cursor-account.json"),
|
||||
netproxy.NewHTTPClient(publicAPITimeout),
|
||||
)
|
||||
service.store = serverconfig.NewStore(service.configPath, service.logsRoot)
|
||||
host, err := backend.NewHost(service.store)
|
||||
host, err := backend.NewHost(service.store, service.cursorAccount)
|
||||
if err != nil {
|
||||
logger.Errorf("init backend host failed: %v", err)
|
||||
} else {
|
||||
@@ -101,7 +109,7 @@ func (s *ProxyService) ensureBackendHost() error {
|
||||
if s.backendHost != nil {
|
||||
return nil
|
||||
}
|
||||
host, err := backend.NewHost(s.store)
|
||||
host, err := backend.NewHost(s.store, s.cursorAccount)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
@@ -0,0 +1,589 @@
|
||||
package cursoraccount
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"crypto/sha256"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"cursor/gen/aiserverv1"
|
||||
"cursor/internal/backend/server/upstream"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"github.com/pkg/browser"
|
||||
"google.golang.org/protobuf/proto"
|
||||
)
|
||||
|
||||
const (
|
||||
StateSignedOut = "signed_out"
|
||||
StateWaiting = "waiting"
|
||||
StateSignedIn = "signed_in"
|
||||
StateError = "error"
|
||||
|
||||
websiteURL = "https://cursor.com"
|
||||
backendURL = "https://api2.cursor.sh"
|
||||
authClientID = "KbZUR41cY7W6zRSdpSUJ7I7mLYBKOCmB"
|
||||
loginTimeout = 10 * time.Minute
|
||||
pollInterval = time.Second
|
||||
refreshMargin = 2 * time.Minute
|
||||
)
|
||||
|
||||
var ErrNotSignedIn = errors.New("尚未在 cursor-byok 中登录 Cursor 账号")
|
||||
|
||||
// Status 是可安全返回给前端的脱敏账号状态。
|
||||
type Status struct {
|
||||
State string `json:"state"`
|
||||
AuthID string `json:"authId"`
|
||||
Email string `json:"email"`
|
||||
Error string `json:"error"`
|
||||
}
|
||||
|
||||
type credentials struct {
|
||||
AccessToken string `json:"accessToken"`
|
||||
RefreshToken string `json:"refreshToken"`
|
||||
AuthID string `json:"authId"`
|
||||
Email string `json:"email,omitempty"`
|
||||
}
|
||||
|
||||
type pollResponse struct {
|
||||
AccessToken string `json:"accessToken"`
|
||||
RefreshToken string `json:"refreshToken"`
|
||||
AuthID string `json:"authId"`
|
||||
}
|
||||
|
||||
type refreshResponse struct {
|
||||
AccessToken string `json:"access_token"`
|
||||
RefreshToken string `json:"refresh_token"`
|
||||
ShouldLogout bool `json:"shouldLogout"`
|
||||
}
|
||||
|
||||
// Manager 持有 cursor-byok 自己的 Cursor 登录态,不读写 Cursor 客户端状态库。
|
||||
type Manager struct {
|
||||
path string
|
||||
client *http.Client
|
||||
|
||||
mu sync.RWMutex
|
||||
credentials credentials
|
||||
state string
|
||||
lastError string
|
||||
loginCancel context.CancelFunc
|
||||
loginGeneration uint64
|
||||
|
||||
refreshMu sync.Mutex
|
||||
}
|
||||
|
||||
func NewManager(path string, client *http.Client) *Manager {
|
||||
if client == nil {
|
||||
client = &http.Client{Timeout: 15 * time.Second}
|
||||
}
|
||||
manager := &Manager{
|
||||
path: strings.TrimSpace(path),
|
||||
client: client,
|
||||
state: StateSignedOut,
|
||||
}
|
||||
if err := manager.load(); err != nil {
|
||||
manager.state = StateError
|
||||
manager.lastError = fmt.Sprintf("读取 Cursor 账号凭据失败: %v", err)
|
||||
}
|
||||
return manager
|
||||
}
|
||||
|
||||
func (manager *Manager) Status() Status {
|
||||
if manager == nil {
|
||||
return Status{State: StateSignedOut}
|
||||
}
|
||||
manager.mu.RLock()
|
||||
defer manager.mu.RUnlock()
|
||||
return Status{
|
||||
State: manager.state,
|
||||
AuthID: manager.credentials.AuthID,
|
||||
Email: manager.credentials.Email,
|
||||
Error: manager.lastError,
|
||||
}
|
||||
}
|
||||
|
||||
// EnsureEmail backfills a human-readable identity for credentials saved by
|
||||
// builds that only persisted authId. Profile lookup failure does not invalidate
|
||||
// an otherwise usable control-plane login.
|
||||
func (manager *Manager) EnsureEmail(ctx context.Context) {
|
||||
if manager == nil || !manager.SignedIn() {
|
||||
return
|
||||
}
|
||||
current, generation := manager.snapshotCredentials()
|
||||
if strings.TrimSpace(current.Email) != "" {
|
||||
return
|
||||
}
|
||||
authorization, err := manager.Authorization(ctx)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
profile, err := manager.fetchProfile(ctx, authorization)
|
||||
if err != nil || strings.TrimSpace(profile.GetEmail()) == "" {
|
||||
return
|
||||
}
|
||||
current, currentGeneration := manager.snapshotCredentials()
|
||||
if currentGeneration != generation {
|
||||
return
|
||||
}
|
||||
current.Email = strings.TrimSpace(profile.GetEmail())
|
||||
_ = manager.commitCredentials(generation, current)
|
||||
}
|
||||
|
||||
func (manager *Manager) SignedIn() bool {
|
||||
if manager == nil {
|
||||
return false
|
||||
}
|
||||
manager.mu.RLock()
|
||||
defer manager.mu.RUnlock()
|
||||
return manager.state == StateSignedIn && strings.TrimSpace(manager.credentials.AccessToken) != ""
|
||||
}
|
||||
|
||||
// StartLogin 启动官方浏览器 PKCE 登录,并在后台等待登录结果。
|
||||
func (manager *Manager) StartLogin() (Status, error) {
|
||||
if manager == nil {
|
||||
return Status{State: StateError}, fmt.Errorf("Cursor 账号服务未初始化")
|
||||
}
|
||||
verifierBytes := make([]byte, 32)
|
||||
if _, err := rand.Read(verifierBytes); err != nil {
|
||||
return manager.Status(), fmt.Errorf("生成 Cursor 登录校验码失败: %w", err)
|
||||
}
|
||||
verifier := base64.RawURLEncoding.EncodeToString(verifierBytes)
|
||||
challengeBytes := sha256.Sum256([]byte(verifier))
|
||||
challenge := base64.RawURLEncoding.EncodeToString(challengeBytes[:])
|
||||
loginID := uuid.NewString()
|
||||
|
||||
loginURL, err := buildLoginURL(loginID, challenge)
|
||||
if err != nil {
|
||||
return manager.Status(), err
|
||||
}
|
||||
ctx, cancel := context.WithTimeout(context.Background(), loginTimeout)
|
||||
|
||||
manager.mu.Lock()
|
||||
if manager.loginCancel != nil {
|
||||
manager.loginCancel()
|
||||
}
|
||||
manager.loginGeneration++
|
||||
generation := manager.loginGeneration
|
||||
manager.loginCancel = cancel
|
||||
manager.state = StateWaiting
|
||||
manager.lastError = ""
|
||||
manager.mu.Unlock()
|
||||
|
||||
if err := browser.OpenURL(loginURL); err != nil {
|
||||
cancel()
|
||||
manager.finishWithError(generation, fmt.Sprintf("打开 Cursor 登录页面失败: %v", err))
|
||||
return manager.Status(), err
|
||||
}
|
||||
|
||||
go manager.pollLogin(ctx, generation, loginID, verifier)
|
||||
return manager.Status(), nil
|
||||
}
|
||||
|
||||
// Disconnect 只清除 cursor-byok 自己保存的账号,不调用 Cursor 客户端 logout。
|
||||
func (manager *Manager) Disconnect() (Status, error) {
|
||||
if manager == nil {
|
||||
return Status{State: StateSignedOut}, nil
|
||||
}
|
||||
manager.mu.Lock()
|
||||
manager.loginGeneration++
|
||||
if manager.loginCancel != nil {
|
||||
manager.loginCancel()
|
||||
manager.loginCancel = nil
|
||||
}
|
||||
manager.credentials = credentials{}
|
||||
manager.state = StateSignedOut
|
||||
manager.lastError = ""
|
||||
manager.mu.Unlock()
|
||||
|
||||
err := os.Remove(manager.path)
|
||||
if err != nil && !errors.Is(err, os.ErrNotExist) {
|
||||
manager.mu.Lock()
|
||||
manager.state = StateError
|
||||
manager.lastError = fmt.Sprintf("清除 Cursor 账号凭据失败: %v", err)
|
||||
manager.mu.Unlock()
|
||||
return manager.Status(), err
|
||||
}
|
||||
return manager.Status(), nil
|
||||
}
|
||||
|
||||
func (manager *Manager) Shutdown() {
|
||||
if manager == nil {
|
||||
return
|
||||
}
|
||||
manager.mu.Lock()
|
||||
manager.loginGeneration++
|
||||
if manager.loginCancel != nil {
|
||||
manager.loginCancel()
|
||||
manager.loginCancel = nil
|
||||
}
|
||||
manager.mu.Unlock()
|
||||
}
|
||||
|
||||
// Authorization 返回官方控制面请求使用的真实 Cursor Bearer 身份。
|
||||
func (manager *Manager) Authorization(ctx context.Context) (string, error) {
|
||||
if manager == nil {
|
||||
return "", ErrNotSignedIn
|
||||
}
|
||||
manager.refreshMu.Lock()
|
||||
defer manager.refreshMu.Unlock()
|
||||
|
||||
creds, generation := manager.snapshotCredentials()
|
||||
if strings.TrimSpace(creds.AccessToken) == "" {
|
||||
return "", ErrNotSignedIn
|
||||
}
|
||||
if !tokenNeedsRefresh(creds.AccessToken, time.Now()) {
|
||||
return bearer(creds.AccessToken), nil
|
||||
}
|
||||
if strings.TrimSpace(creds.RefreshToken) == "" {
|
||||
manager.setAuthorizationError(generation, "Cursor 登录已过期,请重新登录")
|
||||
return "", fmt.Errorf("Cursor 登录已过期且没有刷新令牌")
|
||||
}
|
||||
|
||||
updated, shouldLogout, err := manager.refresh(ctx, creds)
|
||||
if err != nil {
|
||||
manager.setAuthorizationError(generation, fmt.Sprintf("刷新 Cursor 登录失败: %v", err))
|
||||
return "", err
|
||||
}
|
||||
if shouldLogout {
|
||||
manager.invalidateAuthorization(generation, "Cursor 登录已失效,请重新登录")
|
||||
return "", ErrNotSignedIn
|
||||
}
|
||||
if err := manager.commitCredentials(generation, updated); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return bearer(updated.AccessToken), nil
|
||||
}
|
||||
|
||||
func (manager *Manager) pollLogin(ctx context.Context, generation uint64, loginID string, verifier string) {
|
||||
defer func() {
|
||||
manager.mu.Lock()
|
||||
if manager.loginGeneration == generation {
|
||||
manager.loginCancel = nil
|
||||
}
|
||||
manager.mu.Unlock()
|
||||
}()
|
||||
|
||||
for {
|
||||
result, pending, err := manager.pollOnce(ctx, loginID, verifier)
|
||||
if err == nil && !pending {
|
||||
creds := credentials{
|
||||
AccessToken: strings.TrimSpace(result.AccessToken),
|
||||
RefreshToken: strings.TrimSpace(result.RefreshToken),
|
||||
AuthID: strings.TrimSpace(result.AuthID),
|
||||
}
|
||||
if creds.AccessToken == "" {
|
||||
manager.finishWithError(generation, "Cursor 登录响应缺少 access token")
|
||||
return
|
||||
}
|
||||
if profile, profileErr := manager.fetchProfile(ctx, bearer(creds.AccessToken)); profileErr == nil {
|
||||
creds.Email = strings.TrimSpace(profile.GetEmail())
|
||||
}
|
||||
_ = manager.commitCredentials(generation, creds)
|
||||
return
|
||||
}
|
||||
if err != nil && !isRetryablePollError(err) {
|
||||
manager.finishWithError(generation, fmt.Sprintf("Cursor 登录失败: %v", err))
|
||||
return
|
||||
}
|
||||
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
if errors.Is(ctx.Err(), context.DeadlineExceeded) {
|
||||
manager.finishWithError(generation, "Cursor 登录等待超时,请重试")
|
||||
}
|
||||
return
|
||||
case <-time.After(pollInterval):
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (manager *Manager) fetchProfile(ctx context.Context, authorization string) (*aiserverv1.GetMeResponse, error) {
|
||||
body, err := proto.Marshal(&aiserverv1.GetMeRequest{})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodPost, backendURL+"/aiserver.v1.DashboardService/GetMe", bytes.NewReader(body))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
req.Header.Set("authorization", authorization)
|
||||
req.Header.Set("x-cursor-checksum", upstream.BuildCursorChecksum(authorization))
|
||||
req.Header.Set("content-type", "application/proto")
|
||||
req.Header.Set("accept", "application/proto")
|
||||
req.Header.Set("connect-protocol-version", "1")
|
||||
resp, err := manager.client.Do(req)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
responseBody, err := io.ReadAll(io.LimitReader(resp.Body, 1024*1024))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
|
||||
return nil, fmt.Errorf("GetMe 返回 HTTP %d", resp.StatusCode)
|
||||
}
|
||||
profile := &aiserverv1.GetMeResponse{}
|
||||
if err := proto.Unmarshal(responseBody, profile); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return profile, nil
|
||||
}
|
||||
|
||||
func (manager *Manager) pollOnce(ctx context.Context, loginID string, verifier string) (pollResponse, bool, error) {
|
||||
endpoint, err := url.Parse(backendURL + "/auth/poll")
|
||||
if err != nil {
|
||||
return pollResponse{}, false, err
|
||||
}
|
||||
query := endpoint.Query()
|
||||
query.Set("uuid", loginID)
|
||||
query.Set("verifier", verifier)
|
||||
endpoint.RawQuery = query.Encode()
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, endpoint.String(), nil)
|
||||
if err != nil {
|
||||
return pollResponse{}, false, err
|
||||
}
|
||||
resp, err := manager.client.Do(req)
|
||||
if err != nil {
|
||||
return pollResponse{}, false, err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
if resp.StatusCode == http.StatusNotFound {
|
||||
_, _ = io.Copy(io.Discard, io.LimitReader(resp.Body, 64*1024))
|
||||
return pollResponse{}, true, nil
|
||||
}
|
||||
body, err := io.ReadAll(io.LimitReader(resp.Body, 1024*1024))
|
||||
if err != nil {
|
||||
return pollResponse{}, false, err
|
||||
}
|
||||
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
|
||||
return pollResponse{}, false, fmt.Errorf("登录服务返回 HTTP %d", resp.StatusCode)
|
||||
}
|
||||
result := pollResponse{}
|
||||
if err := json.Unmarshal(body, &result); err != nil {
|
||||
return pollResponse{}, false, fmt.Errorf("解析登录响应失败: %w", err)
|
||||
}
|
||||
return result, false, nil
|
||||
}
|
||||
|
||||
func (manager *Manager) refresh(ctx context.Context, current credentials) (credentials, bool, error) {
|
||||
payload, err := json.Marshal(map[string]string{
|
||||
"grant_type": "refresh_token",
|
||||
"client_id": authClientID,
|
||||
"refresh_token": current.RefreshToken,
|
||||
})
|
||||
if err != nil {
|
||||
return credentials{}, false, err
|
||||
}
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodPost, backendURL+"/oauth/token", bytes.NewReader(payload))
|
||||
if err != nil {
|
||||
return credentials{}, false, err
|
||||
}
|
||||
req.Header.Set("content-type", "application/json")
|
||||
resp, err := manager.client.Do(req)
|
||||
if err != nil {
|
||||
return credentials{}, false, err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
body, err := io.ReadAll(io.LimitReader(resp.Body, 1024*1024))
|
||||
if err != nil {
|
||||
return credentials{}, false, err
|
||||
}
|
||||
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
|
||||
return credentials{}, false, fmt.Errorf("刷新服务返回 HTTP %d", resp.StatusCode)
|
||||
}
|
||||
result := refreshResponse{}
|
||||
if err := json.Unmarshal(body, &result); err != nil {
|
||||
return credentials{}, false, fmt.Errorf("解析刷新响应失败: %w", err)
|
||||
}
|
||||
if result.ShouldLogout {
|
||||
return credentials{}, true, nil
|
||||
}
|
||||
if strings.TrimSpace(result.AccessToken) == "" {
|
||||
return credentials{}, false, fmt.Errorf("刷新响应缺少 access token")
|
||||
}
|
||||
current.AccessToken = strings.TrimSpace(result.AccessToken)
|
||||
if strings.TrimSpace(result.RefreshToken) != "" {
|
||||
current.RefreshToken = strings.TrimSpace(result.RefreshToken)
|
||||
}
|
||||
return current, false, nil
|
||||
}
|
||||
|
||||
func (manager *Manager) load() error {
|
||||
if manager.path == "" {
|
||||
return fmt.Errorf("Cursor 账号凭据路径为空")
|
||||
}
|
||||
data, err := os.ReadFile(manager.path)
|
||||
if errors.Is(err, os.ErrNotExist) {
|
||||
return nil
|
||||
}
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
loaded := credentials{}
|
||||
if err := json.Unmarshal(data, &loaded); err != nil {
|
||||
return err
|
||||
}
|
||||
loaded.AccessToken = strings.TrimSpace(loaded.AccessToken)
|
||||
loaded.RefreshToken = strings.TrimSpace(loaded.RefreshToken)
|
||||
loaded.AuthID = strings.TrimSpace(loaded.AuthID)
|
||||
loaded.Email = strings.TrimSpace(loaded.Email)
|
||||
if loaded.AccessToken == "" {
|
||||
return nil
|
||||
}
|
||||
manager.credentials = loaded
|
||||
manager.state = StateSignedIn
|
||||
return nil
|
||||
}
|
||||
|
||||
func (manager *Manager) save(value credentials) error {
|
||||
if manager.path == "" {
|
||||
return fmt.Errorf("Cursor 账号凭据路径为空")
|
||||
}
|
||||
if err := os.MkdirAll(filepath.Dir(manager.path), 0o700); err != nil {
|
||||
return err
|
||||
}
|
||||
data, err := json.MarshalIndent(value, "", " ")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
tempPath := manager.path + ".tmp"
|
||||
if err := os.WriteFile(tempPath, append(data, '\n'), 0o600); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := os.Chmod(tempPath, 0o600); err != nil {
|
||||
_ = os.Remove(tempPath)
|
||||
return err
|
||||
}
|
||||
if err := os.Rename(tempPath, manager.path); err != nil {
|
||||
_ = os.Remove(tempPath)
|
||||
return err
|
||||
}
|
||||
return os.Chmod(manager.path, 0o600)
|
||||
}
|
||||
|
||||
func (manager *Manager) snapshotCredentials() (credentials, uint64) {
|
||||
manager.mu.RLock()
|
||||
defer manager.mu.RUnlock()
|
||||
return manager.credentials, manager.loginGeneration
|
||||
}
|
||||
|
||||
func (manager *Manager) finishWithError(generation uint64, message string) {
|
||||
manager.mu.Lock()
|
||||
defer manager.mu.Unlock()
|
||||
if manager.loginGeneration != generation {
|
||||
return
|
||||
}
|
||||
manager.state = StateError
|
||||
manager.lastError = strings.TrimSpace(message)
|
||||
}
|
||||
|
||||
func (manager *Manager) commitCredentials(generation uint64, value credentials) error {
|
||||
manager.mu.Lock()
|
||||
defer manager.mu.Unlock()
|
||||
if manager.loginGeneration != generation {
|
||||
return ErrNotSignedIn
|
||||
}
|
||||
if err := manager.save(value); err != nil {
|
||||
manager.state = StateError
|
||||
manager.lastError = fmt.Sprintf("保存 Cursor 登录凭据失败: %v", err)
|
||||
return err
|
||||
}
|
||||
manager.credentials = value
|
||||
manager.state = StateSignedIn
|
||||
manager.lastError = ""
|
||||
return nil
|
||||
}
|
||||
|
||||
func (manager *Manager) setAuthorizationError(generation uint64, message string) {
|
||||
manager.mu.Lock()
|
||||
defer manager.mu.Unlock()
|
||||
if manager.loginGeneration != generation {
|
||||
return
|
||||
}
|
||||
manager.state = StateError
|
||||
manager.lastError = strings.TrimSpace(message)
|
||||
}
|
||||
|
||||
func (manager *Manager) invalidateAuthorization(generation uint64, message string) {
|
||||
manager.mu.Lock()
|
||||
defer manager.mu.Unlock()
|
||||
if manager.loginGeneration != generation {
|
||||
return
|
||||
}
|
||||
manager.loginGeneration++
|
||||
manager.credentials = credentials{}
|
||||
manager.state = StateError
|
||||
manager.lastError = strings.TrimSpace(message)
|
||||
_ = os.Remove(manager.path)
|
||||
}
|
||||
|
||||
func buildLoginURL(loginID string, challenge string) (string, error) {
|
||||
parsed, err := url.Parse(websiteURL + "/loginDeepControl")
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
query := parsed.Query()
|
||||
query.Set("challenge", challenge)
|
||||
query.Set("uuid", loginID)
|
||||
query.Set("mode", "login")
|
||||
query.Set("supportsSelectedTeamLogin", "true")
|
||||
parsed.RawQuery = query.Encode()
|
||||
return parsed.String(), nil
|
||||
}
|
||||
|
||||
func bearer(token string) string {
|
||||
value := strings.TrimSpace(token)
|
||||
if strings.HasPrefix(strings.ToLower(value), "bearer ") {
|
||||
return value
|
||||
}
|
||||
return "Bearer " + value
|
||||
}
|
||||
|
||||
func tokenNeedsRefresh(token string, now time.Time) bool {
|
||||
parts := strings.Split(strings.TrimSpace(token), ".")
|
||||
if len(parts) < 2 {
|
||||
return false
|
||||
}
|
||||
payload, err := base64.RawURLEncoding.DecodeString(parts[1])
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
claims := struct {
|
||||
ExpiresAt json.Number `json:"exp"`
|
||||
}{}
|
||||
decoder := json.NewDecoder(bytes.NewReader(payload))
|
||||
decoder.UseNumber()
|
||||
if err := decoder.Decode(&claims); err != nil || claims.ExpiresAt == "" {
|
||||
return false
|
||||
}
|
||||
expiresAt, err := claims.ExpiresAt.Int64()
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
return !now.Add(refreshMargin).Before(time.Unix(expiresAt, 0))
|
||||
}
|
||||
|
||||
func isRetryablePollError(err error) bool {
|
||||
if err == nil {
|
||||
return false
|
||||
}
|
||||
var urlErr *url.Error
|
||||
if errors.As(err, &urlErr) {
|
||||
return true
|
||||
}
|
||||
message := strings.ToLower(err.Error())
|
||||
return strings.Contains(message, "http 429") || strings.Contains(message, "http 5")
|
||||
}
|
||||
Reference in New Issue
Block a user