diff --git a/api/ai_session.go b/api/ai_session.go index 81e3df9f..bae68332 100644 --- a/api/ai_session.go +++ b/api/ai_session.go @@ -72,7 +72,7 @@ func registerAISessionRoutes(app *fiber.App, group *huma.Group) { }, func(op *huma.Operation) { op.OperationID = "ai-session-list-by-project" op.Summary = "获取项目的AI助手会话列表" - op.Description = "返回指定项目目录下的 Claude Code 和 Codex 会话信息" + op.Description = "返回指定项目目录下的 Claude Code、Codex 和 Pi 会话信息" op.Tags = []string{aiSessionTag} }) @@ -96,7 +96,7 @@ func registerAISessionRoutes(app *fiber.App, group *huma.Group) { }, func(op *huma.Operation) { op.OperationID = "ai-session-list-by-path" op.Summary = "通过路径获取AI助手会话列表" - op.Description = "根据目录路径返回 Claude Code 和 Codex 会话信息" + op.Description = "根据目录路径返回 Claude Code、Codex 和 Pi 会话信息" op.Tags = []string{aiSessionTag} }) diff --git a/api/api.go b/api/api.go index a624c517..a4224ef2 100644 --- a/api/api.go +++ b/api/api.go @@ -122,10 +122,11 @@ func Init(ctx context.Context, cfg *utils.AppConfig, assets embed.FS, info *AppI theLogger.Error("failed to initialize web session manager", zap.Error(err)) return err } + defer webSessionManager.StopAllPiRuntimes() registerAuthRoutes(app, cfg) registerHealthRoutes(app, humaAPI) - registerProjectRoutes(v1) + registerProjectRoutes(v1, webSessionManager) registerWorktreeRoutes(v1, cfg) registerBranchRoutes(v1) registerTaskRoutes(v1) diff --git a/api/project.go b/api/project.go index 91b4fc9e..e5369013 100644 --- a/api/project.go +++ b/api/project.go @@ -9,6 +9,7 @@ import ( "code-kanban/api/h" "code-kanban/model" + "code-kanban/service/websession" ) const projectTag = "project-项目管理" @@ -43,7 +44,7 @@ type projectAccessInput struct { ID string `path:"id"` } -func registerProjectRoutes(group *huma.Group) { +func registerProjectRoutes(group *huma.Group, webSessionManager *websession.Manager) { service := model.NewProjectService() huma.Post(group, "/projects/create", func(ctx context.Context, input *createProjectInput) (*h.ItemResponse[model.Project], error) { @@ -227,6 +228,9 @@ func registerProjectRoutes(group *huma.Group) { huma.Post(group, "/projects/{id}/delete", func(ctx context.Context, input *struct { ID string `path:"id"` }) (*h.MessageResponse, error) { + if webSessionManager != nil { + webSessionManager.StopProjectPiRuntimes(input.ID) + } if err := service.DeleteProject(ctx, input.ID); err != nil { if errors.Is(err, model.ErrDBNotInitialized) { return nil, huma.Error503ServiceUnavailable("database is not initialized") diff --git a/api/project_agent_trust_test.go b/api/project_agent_trust_test.go new file mode 100644 index 00000000..958ea01e --- /dev/null +++ b/api/project_agent_trust_test.go @@ -0,0 +1,97 @@ +package api + +import ( + "io" + "net/http" + "net/http/httptest" + "path/filepath" + "strings" + "testing" + + "code-kanban/api/h" + "code-kanban/model" + "code-kanban/model/tables" + "code-kanban/service" + "code-kanban/service/websession" + "code-kanban/utils" + + "github.com/gofiber/fiber/v2" + "go.uber.org/zap" +) + +func TestProjectPiTrustRoutes(t *testing.T) { + model.DBClose() + if err := model.InitWithDSN(filepath.Join(t.TempDir(), "agent-trust.db"), 0, true); err != nil { + t.Fatalf("InitWithDSN: %v", err) + } + t.Cleanup(model.DBClose) + + project := &tables.ProjectTable{Name: "Trust API", Path: t.TempDir()} + project.Init() + if err := model.GetDB().Create(project).Error; err != nil { + t.Fatalf("create project: %v", err) + } + manager, err := websession.NewManager(websession.Config{DataDir: t.TempDir()}, zap.NewNop()) + if err != nil { + t.Fatalf("NewManager: %v", err) + } + + app := fiber.New(fiber.Config{Immutable: true}) + _, group := h.NewAPI(app, &utils.AppConfig{}) + registerWebSessionRoutes(app, group, manager, zap.NewNop()) + path := "/api/v1/projects/" + project.ID + "/agent-trust/pi" + + assertProjectPiTrustResponse(t, app, http.MethodGet, path, "", http.StatusOK, `"trusted":false`) + assertProjectPiTrustResponse( + t, + app, + http.MethodPost, + path, + `{"trusted":true,"path":"D:/forged"}`, + http.StatusOK, + `"trusted":true`, + ) + var trust tables.ProjectAgentTrustTable + if err := model.GetDB().Where("project_id = ? AND agent = ?", project.ID, "pi").First(&trust).Error; err != nil { + t.Fatalf("load trust record: %v", err) + } + wantTrustedPath, err := service.CanonicalAgentTrustPath(project.Path) + if err != nil { + t.Fatalf("canonical project path: %v", err) + } + if trust.TrustedPath != wantTrustedPath { + t.Fatalf("trusted path = %q, want server project path %q", trust.TrustedPath, wantTrustedPath) + } + assertProjectPiTrustResponse(t, app, http.MethodDelete, path, "", http.StatusOK, `"trusted":false`) +} + +func assertProjectPiTrustResponse( + t *testing.T, + app *fiber.App, + method string, + path string, + body string, + wantStatus int, + wantFragment string, +) { + t.Helper() + request := httptest.NewRequest(method, path, strings.NewReader(body)) + if body != "" { + request.Header.Set(fiber.HeaderContentType, fiber.MIMEApplicationJSON) + } + response, err := app.Test(request) + if err != nil { + t.Fatalf("%s %s: %v", method, path, err) + } + defer response.Body.Close() + payload, err := io.ReadAll(response.Body) + if err != nil { + t.Fatalf("read response: %v", err) + } + if response.StatusCode != wantStatus { + t.Fatalf("%s %s status = %d, want %d: %s", method, path, response.StatusCode, wantStatus, payload) + } + if !strings.Contains(string(payload), wantFragment) { + t.Fatalf("%s %s response %s does not contain %s", method, path, payload, wantFragment) + } +} diff --git a/api/web_session.go b/api/web_session.go index cdb3f087..584c96a0 100644 --- a/api/web_session.go +++ b/api/web_session.go @@ -12,6 +12,7 @@ import ( "code-kanban/api/h" "code-kanban/model" + "code-kanban/service" "code-kanban/service/websession" "github.com/danielgtaylor/huma/v2" @@ -42,6 +43,12 @@ type webSessionCountsResponse struct { } `json:"body"` } +type piTreeNavigateBody struct { + TargetID string `json:"targetId" minLength:"1"` + Revision string `json:"revision" minLength:"1"` + Summarize *bool `json:"summarize,omitempty"` +} + func registerWebSessionRoutes(app *fiber.App, group *huma.Group, manager *websession.Manager, logger *zap.Logger) { ctrl := &webSessionController{ manager: manager, @@ -58,6 +65,84 @@ func registerWebSessionRoutes(app *fiber.App, group *huma.Group, manager *webses } func (c *webSessionController) registerHTTP(app *fiber.App, group *huma.Group) { + huma.Get(group, "/projects/{projectId}/agent-trust/pi", func( + ctx context.Context, + input *struct { + ProjectID string `path:"projectId"` + }, + ) (*h.ItemResponse[service.ProjectAgentTrustStatus], error) { + item, err := c.manager.GetProjectPiTrust(ctx, input.ProjectID) + if err != nil { + switch { + case errors.Is(err, model.ErrProjectNotFound): + return nil, huma.Error404NotFound("project not found") + case errors.Is(err, model.ErrDBNotInitialized): + return nil, huma.Error503ServiceUnavailable("database is not initialized") + default: + return nil, huma.Error500InternalServerError("failed to load Pi project trust", err) + } + } + resp := h.NewItemResponse(item) + resp.Status = http.StatusOK + return resp, nil + }, func(op *huma.Operation) { + op.OperationID = "project-agent-trust-pi-get" + op.Summary = "获取项目 Pi 授权状态" + op.Tags = []string{webSessionTag} + }) + + huma.Post(group, "/projects/{projectId}/agent-trust/pi", func( + ctx context.Context, + input *struct { + ProjectID string `path:"projectId"` + }, + ) (*h.ItemResponse[service.ProjectAgentTrustStatus], error) { + item, err := c.manager.TrustProjectForPi(ctx, input.ProjectID) + if err != nil { + switch { + case errors.Is(err, model.ErrProjectNotFound): + return nil, huma.Error404NotFound("project not found") + case errors.Is(err, model.ErrDBNotInitialized): + return nil, huma.Error503ServiceUnavailable("database is not initialized") + default: + return nil, huma.Error400BadRequest(err.Error()) + } + } + resp := h.NewItemResponse(item) + resp.Status = http.StatusOK + return resp, nil + }, func(op *huma.Operation) { + op.OperationID = "project-agent-trust-pi-create" + op.Summary = "授权项目加载 Pi 本地资源" + op.Tags = []string{webSessionTag} + }) + + huma.Delete(group, "/projects/{projectId}/agent-trust/pi", func( + ctx context.Context, + input *struct { + ProjectID string `path:"projectId"` + }, + ) (*h.ItemResponse[service.ProjectAgentTrustStatus], error) { + item, err := c.manager.RevokeProjectPiTrust(ctx, input.ProjectID) + if err != nil { + switch { + case errors.Is(err, model.ErrProjectNotFound): + return nil, huma.Error404NotFound("project not found") + case errors.Is(err, model.ErrDBNotInitialized): + return nil, huma.Error503ServiceUnavailable("database is not initialized") + default: + return nil, huma.Error500InternalServerError("failed to revoke Pi project trust", err) + } + } + resp := h.NewItemResponse(item) + resp.Status = http.StatusOK + return resp, nil + }, func(op *huma.Operation) { + op.OperationID = "project-agent-trust-pi-delete" + op.Summary = "撤销项目 Pi 授权" + op.Tags = []string{webSessionTag} + }) + huma.Get(group, "/projects/{projectId}/web-sessions", func( ctx context.Context, input *struct { @@ -129,6 +214,104 @@ func (c *webSessionController) registerHTTP(app *fiber.App, group *huma.Group) { op.Tags = []string{webSessionTag} }) + huma.Get(group, "/projects/{projectId}/web-sessions/{sessionId}/tree", func( + ctx context.Context, + input *struct { + ProjectID string `path:"projectId"` + SessionID string `path:"sessionId"` + }, + ) (*h.ItemResponse[websession.PiTreeSnapshot], error) { + if err := c.requireProjectSession(ctx, input.ProjectID, input.SessionID); err != nil { + return nil, err + } + item, err := c.manager.GetPiSessionTree(ctx, input.SessionID) + if err != nil { + return nil, piTreeHTTPError(err) + } + resp := h.NewItemResponse(item) + resp.Status = http.StatusOK + return resp, nil + }, func(op *huma.Operation) { + op.OperationID = "web-session-tree-get" + op.Summary = "获取 Pi 会话历史树" + op.Tags = []string{webSessionTag} + }) + + huma.Post(group, "/projects/{projectId}/web-sessions/{sessionId}/tree/navigate", func( + ctx context.Context, + input *struct { + ProjectID string `path:"projectId"` + SessionID string `path:"sessionId"` + Body piTreeNavigateBody + }, + ) (*h.ItemResponse[websession.PiTreeNavigateResult], error) { + if err := c.requireProjectSession(ctx, input.ProjectID, input.SessionID); err != nil { + return nil, err + } + item, err := c.manager.NavigatePiSessionTree(ctx, input.SessionID, websession.PiTreeNavigateInput{ + TargetID: input.Body.TargetID, Revision: input.Body.Revision, + Summarize: input.Body.Summarize != nil && *input.Body.Summarize, + }) + if err != nil { + return nil, piTreeHTTPError(err) + } + resp := h.NewItemResponse(item) + resp.Status = http.StatusOK + return resp, nil + }, func(op *huma.Operation) { + op.OperationID = "web-session-tree-navigate" + op.Summary = "切换 Pi 会话历史分支" + op.Tags = []string{webSessionTag} + }) + + huma.Post(group, "/projects/{projectId}/web-sessions/{sessionId}/tree/fork", func( + ctx context.Context, + input *struct { + ProjectID string `path:"projectId"` + SessionID string `path:"sessionId"` + Body websession.PiTreeForkInput + }, + ) (*h.ItemResponse[websession.PiTreeCreateResult], error) { + if err := c.requireProjectSession(ctx, input.ProjectID, input.SessionID); err != nil { + return nil, err + } + item, err := c.manager.ForkPiSessionTree(ctx, input.SessionID, input.Body) + if err != nil { + return nil, piTreeHTTPError(err) + } + resp := h.NewItemResponse(item) + resp.Status = http.StatusCreated + return resp, nil + }, func(op *huma.Operation) { + op.OperationID = "web-session-tree-fork" + op.Summary = "从 Pi 历史节点创建新会话" + op.Tags = []string{webSessionTag} + }) + + huma.Post(group, "/projects/{projectId}/web-sessions/{sessionId}/tree/clone", func( + ctx context.Context, + input *struct { + ProjectID string `path:"projectId"` + SessionID string `path:"sessionId"` + Body websession.PiTreeCloneInput + }, + ) (*h.ItemResponse[websession.PiTreeCreateResult], error) { + if err := c.requireProjectSession(ctx, input.ProjectID, input.SessionID); err != nil { + return nil, err + } + item, err := c.manager.ClonePiSessionTree(ctx, input.SessionID, input.Body) + if err != nil { + return nil, piTreeHTTPError(err) + } + resp := h.NewItemResponse(item) + resp.Status = http.StatusCreated + return resp, nil + }, func(op *huma.Operation) { + op.OperationID = "web-session-tree-clone" + op.Summary = "克隆当前 Pi 会话分支" + op.Tags = []string{webSessionTag} + }) + huma.Get(group, "/projects/{projectId}/web-sessions/{sessionId}/history", func( ctx context.Context, input *struct { @@ -217,13 +400,13 @@ func (c *webSessionController) registerHTTP(app *fiber.App, group *huma.Group) { ProjectID string `path:"projectId"` }, ) (*h.ItemResponse[websession.ImportSourceList], error) { - item, err := c.manager.ListCodexImportSources(ctx, input.ProjectID) + item, err := c.manager.ListImportSources(ctx, input.ProjectID) if err != nil { switch { case errors.Is(err, model.ErrProjectNotFound): return nil, huma.Error404NotFound("project not found") default: - return nil, huma.Error500InternalServerError("failed to list codex import sources", err) + return nil, huma.Error500InternalServerError("failed to list import sources", err) } } resp := h.NewItemResponse(item) @@ -231,15 +414,15 @@ func (c *webSessionController) registerHTTP(app *fiber.App, group *huma.Group) { return resp, nil }, func(op *huma.Operation) { op.OperationID = "web-session-import-sources" - op.Summary = "获取 Codex 导入源列表" + op.Summary = "获取 AI 会话导入源列表" op.Tags = []string{webSessionTag} }) huma.Get(group, "/web-sessions/runtime-config", func( ctx context.Context, _ *struct{}, - ) (*h.ItemResponse[websession.CodexRuntimeConfig], error) { - resp := h.NewItemResponse(c.manager.GetCodexRuntimeConfigWithModels()) + ) (*h.ItemResponse[websession.WebSessionRuntimeConfig], error) { + resp := h.NewItemResponse(c.manager.GetWebSessionRuntimeConfigWithModels()) resp.Status = http.StatusOK return resp, nil }, func(op *huma.Operation) { @@ -374,6 +557,7 @@ func (c *webSessionController) registerHTTP(app *fiber.App, group *huma.Group) { input *struct { ProjectID string `path:"projectId"` Body struct { + Agent string `json:"agent,omitempty"` AISessionID string `json:"aiSessionId"` SessionID string `json:"sessionId,omitempty"` Mode string `json:"mode,omitempty"` @@ -384,27 +568,36 @@ func (c *webSessionController) registerHTTP(app *fiber.App, group *huma.Group) { item websession.ImportResult err error ) - if strings.TrimSpace(input.Body.SessionID) != "" { - item, err = c.manager.ImportCodexSessionBySessionID( - ctx, - input.ProjectID, - input.Body.SessionID, - websession.SyncMode(input.Body.Mode), - ) - } else { - item, err = c.manager.ImportCodexSession( - ctx, - input.ProjectID, - input.Body.AISessionID, - websession.SyncMode(input.Body.Mode), - ) + agent := websession.Agent(strings.ToLower(strings.TrimSpace(input.Body.Agent))) + if agent == "" { + agent = websession.AgentCodex + } + switch agent { + case websession.AgentPi: + if strings.TrimSpace(input.Body.SessionID) != "" { + item, err = c.manager.ImportPiSessionBySessionID(ctx, input.ProjectID, input.Body.SessionID) + } else { + item, err = c.manager.ImportPiSession(ctx, input.ProjectID, input.Body.AISessionID) + } + case websession.AgentCodex: + if strings.TrimSpace(input.Body.SessionID) != "" { + item, err = c.manager.ImportCodexSessionBySessionID( + ctx, input.ProjectID, input.Body.SessionID, websession.SyncMode(input.Body.Mode), + ) + } else { + item, err = c.manager.ImportCodexSession( + ctx, input.ProjectID, input.Body.AISessionID, websession.SyncMode(input.Body.Mode), + ) + } + default: + return nil, huma.Error400BadRequest("unsupported import agent") } if err != nil { switch { case errors.Is(err, model.ErrProjectNotFound): return nil, huma.Error404NotFound("project not found") case errors.Is(err, gorm.ErrRecordNotFound): - return nil, huma.Error404NotFound("codex session not found") + return nil, huma.Error404NotFound("agent session not found") default: return nil, huma.Error400BadRequest(err.Error()) } @@ -414,7 +607,7 @@ func (c *webSessionController) registerHTTP(app *fiber.App, group *huma.Group) { return resp, nil }, func(op *huma.Operation) { op.OperationID = "web-session-import" - op.Summary = "导入 Codex 历史会话" + op.Summary = "导入 Agent 历史会话" op.Tags = []string{webSessionTag} }) @@ -857,6 +1050,34 @@ func looksLikeWindowsAbsolutePath(value string) bool { return (first >= 'A' && first <= 'Z') || (first >= 'a' && first <= 'z') } +func (c *webSessionController) requireProjectSession(ctx context.Context, projectID, sessionID string) error { + record, err := c.manager.GetSession(ctx, sessionID) + if err != nil || record.ProjectID != projectID { + return huma.Error404NotFound("session not found") + } + if !c.manager.SupportsPiSessionTree() { + return huma.Error403Forbidden("Pi session tree is not supported") + } + return nil +} + +func piTreeHTTPError(err error) error { + if errors.Is(err, model.ErrDBNotInitialized) { + return huma.Error503ServiceUnavailable("database is not available") + } + publicErr := websession.ClassifyPiTreeError(err) + switch publicErr.Code { + case "conflict", "invalid_state": + return huma.Error409Conflict(publicErr.Message) + case "bad_req": + return huma.Error400BadRequest(publicErr.Message) + case "forbidden": + return huma.Error403Forbidden(publicErr.Message) + default: + return huma.Error500InternalServerError(publicErr.Message) + } +} + func (c *webSessionController) registerWebsocket(app *fiber.App) { commandHandler := fasthttpadaptor.NewFastHTTPHandler(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { c.serveCommandWebsocket(w, r) diff --git a/api/web_session_tree_test.go b/api/web_session_tree_test.go new file mode 100644 index 00000000..b7613dd0 --- /dev/null +++ b/api/web_session_tree_test.go @@ -0,0 +1,127 @@ +package api + +import ( + "errors" + "io" + "net/http" + "net/http/httptest" + "path/filepath" + "strings" + "testing" + "time" + + "code-kanban/api/h" + "code-kanban/model" + "code-kanban/model/tables" + "code-kanban/service" + "code-kanban/service/websession" + "code-kanban/utils" + + "github.com/danielgtaylor/huma/v2" + "github.com/gofiber/fiber/v2" + "go.uber.org/zap" +) + +func TestWebSessionTreeRoutesFailClosedAndHideCrossProjectSessions(t *testing.T) { + model.DBClose() + if err := model.InitWithDSN(filepath.Join(t.TempDir(), "web-session-tree.db"), 0, true); err != nil { + t.Fatalf("InitWithDSN: %v", err) + } + t.Cleanup(model.DBClose) + + session := tables.WebSessionTable{ + ProjectID: "project-tree", Agent: "pi", Backend: "pi_rpc", Title: "Tree source", + WorkflowMode: "default", PermissionLevel: "yolo", Cwd: t.TempDir(), Status: "idle", + ActivityAt: time.Now(), NativeSessionID: stringPointer("native-tree"), ThreadPath: stringPointer(filepath.Join(t.TempDir(), "native-tree.jsonl")), + } + session.Init() + if err := model.GetDB().Create(&session).Error; err != nil { + t.Fatalf("create session: %v", err) + } + manager, err := websession.NewManager(websession.Config{DataDir: t.TempDir(), PiPath: filepath.Join(t.TempDir(), "missing-pi")}, zap.NewNop()) + if err != nil { + t.Fatalf("NewManager: %v", err) + } + app := fiber.New(fiber.Config{Immutable: true}) + _, group := h.NewAPI(app, &utils.AppConfig{}) + registerWebSessionRoutes(app, group, manager, zap.NewNop()) + + base := "/api/v1/projects/project-tree/web-sessions/" + session.ID + "/tree" + for _, request := range []struct { + method string + path string + body string + }{ + {http.MethodGet, base, ""}, + {http.MethodPost, base + "/navigate", `{"targetId":"item","revision":"rev"}`}, + {http.MethodPost, base + "/fork", `{"targetId":"item","revision":"rev"}`}, + {http.MethodPost, base + "/clone", `{"revision":"rev"}`}, + } { + response, payload := requestWebSessionTree(t, app, request.method, request.path, request.body) + if response.StatusCode != http.StatusForbidden { + response.Body.Close() + t.Fatalf("%s %s status = %d, want 403: %s", request.method, request.path, response.StatusCode, payload) + } + response.Body.Close() + } + + crossProject := strings.Replace(base, "/projects/project-tree/", "/projects/project-other/", 1) + response, payload := requestWebSessionTree(t, app, http.MethodGet, crossProject, "") + defer response.Body.Close() + if response.StatusCode != http.StatusNotFound { + t.Fatalf("cross-project status = %d, want 404: %s", response.StatusCode, payload) + } +} + +func TestPiTreeHTTPErrorMapping(t *testing.T) { + tests := []struct { + name string + err error + want int + }{ + {"revision", websession.ErrPiTreeRevisionConflict, http.StatusConflict}, + {"active", errors.New("cannot navigate an active Pi web session"), http.StatusConflict}, + {"pending", errors.New("cannot navigate while messages are pending"), http.StatusConflict}, + {"input", errors.New("Pi tree revision is required"), http.StatusBadRequest}, + {"target", errors.New("Pi fork target is not a user message"), http.StatusBadRequest}, + {"trust", service.ErrProjectAgentTrustRequired, http.StatusForbidden}, + {"db", model.ErrDBNotInitialized, http.StatusServiceUnavailable}, + {"integrity", errors.New("Pi session tree contains a duplicate node id"), http.StatusInternalServerError}, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + var statusErr huma.StatusError + if !errors.As(piTreeHTTPError(test.err), &statusErr) { + t.Fatalf("error does not implement huma.StatusError") + } + if statusErr.GetStatus() != test.want { + t.Fatalf("status = %d, want %d", statusErr.GetStatus(), test.want) + } + if strings.Contains(piTreeHTTPError(test.err).Error(), "duplicate node id") { + t.Fatal("internal Pi tree detail leaked through HTTP error") + } + }) + } +} + +func requestWebSessionTree(t *testing.T, app *fiber.App, method, target, body string) (*http.Response, string) { + t.Helper() + request := httptest.NewRequest(method, target, strings.NewReader(body)) + if body != "" { + request.Header.Set(fiber.HeaderContentType, fiber.MIMEApplicationJSON) + } + response, err := app.Test(request) + if err != nil { + t.Fatalf("app.Test: %v", err) + } + payload, err := io.ReadAll(response.Body) + if err != nil { + response.Body.Close() + t.Fatalf("read response: %v", err) + } + response.Body.Close() + response.Body = io.NopCloser(strings.NewReader(string(payload))) + return response, string(payload) +} + +func stringPointer(value string) *string { return &value } diff --git a/model/db_migrate.go b/model/db_migrate.go index e01b4d4b..5e7de965 100644 --- a/model/db_migrate.go +++ b/model/db_migrate.go @@ -13,6 +13,7 @@ func GetAllModels() []any { &tables.UserTable{}, &tables.UserAccessTokenTable{}, &tables.ProjectTable{}, + &tables.ProjectAgentTrustTable{}, &tables.WorktreeTable{}, &tables.TaskTable{}, &tables.TaskCommentTable{}, diff --git a/model/project.go b/model/project.go index c04e4267..586306fe 100644 --- a/model/project.go +++ b/model/project.go @@ -10,10 +10,12 @@ import ( "strings" "time" + "code-kanban/model/tables" "code-kanban/utils" "code-kanban/utils/git" "go.uber.org/zap" + "gorm.io/gorm" ) var ( @@ -212,25 +214,30 @@ func (s *ProjectService) DeleteProject(ctx context.Context, id string) error { if ctx == nil { ctx = context.Background() } - - q, err := resolveQueries(nil) - if err != nil { - return err + if db == nil { + return ErrDBNotInitialized } - now := time.Now() - affected, err := q.ProjectSoftDelete(ctx, &ProjectSoftDeleteParams{ - DeletedAt: &now, - UpdatedAt: now, - Id: id, + return db.WithContext(ctx).Transaction(func(tx *gorm.DB) error { + var project tables.ProjectTable + if err := tx.Where("id = ?", id).First(&project).Error; err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return ErrProjectNotFound + } + return err + } + + now := time.Now() + if err := tx.Model(&project).Updates(map[string]any{ + "deleted_at": now, + "updated_at": now, + }).Error; err != nil { + return err + } + return tx.Unscoped(). + Where("project_id = ?", id). + Delete(&tables.ProjectAgentTrustTable{}).Error }) - if err != nil { - return err - } - if affected == 0 { - return ErrProjectNotFound - } - return nil } // UpdateProject modifies project metadata such as name and description. diff --git a/model/sqlc_gen/schema_postgres.sql b/model/sqlc_gen/schema_postgres.sql index 1ca7a134..7c321131 100644 --- a/model/sqlc_gen/schema_postgres.sql +++ b/model/sqlc_gen/schema_postgres.sql @@ -2,4 +2,20 @@ -- 生成时间: 2025-10-07 04:00:01 -- 数据库方言: postgres --- 预留文件:当前模板尚未提供 postgres 专用建表语句,按需补充。 +-- 当前模板尚未提供完整的 postgres 专用建表语句。 + +CREATE TABLE "project_agent_trusts" ( + "id" text NOT NULL, + "created_at" timestamp with time zone, + "updated_at" timestamp with time zone, + "deleted_at" timestamp with time zone, + "project_id" text NOT NULL, + "agent" text NOT NULL, + "trusted_path" text NOT NULL, + "trusted_at" timestamp with time zone NOT NULL, + "revoked_at" timestamp with time zone, + PRIMARY KEY ("id") +); +CREATE INDEX "idx_project_agent_trusts_revoked_at" ON "project_agent_trusts" ("revoked_at"); +CREATE UNIQUE INDEX "idx_project_agent_trust" ON "project_agent_trusts" ("project_id", "agent"); +CREATE INDEX "idx_project_agent_trusts_deleted_at" ON "project_agent_trusts" ("deleted_at"); diff --git a/model/sqlc_gen/schema_sqlite.sql b/model/sqlc_gen/schema_sqlite.sql index a75eb407..de252d1e 100644 --- a/model/sqlc_gen/schema_sqlite.sql +++ b/model/sqlc_gen/schema_sqlite.sql @@ -1,7 +1,7 @@ -- 数据库建表语句 -- 生成时间: 2026-07-26 22:41:41 -- 数据库方言: sqlite --- 总共 103 条语句 +-- 总共 108 条语句 CREATE TABLE "users" ("id" text NOT NULL,"created_at" datetime,"updated_at" datetime,"deleted_at" datetime,"nickname" text,"avatar" text,"brief" text,"username" text NOT NULL,"password" text NOT NULL,"salt" text NOT NULL,"disabled" numeric NOT NULL DEFAULT false,PRIMARY KEY ("id")); @@ -21,6 +21,12 @@ CREATE INDEX "idx_projects_name" ON "projects"("name"); CREATE INDEX "idx_projects_deleted_at" ON "projects"("deleted_at"); +CREATE TABLE "project_agent_trusts" ("id" text NOT NULL,"created_at" datetime,"updated_at" datetime,"deleted_at" datetime,"project_id" text NOT NULL,"agent" text NOT NULL,"trusted_path" text NOT NULL,"trusted_at" datetime NOT NULL,"revoked_at" datetime,PRIMARY KEY ("id")); +CREATE INDEX "idx_project_agent_trusts_revoked_at" ON "project_agent_trusts"("revoked_at"); +CREATE UNIQUE INDEX "idx_project_agent_trust" ON "project_agent_trusts"("project_id","agent"); +CREATE INDEX "idx_project_agent_trusts_deleted_at" ON "project_agent_trusts"("deleted_at"); + + CREATE TABLE "worktrees" ("id" text NOT NULL,"created_at" datetime,"updated_at" datetime,"deleted_at" datetime,"project_id" text NOT NULL,"branch_name" text NOT NULL,"path" text NOT NULL,"is_main" boolean DEFAULT false,"is_bare" boolean DEFAULT false,"head_commit" text,"head_commit_message" text,"head_commit_date" datetime,"status_ahead" integer DEFAULT 0,"status_behind" integer DEFAULT 0,"status_modified" integer DEFAULT 0,"status_staged" integer DEFAULT 0,"status_untracked" integer DEFAULT 0,"status_conflicts" integer DEFAULT 0,"status_updated_at" datetime,PRIMARY KEY ("id")); CREATE UNIQUE INDEX "idx_worktrees_path" ON "worktrees"("path") WHERE deleted_at IS NULL; CREATE INDEX "idx_worktrees_branch_name" ON "worktrees"("branch_name"); @@ -54,7 +60,7 @@ CREATE UNIQUE INDEX "idx_session_type" ON "ai_sessions"("session_id","type"); CREATE INDEX "idx_ai_sessions_deleted_at" ON "ai_sessions"("deleted_at"); -CREATE TABLE "web_sessions" ("id" text NOT NULL,"created_at" datetime,"updated_at" datetime,"deleted_at" datetime,"project_id" text NOT NULL,"worktree_id" text,"order_index" real NOT NULL DEFAULT 0,"agent" text NOT NULL,"claude_runtime" text NOT NULL DEFAULT "claude","backend" text NOT NULL DEFAULT "legacy_exec","title" text NOT NULL,"title_auto" boolean NOT NULL DEFAULT false,"model" text,"reasoning_effort" text,"workflow_mode" text NOT NULL DEFAULT "default","permission_level" text NOT NULL DEFAULT "elevated","active_call_timeout_enabled" boolean,"auto_retry_enabled" boolean NOT NULL DEFAULT false,"auto_retry_scope" text NOT NULL DEFAULT "network_only","auto_retry_preset" text NOT NULL DEFAULT "gentle_stop","auto_retry_max_attempts" integer NOT NULL DEFAULT 0,"auto_retry_dispatch_pending_on_failure" boolean NOT NULL DEFAULT false,"permission_mode" text,"cwd" text NOT NULL,"native_session_id" text,"cyber_policy_flagged" boolean NOT NULL DEFAULT false,"status" text NOT NULL,"assistant_state" text,"has_unread" boolean NOT NULL DEFAULT false,"archived_at" datetime,"activity_at" datetime,"status_updated_at" datetime,"assistant_state_updated_at" datetime,"source_kind" text NOT NULL DEFAULT "codex_app_server","sync_state" text NOT NULL DEFAULT "missing","last_sync_mode" text,"source_created_at" datetime,"source_updated_at" datetime,"last_synced_at" datetime,"thread_path" text,"thread_preview" text,"turn_count" integer NOT NULL DEFAULT 0,"item_count" integer NOT NULL DEFAULT 0,"last_message_at" datetime,"last_event_seq" integer NOT NULL DEFAULT 0,"snapshot_revision" integer NOT NULL DEFAULT 1,"goal_objective" text,"goal_status" text,"goal_token_budget" integer,"goal_tokens_used" integer NOT NULL DEFAULT 0,"goal_time_used_seconds" integer NOT NULL DEFAULT 0,"goal_created_at" datetime,"goal_updated_at" datetime,"total_input_tokens" integer NOT NULL DEFAULT 0,"total_cached_input_tokens" integer NOT NULL DEFAULT 0,"total_output_tokens" integer NOT NULL DEFAULT 0,"total_cost" real NOT NULL DEFAULT 0,"last_completed_input_tokens" integer NOT NULL DEFAULT 0,"last_completed_cached_input_tokens" integer NOT NULL DEFAULT 0,"last_completed_output_tokens" integer NOT NULL DEFAULT 0,"latest_turn_input_tokens" integer NOT NULL DEFAULT 0,"latest_turn_cached_input_tokens" integer NOT NULL DEFAULT 0,"latest_turn_output_tokens" integer NOT NULL DEFAULT 0,"latest_turn_usage_updated_at" datetime,"latest_token_count_input_tokens" integer NOT NULL DEFAULT 0,"latest_token_count_cached_input_tokens" integer NOT NULL DEFAULT 0,"latest_token_count_output_tokens" integer NOT NULL DEFAULT 0,"latest_token_count_total_tokens" integer NOT NULL DEFAULT 0,"latest_token_count_updated_at" datetime,"session_context_window_tokens" integer NOT NULL DEFAULT 0,"session_context_window_observed_at" datetime,"context_baseline_input_tokens" integer NOT NULL DEFAULT 0,"context_baseline_cached_input_tokens" integer NOT NULL DEFAULT 0,"context_baseline_output_tokens" integer NOT NULL DEFAULT 0,"last_context_compaction_at" datetime,"auto_retry_attempt" integer NOT NULL DEFAULT 0,"auto_retry_next_at" datetime,"auto_retry_last_error_code" text,"last_error" text,"sync_error" text,PRIMARY KEY ("id")); +CREATE TABLE "web_sessions" ("id" text NOT NULL,"created_at" datetime,"updated_at" datetime,"deleted_at" datetime,"project_id" text NOT NULL,"worktree_id" text,"order_index" real NOT NULL DEFAULT 0,"agent" text NOT NULL,"claude_runtime" text NOT NULL DEFAULT "claude","backend" text NOT NULL DEFAULT "legacy_exec","title" text NOT NULL,"title_auto" boolean NOT NULL DEFAULT false,"model" text,"reasoning_effort" text,"workflow_mode" text NOT NULL DEFAULT "default","permission_level" text NOT NULL DEFAULT "elevated","active_call_timeout_enabled" boolean,"auto_retry_enabled" boolean NOT NULL DEFAULT false,"auto_retry_scope" text NOT NULL DEFAULT "network_only","auto_retry_preset" text NOT NULL DEFAULT "gentle_stop","auto_retry_max_attempts" integer NOT NULL DEFAULT 0,"auto_retry_dispatch_pending_on_failure" boolean NOT NULL DEFAULT false,"permission_mode" text,"cwd" text NOT NULL,"native_session_id" text,"native_leaf_id" text,"source_revision" text,"cyber_policy_flagged" boolean NOT NULL DEFAULT false,"status" text NOT NULL,"assistant_state" text,"has_unread" boolean NOT NULL DEFAULT false,"archived_at" datetime,"activity_at" datetime,"status_updated_at" datetime,"assistant_state_updated_at" datetime,"source_kind" text NOT NULL DEFAULT "codex_app_server","sync_state" text NOT NULL DEFAULT "missing","last_sync_mode" text,"source_created_at" datetime,"source_updated_at" datetime,"last_synced_at" datetime,"thread_path" text,"thread_preview" text,"turn_count" integer NOT NULL DEFAULT 0,"item_count" integer NOT NULL DEFAULT 0,"last_message_at" datetime,"last_event_seq" integer NOT NULL DEFAULT 0,"snapshot_revision" integer NOT NULL DEFAULT 1,"goal_objective" text,"goal_status" text,"goal_token_budget" integer,"goal_tokens_used" integer NOT NULL DEFAULT 0,"goal_time_used_seconds" integer NOT NULL DEFAULT 0,"goal_created_at" datetime,"goal_updated_at" datetime,"total_input_tokens" integer NOT NULL DEFAULT 0,"total_cached_input_tokens" integer NOT NULL DEFAULT 0,"total_output_tokens" integer NOT NULL DEFAULT 0,"total_cost" real NOT NULL DEFAULT 0,"last_completed_input_tokens" integer NOT NULL DEFAULT 0,"last_completed_cached_input_tokens" integer NOT NULL DEFAULT 0,"last_completed_output_tokens" integer NOT NULL DEFAULT 0,"latest_turn_input_tokens" integer NOT NULL DEFAULT 0,"latest_turn_cached_input_tokens" integer NOT NULL DEFAULT 0,"latest_turn_output_tokens" integer NOT NULL DEFAULT 0,"latest_turn_usage_updated_at" datetime,"latest_token_count_input_tokens" integer NOT NULL DEFAULT 0,"latest_token_count_cached_input_tokens" integer NOT NULL DEFAULT 0,"latest_token_count_output_tokens" integer NOT NULL DEFAULT 0,"latest_token_count_total_tokens" integer NOT NULL DEFAULT 0,"latest_token_count_updated_at" datetime,"session_context_window_tokens" integer NOT NULL DEFAULT 0,"session_context_window_observed_at" datetime,"context_baseline_input_tokens" integer NOT NULL DEFAULT 0,"context_baseline_cached_input_tokens" integer NOT NULL DEFAULT 0,"context_baseline_output_tokens" integer NOT NULL DEFAULT 0,"last_context_compaction_at" datetime,"auto_retry_attempt" integer NOT NULL DEFAULT 0,"auto_retry_next_at" datetime,"auto_retry_last_error_code" text,"last_error" text,"sync_error" text,PRIMARY KEY ("id")); CREATE INDEX "idx_web_sessions_source_updated_at" ON "web_sessions"("source_updated_at"); CREATE INDEX "idx_web_sessions_sync_state" ON "web_sessions"("sync_state"); CREATE INDEX "idx_web_sessions_status_updated_at" ON "web_sessions"("status_updated_at"); diff --git a/model/tables/ai_session.go b/model/tables/ai_session.go index 8e1d00f3..6d339e0f 100644 --- a/model/tables/ai_session.go +++ b/model/tables/ai_session.go @@ -12,6 +12,7 @@ type AISessionType string const ( AISessionTypeClaudeCode AISessionType = "claude_code" AISessionTypeCodex AISessionType = "codex" + AISessionTypePi AISessionType = "pi" ) // AISessionTable stores cached AI assistant session metadata. @@ -23,7 +24,7 @@ type AISessionTable struct { // SessionID is the unique identifier from the AI assistant SessionID string `gorm:"type:text;not null;uniqueIndex:idx_session_type" json:"sessionId"` - // Type identifies which AI assistant (claude_code, codex) + // Type identifies which AI assistant (claude_code, codex, pi) Type AISessionType `gorm:"type:text;not null;uniqueIndex:idx_session_type" json:"type"` // ProjectPath is the working directory associated with this session diff --git a/model/tables/project_agent_trust.go b/model/tables/project_agent_trust.go new file mode 100644 index 00000000..7d45c2ff --- /dev/null +++ b/model/tables/project_agent_trust.go @@ -0,0 +1,23 @@ +package tables + +import ( + "time" + + "code-kanban/utils/model_base" +) + +// ProjectAgentTrustTable records explicit approval for an agent to load +// project-local resources from a specific project path. +type ProjectAgentTrustTable struct { + model_base.StringPKBaseModel + + ProjectID string `gorm:"type:text;not null;uniqueIndex:idx_project_agent_trust,priority:1" json:"projectId"` + Agent string `gorm:"type:text;not null;uniqueIndex:idx_project_agent_trust,priority:2" json:"agent"` + TrustedPath string `gorm:"type:text;not null" json:"trustedPath"` + TrustedAt time.Time `gorm:"type:datetime;not null" json:"trustedAt"` + RevokedAt *time.Time `gorm:"type:datetime;index" json:"revokedAt"` +} + +func (ProjectAgentTrustTable) TableName() string { + return "project_agent_trusts" +} diff --git a/model/tables/web_session.go b/model/tables/web_session.go index b17a3f5b..608a9264 100644 --- a/model/tables/web_session.go +++ b/model/tables/web_session.go @@ -34,6 +34,8 @@ type WebSessionTable struct { Cwd string `gorm:"type:text;not null" json:"cwd"` NativeSessionID *string `gorm:"type:text" json:"nativeSessionId"` + NativeLeafID *string `gorm:"type:text" json:"nativeLeafId"` + SourceRevision *string `gorm:"type:text" json:"sourceRevision"` CyberPolicyFlagged bool `gorm:"type:boolean;not null;default:false" json:"cyberPolicyFlagged"` Status string `gorm:"type:text;not null;index" json:"status"` AssistantState string `gorm:"type:text;index" json:"assistantState"` diff --git a/packages/node-sdk/src/client.js b/packages/node-sdk/src/client.js index d7b50be3..38799ac9 100644 --- a/packages/node-sdk/src/client.js +++ b/packages/node-sdk/src/client.js @@ -1,6 +1,6 @@ import { readFile } from "node:fs/promises"; -import { buildAgentLaunchSpec } from "./command-builder.js"; +import { AGENTS, buildAgentLaunchSpec } from "./command-builder.js"; import { CodeKanbanConfigError, CodeKanbanHttpError, @@ -14,7 +14,11 @@ import { WEB_SESSION_EVENTS_WS_PATH, analyzeWebSession, ensureImageMimeType, + normalizePiTreeCreateResult, + normalizePiTreeNavigateResult, + normalizePiTreeSnapshot, normalizeWebSessionAttachment, + normalizeWebSessionRuntimeConfig, } from "./web-session-shared.js"; import { ensureArrayOfStrings, @@ -603,6 +607,50 @@ export class CodeKanbanClient { return extractPayloadItem(response); } + async getProjectPiTrust(input = {}) { + const projectId = await this.resolveProjectId({ + projectId: input.projectId, + projectName: input.projectName, + projectIndex: input.projectIndex, + path: input.path, + ensureProject: input.ensureProject !== false, + }); + const response = await this.requestJson( + `/api/v1/projects/${projectId}/agent-trust/pi`, + ); + return extractPayloadItem(response); + } + + async trustProjectForPi(input = {}) { + const projectId = await this.resolveProjectId({ + projectId: input.projectId, + projectName: input.projectName, + projectIndex: input.projectIndex, + path: input.path, + ensureProject: input.ensureProject !== false, + }); + const response = await this.requestJson( + `/api/v1/projects/${projectId}/agent-trust/pi`, + { method: "POST" }, + ); + return extractPayloadItem(response); + } + + async revokeProjectPiTrust(input = {}) { + const projectId = await this.resolveProjectId({ + projectId: input.projectId, + projectName: input.projectName, + projectIndex: input.projectIndex, + path: input.path, + ensureProject: input.ensureProject !== false, + }); + const response = await this.requestJson( + `/api/v1/projects/${projectId}/agent-trust/pi`, + { method: "DELETE" }, + ); + return extractPayloadItem(response); + } + async listWorktrees(projectId) { ensureString(projectId, "projectId"); const response = await this.requestJson( @@ -856,8 +904,10 @@ export class CodeKanbanClient { : Promise.resolve({ hasClaudeCode: false, hasCodex: false, + hasPi: false, claudeSessions: [], codexSessions: [], + piSessions: [], }), ]); @@ -966,6 +1016,57 @@ export class CodeKanbanClient { return extractPayloadItems(response); } + async listWebSessionImportSources({ projectId, projectName, projectIndex, path, refresh = false } = {}) { + const resolvedProjectId = await this.resolveProjectId({ + projectId, + projectName, + projectIndex, + path, + ensureProject: true, + }); + const query = refresh === true ? "?refresh=true" : ""; + const response = await this.requestJson( + `/api/v1/projects/${resolvedProjectId}/web-sessions/import-sources${query}`, + ); + return extractPayloadItem(response) || { items: [], scanPhase: "complete" }; + } + + async importWebSession(input = {}) { + const projectId = await this.resolveProjectId({ + projectId: input.projectId, + projectName: input.projectName, + projectIndex: input.projectIndex, + path: input.path, + ensureProject: true, + }); + const agent = ensureOptionalString(input.agent) || "codex"; + if (agent !== "codex" && agent !== "pi") { + throw new CodeKanbanValidationError("agent must be codex or pi"); + } + const sessionId = ensureOptionalString(input.sessionId); + const aiSessionId = ensureOptionalString(input.aiSessionId); + if (!sessionId && !aiSessionId) { + throw new CodeKanbanValidationError("sessionId or aiSessionId is required"); + } + const mode = ensureOptionalString(input.mode); + if (mode && mode !== "fast" && mode !== "deep") { + throw new CodeKanbanValidationError("mode must be fast or deep"); + } + const response = await this.requestJson( + `/api/v1/projects/${projectId}/web-sessions/import`, + { + method: "POST", + body: { + agent, + ...(sessionId ? { sessionId } : {}), + ...(aiSessionId ? { aiSessionId } : {}), + ...(mode ? { mode } : {}), + }, + }, + ); + return extractPayloadItem(response); + } + async createWebSession(input = {}) { const projectId = await this.resolveProjectId({ projectId: input.projectId, @@ -976,6 +1077,9 @@ export class CodeKanbanClient { }); const agent = ensureString(input.agent, "agent"); + if (!AGENTS.includes(agent)) { + throw new CodeKanbanValidationError(`agent must be one of: ${AGENTS.join(", ")}`); + } const permissionMode = ensureOptionalString( input.permissionMode, ).toLowerCase(); @@ -1036,6 +1140,111 @@ export class CodeKanbanClient { return extractPayloadItem(response); } + async getWebSessionTree({ projectId, projectName, projectIndex, path, sessionId }) { + const resolvedProjectId = await this.resolveProjectId({ + projectId, + projectName, + projectIndex, + path, + ensureProject: true, + }); + const resolvedSessionId = ensureString(sessionId, "sessionId"); + const response = await this.requestJson( + `/api/v1/projects/${resolvedProjectId}/web-sessions/${resolvedSessionId}/tree`, + ); + return normalizePiTreeSnapshot(extractPayloadItem(response)); + } + + async navigateWebSessionTree({ + projectId, + projectName, + projectIndex, + path, + sessionId, + targetId, + revision, + summarize = false, + }) { + const resolvedProjectId = await this.resolveProjectId({ + projectId, + projectName, + projectIndex, + path, + ensureProject: true, + }); + const resolvedSessionId = ensureString(sessionId, "sessionId"); + const response = await this.requestJson( + `/api/v1/projects/${resolvedProjectId}/web-sessions/${resolvedSessionId}/tree/navigate`, + { + method: "POST", + body: { + targetId: ensureString(targetId, "targetId"), + revision: ensureString(revision, "revision"), + summarize: summarize === true, + }, + }, + ); + return normalizePiTreeNavigateResult(extractPayloadItem(response)); + } + + async forkWebSessionTree({ + projectId, + projectName, + projectIndex, + path, + sessionId, + targetId, + revision, + }) { + const resolvedProjectId = await this.resolveProjectId({ + projectId, + projectName, + projectIndex, + path, + ensureProject: true, + }); + const resolvedSessionId = ensureString(sessionId, "sessionId"); + const response = await this.requestJson( + `/api/v1/projects/${resolvedProjectId}/web-sessions/${resolvedSessionId}/tree/fork`, + { + method: "POST", + body: { + targetId: ensureString(targetId, "targetId"), + revision: ensureString(revision, "revision"), + }, + }, + ); + return normalizePiTreeCreateResult(extractPayloadItem(response)); + } + + async cloneWebSessionTree({ + projectId, + projectName, + projectIndex, + path, + sessionId, + revision, + }) { + const resolvedProjectId = await this.resolveProjectId({ + projectId, + projectName, + projectIndex, + path, + ensureProject: true, + }); + const resolvedSessionId = ensureString(sessionId, "sessionId"); + const response = await this.requestJson( + `/api/v1/projects/${resolvedProjectId}/web-sessions/${resolvedSessionId}/tree/clone`, + { + method: "POST", + body: { + revision: ensureString(revision, "revision"), + }, + }, + ); + return normalizePiTreeCreateResult(extractPayloadItem(response)); + } + async getWebSessionHistory({ projectId, projectName, @@ -1265,7 +1474,7 @@ export class CodeKanbanClient { const response = await this.requestJson( "/api/v1/web-sessions/runtime-config", ); - return extractPayloadItem(response); + return normalizeWebSessionRuntimeConfig(extractPayloadItem(response)); } async uploadWebSessionAttachment({ @@ -1351,6 +1560,12 @@ export class CodeKanbanClient { ); } + async compactWebSession({ sessionId }) { + return await this.withWebSessionCommandChannel((channel) => + channel.compact(sessionId), + ); + } + async removeWebSessionPendingInput({ sessionId, pendingId }) { return await this.withWebSessionCommandChannel((channel) => channel.removePendingInput(sessionId, { diff --git a/packages/node-sdk/src/command-builder.js b/packages/node-sdk/src/command-builder.js index d0f91333..2fe558af 100644 --- a/packages/node-sdk/src/command-builder.js +++ b/packages/node-sdk/src/command-builder.js @@ -4,7 +4,7 @@ import { ensureArrayOfStrings, ensureOptionalString, toCommandString } from './u export const SANDBOX_MODES = ['read-only', 'workspace-write', 'danger-full-access']; export const APPROVAL_POLICIES = ['untrusted', 'on-request', 'never']; export const WORKFLOW_PROFILES = ['plan', 'standard', 'yolo']; -export const AGENTS = ['codex', 'claude']; +export const AGENTS = ['codex', 'claude', 'pi']; export const CLAUDE_RUNTIMES = ['claude', 'ccr']; const KNOWN_STRUCTURED_FLAGS = new Set([ @@ -80,6 +80,20 @@ export function buildAgentLaunchSpec(options = {}) { }; } + if (agent === 'pi') { + if (options.permissions) { + throw new CodeKanbanValidationError('structured permissions are not supported for pi'); + } + const argv = ['pi', ...extraArgs]; + return { + agent, + profile, + argv, + command: toCommandString(argv), + prompt: composeWorkflowPrompt({ profile, prompt: options.prompt }), + }; + } + const permissions = options.permissions || {}; const conflicts = detectStructuredFlagConflicts(extraArgs); if ( diff --git a/packages/node-sdk/src/index.js b/packages/node-sdk/src/index.js index 353db5dc..7d774882 100644 --- a/packages/node-sdk/src/index.js +++ b/packages/node-sdk/src/index.js @@ -17,4 +17,7 @@ export { export { TerminalConnection } from './terminal-connection.js'; export { WebSessionCommandChannel } from './web-session-command-channel.js'; export { WebSessionEventStream } from './web-session-event-stream.js'; -export { analyzeWebSession } from './web-session-shared.js'; +export { + analyzeWebSession, + normalizeWebSessionRuntimeConfig, +} from './web-session-shared.js'; diff --git a/packages/node-sdk/src/web-session-command-channel.js b/packages/node-sdk/src/web-session-command-channel.js index e6817aca..1f0b5022 100644 --- a/packages/node-sdk/src/web-session-command-channel.js +++ b/packages/node-sdk/src/web-session-command-channel.js @@ -9,6 +9,9 @@ import { buildWebSessionHeartbeatFrame, decodeWebSessionSocketMessage, isWebSessionHeartbeatFrame, + normalizePiTreeCreateResult, + normalizePiTreeNavigateResult, + normalizePiTreeSnapshot, normalizeWebSessionFrame, } from "./web-session-shared.js"; import { @@ -207,6 +210,51 @@ export class WebSessionCommandChannel { return await this._executeCommand({ operation: "abort", sessionId }); } + async compact(sessionId) { + return await this._executeCommand({ operation: "compact", sessionId }); + } + + async getTree(sessionId) { + const ack = await this._executeCommand({ operation: "tree_get", sessionId }); + return normalizePiTreeSnapshot(ack.payload); + } + + async navigateTree(sessionId, input = {}) { + const ack = await this._executeCommand({ + operation: "tree_nav", + sessionId, + payload: { + tid: ensureString(input.targetId, "targetId"), + rev: ensureString(input.revision, "revision"), + sum: input.summarize === true, + }, + }); + return normalizePiTreeNavigateResult(ack.payload); + } + + async forkTree(sessionId, input = {}) { + const ack = await this._executeCommand({ + operation: "tree_fork", + sessionId, + payload: { + tid: ensureString(input.targetId, "targetId"), + rev: ensureString(input.revision, "revision"), + }, + }); + return normalizePiTreeCreateResult(ack.payload); + } + + async cloneTree(sessionId, input = {}) { + const ack = await this._executeCommand({ + operation: "tree_clone", + sessionId, + payload: { + rev: ensureString(input.revision, "revision"), + }, + }); + return normalizePiTreeCreateResult(ack.payload); + } + async approve(sessionId) { return await this._executeCommand({ operation: "approve", sessionId }); } diff --git a/packages/node-sdk/src/web-session-shared.js b/packages/node-sdk/src/web-session-shared.js index 49c13d18..5defb459 100644 --- a/packages/node-sdk/src/web-session-shared.js +++ b/packages/node-sdk/src/web-session-shared.js @@ -64,6 +64,119 @@ function booleanValue(value) { return value === true; } +function unavailableAgentCapability() { + return { + installed: false, + version: null, + supportsWebSession: false, + supportsTree: false, + supportsImages: false, + supportsCompaction: false, + supportsSteer: false, + supportsFollowUp: false, + supportsGoal: false, + supportsSubAgentRegistry: false, + permissionModes: [], + }; +} + +function normalizeAgentCapability(value, fallback = unavailableAgentCapability()) { + if (!value || typeof value !== "object") { + return { ...fallback }; + } + return { + ...fallback, + installed: booleanValue(value.installed), + version: trimmedString(value.version) || null, + supportsWebSession: booleanValue(value.supportsWebSession), + supportsTree: booleanValue(value.supportsTree), + supportsImages: booleanValue(value.supportsImages), + supportsCompaction: booleanValue(value.supportsCompaction), + supportsSteer: booleanValue(value.supportsSteer), + supportsFollowUp: booleanValue(value.supportsFollowUp), + supportsGoal: booleanValue(value.supportsGoal), + supportsSubAgentRegistry: booleanValue(value.supportsSubAgentRegistry), + permissionModes: Array.isArray(value.permissionModes) + ? value.permissionModes + .map((mode) => ({ + id: trimmedString(mode?.id), + available: booleanValue(mode?.available), + })) + .filter((mode) => mode.id) + : [], + }; +} + +export function normalizeWebSessionRuntimeConfig(value) { + const config = value && typeof value === "object" ? value : {}; + const hasCodex = booleanValue(config.hasCodex); + const hasClaudeCode = booleanValue(config.hasClaudeCode); + const hasPi = booleanValue(config.hasPi); + const supportsCodexWebSession = hasCodex && config.supportsWebSession !== false; + const legacyCapabilities = { + claude: { + ...unavailableAgentCapability(), + installed: hasClaudeCode, + supportsWebSession: hasClaudeCode, + supportsImages: true, + supportsCompaction: true, + supportsSteer: true, + supportsFollowUp: true, + }, + codex: { + ...unavailableAgentCapability(), + installed: hasCodex, + version: trimmedString(config.codexVersion) || null, + supportsWebSession: supportsCodexWebSession, + supportsImages: true, + supportsCompaction: true, + supportsSteer: true, + supportsFollowUp: true, + supportsGoal: booleanValue(config.supportsGoalMode), + supportsSubAgentRegistry: booleanValue(config.supportsMultiAgentV2), + }, + pi: { + ...unavailableAgentCapability(), + installed: hasPi, + version: trimmedString(config.piVersion) || null, + supportsWebSession: hasPi && booleanValue(config.supportsPiWebSession), + supportsTree: booleanValue(config.supportsPiWebSession), + supportsImages: booleanValue(config.supportsPiWebSession), + supportsCompaction: booleanValue(config.supportsPiWebSession), + supportsSteer: booleanValue(config.supportsPiWebSession), + supportsFollowUp: booleanValue(config.supportsPiWebSession), + }, + }; + const explicitAgents = + config.agents && typeof config.agents === "object" ? config.agents : {}; + const piModels = Array.isArray(config.piModels) + ? config.piModels + .map((model) => ({ + provider: trimmedString(model?.provider), + id: trimmedString(model?.id), + name: trimmedString(model?.name), + reasoning: booleanValue(model?.reasoning), + input: Array.isArray(model?.input) + ? model.input.map(trimmedString).filter(Boolean) + : [], + contextWindow: numberValue(model?.contextWindow, 0), + ...(nullableNumberValue(model?.maxTokens) == null + ? {} + : { maxTokens: numberValue(model.maxTokens, 0) }), + })) + .filter((model) => model.provider && model.id) + : []; + return { + ...config, + piModels, + agents: { + claude: normalizeAgentCapability(explicitAgents.claude, legacyCapabilities.claude), + codex: normalizeAgentCapability(explicitAgents.codex, legacyCapabilities.codex), + pi: normalizeAgentCapability(explicitAgents.pi, legacyCapabilities.pi), + }, + }; +} + function normalizeUsage(value) { return { inputTokens: numberValue(value?.in, 0), @@ -246,6 +359,9 @@ function normalizePendingInput(value) { ? trimmedString(value.readyAt) || null : null), paused: value?.ps === true || value?.paused === true, + ...(value?.nq === true || value?.nativeQueued === true + ? { nativeQueued: true } + : {}), createdAt: isoFromUnixMilli(value?.ca) || (typeof value?.createdAt === "string" @@ -360,6 +476,8 @@ export function normalizeWebSessionSummaryFromWire(value) { permissionLevel: trimmedString(value?.pl) || "elevated", cwd: trimmedString(value?.cwd), nativeSessionId: trimmedString(value?.nsid) || null, + nativeLeafId: trimmedString(value?.nlid) || null, + sourceRevision: trimmedString(value?.srev) || null, status: trimmedString(value?.st) || "idle", assistantState: trimmedString(value?.ast) || null, hasUnread: booleanValue(value?.unr), @@ -411,6 +529,59 @@ export function normalizeWebSessionSnapshotFromWire(frame) { }; } +function normalizePiTreeNode(value) { + const id = trimmedString(value?.id); + const type = trimmedString(value?.type); + if (!id || !type) { + return null; + } + return { + id, + parentId: trimmedString(value?.parentId) || null, + type, + role: trimmedString(value?.role), + preview: trimmedString(value?.preview), + timestamp: trimmedString(value?.timestamp), + label: trimmedString(value?.label), + active: booleanValue(value?.active), + children: Array.isArray(value?.children) + ? value.children.map(trimmedString).filter(Boolean) + : [], + }; +} + +export function normalizePiTreeSnapshot(value) { + const nodes = Array.isArray(value?.nodes) + ? value.nodes.map(normalizePiTreeNode).filter(Boolean) + : []; + const nodeIds = new Set(nodes.map(node => node.id)); + return { + sessionId: trimmedString(value?.sessionId), + leafId: nodeIds.has(trimmedString(value?.leafId)) + ? trimmedString(value?.leafId) + : null, + revision: trimmedString(value?.revision), + nodes, + }; +} + +export function normalizePiTreeNavigateResult(value) { + return { + tree: normalizePiTreeSnapshot(value?.tree), + editorText: stringValue(value?.editorText), + }; +} + +export function normalizePiTreeCreateResult(value) { + return { + session: value?.s + ? normalizeWebSessionSummaryFromWire(value.s) + : value?.session || null, + tree: normalizePiTreeSnapshot(value?.tree), + editorText: stringValue(value?.editorText), + }; +} + const PROCESS_RESTART_REASON = "process_restart"; const DEFAULT_RECOVERY_MESSAGE = "Session runtime was interrupted. Send a new message to continue."; diff --git a/packages/node-sdk/test/client.test.js b/packages/node-sdk/test/client.test.js index 80443e9c..d854abd6 100644 --- a/packages/node-sdk/test/client.test.js +++ b/packages/node-sdk/test/client.test.js @@ -110,6 +110,72 @@ test('resolveProject creates a project when path is not registered', async () => assert.equal(result.matchedBy, 'created'); }); +test('project Pi trust helpers use server-owned project context', async () => { + const requests = []; + const status = { + projectId: 'p1', + agent: 'pi', + projectPath: 'D:/repo/demo', + trusted: true, + }; + const handlers = new Map([ + ['GET /api/v1/projects/p1/agent-trust/pi', ({ body }) => { + requests.push({ method: 'GET', body }); + return createJsonResponse({ item: { ...status, trusted: false } }); + }], + ['POST /api/v1/projects/p1/agent-trust/pi', ({ body }) => { + requests.push({ method: 'POST', body }); + return createJsonResponse({ item: status }); + }], + ['DELETE /api/v1/projects/p1/agent-trust/pi', ({ body }) => { + requests.push({ method: 'DELETE', body }); + return createJsonResponse({ item: { ...status, trusted: false } }); + }], + ]); + const client = new CodeKanbanClient({ + baseURL: 'http://127.0.0.1:3000', + fetchImpl: createFetchMock(handlers), + WebSocketImpl: FakeWebSocket, + }); + + assert.equal((await client.getProjectPiTrust({ projectId: 'p1' })).trusted, false); + assert.equal((await client.trustProjectForPi({ projectId: 'p1' })).trusted, true); + assert.equal((await client.revokeProjectPiTrust({ projectId: 'p1' })).trusted, false); + assert.deepEqual(requests, [ + { method: 'GET', body: undefined }, + { method: 'POST', body: undefined }, + { method: 'DELETE', body: undefined }, + ]); +}); + +test('Pi Web Session import sends agent identity without accepting a native file path', async () => { + const requests = []; + const handlers = new Map([ + ['GET /api/v1/projects/p1/web-sessions/import-sources', ({ url }) => { + requests.push({ method: 'GET', search: url.search }); + return createJsonResponse({ item: { items: [{ agent: 'pi', sessionId: 'pi-native' }], scanPhase: 'complete' } }); + }], + ['POST /api/v1/projects/p1/web-sessions/import', ({ body }) => { + requests.push({ method: 'POST', body }); + return createJsonResponse({ item: { created: true, session: { id: 'ws1', agent: 'pi' } } }); + }], + ]); + const client = new CodeKanbanClient({ + baseURL: 'http://127.0.0.1:3000', + fetchImpl: createFetchMock(handlers), + WebSocketImpl: FakeWebSocket, + }); + + const sources = await client.listWebSessionImportSources({ projectId: 'p1', refresh: true }); + const imported = await client.importWebSession({ projectId: 'p1', agent: 'pi', sessionId: 'pi-native' }); + assert.equal(sources.items[0].agent, 'pi'); + assert.equal(imported.session.agent, 'pi'); + assert.deepEqual(requests, [ + { method: 'GET', search: '?refresh=true' }, + { method: 'POST', body: { agent: 'pi', sessionId: 'pi-native' } }, + ]); +}); + test('startWorkflow creates a terminal and sends command plus prompt', async () => { FakeWebSocket.instances.length = 0; const handlers = new Map([ @@ -172,7 +238,7 @@ test('listSessions returns terminal and ai summaries', async () => { const handlers = new Map([ ['GET /api/v1/projects/p1', () => createJsonResponse({ item: { id: 'p1', path: 'D:/repo/demo', name: 'demo' } })], ['GET /api/v1/projects/p1/terminals', () => createJsonResponse({ items: [{ id: 't1' }] })], - ['GET /api/v1/projects/p1/ai-sessions', () => createJsonResponse({ item: { hasCodex: true, hasClaudeCode: false, codexSessions: [{ id: 'a1' }], claudeSessions: [] } })], + ['GET /api/v1/projects/p1/ai-sessions', () => createJsonResponse({ item: { hasCodex: true, hasClaudeCode: false, hasPi: true, codexSessions: [{ id: 'a1' }], claudeSessions: [], piSessions: [{ id: 'p1' }] } })], ]); const client = new CodeKanbanClient({ @@ -185,9 +251,59 @@ test('listSessions returns terminal and ai summaries', async () => { assert.equal(result.project.id, 'p1'); assert.equal(result.terminalSessions.length, 1); assert.equal(result.aiSessions.codexSessions.length, 1); + assert.equal(result.aiSessions.piSessions.length, 1); }); +test('Pi web session tree REST helpers use revisioned contracts', async () => { + const tree = { + sessionId: 'ws1', + leafId: 'a1', + revision: 'rev-1', + nodes: [ + { id: 'u1', parentId: null, type: 'message', role: 'user', preview: 'Start', active: true, children: ['a1'] }, + { id: 'a1', parentId: 'u1', type: 'message', role: 'assistant', preview: 'Answer', active: true, children: [] }, + ], + }; + const handlers = new Map([ + ['GET /api/v1/projects/p1/web-sessions/ws1/tree', () => createWrappedJsonResponse({ item: tree })], + ['POST /api/v1/projects/p1/web-sessions/ws1/tree/navigate', ({ body }) => { + assert.deepEqual(body, { targetId: 'u1', revision: 'rev-1', summarize: true }); + return createWrappedJsonResponse({ item: { tree, editorText: 'Start' } }); + }], + ['POST /api/v1/projects/p1/web-sessions/ws1/tree/fork', ({ body }) => { + assert.deepEqual(body, { targetId: 'u1', revision: 'rev-1' }); + return createWrappedJsonResponse({ item: { session: { id: 'ws2', agent: 'pi' }, tree: { ...tree, sessionId: 'ws2' }, editorText: 'Start' } }, 201); + }], + ['POST /api/v1/projects/p1/web-sessions/ws1/tree/clone', ({ body }) => { + assert.deepEqual(body, { revision: 'rev-1' }); + return createWrappedJsonResponse({ item: { session: { id: 'ws3', agent: 'pi' }, tree: { ...tree, sessionId: 'ws3' } } }, 201); + }], + ]); + const client = new CodeKanbanClient({ + baseURL: 'http://127.0.0.1:3000', + fetchImpl: createFetchMock(handlers), + WebSocketImpl: FakeWebSocket, + }); + + const read = await client.getWebSessionTree({ projectId: 'p1', sessionId: 'ws1' }); + assert.equal(read.nodes.length, 2); + const navigated = await client.navigateWebSessionTree({ + projectId: 'p1', sessionId: 'ws1', targetId: 'u1', revision: 'rev-1', summarize: true, + }); + assert.equal(navigated.editorText, 'Start'); + const forked = await client.forkWebSessionTree({ + projectId: 'p1', sessionId: 'ws1', targetId: 'u1', revision: 'rev-1', + }); + assert.equal(forked.session.id, 'ws2'); + assert.equal(forked.tree.sessionId, 'ws2'); + const cloned = await client.cloneWebSessionTree({ + projectId: 'p1', sessionId: 'ws1', revision: 'rev-1', + }); + assert.equal(cloned.session.id, 'ws3'); + assert.equal(cloned.editorText, ''); +}); + test('project file helpers call the file manager endpoints', async () => { const handlers = new Map([ ['GET /api/v1/projects/p1/files/scopes', () => createWrappedJsonResponse({ items: [{ id: 'scope-main', rootPath: '/repo/demo' }] })], diff --git a/packages/node-sdk/test/command-builder.test.js b/packages/node-sdk/test/command-builder.test.js index c2b113e6..b049d2d9 100644 --- a/packages/node-sdk/test/command-builder.test.js +++ b/packages/node-sdk/test/command-builder.test.js @@ -1,7 +1,7 @@ import test from 'node:test'; import assert from 'node:assert/strict'; -import { buildAgentLaunchSpec, composeWorkflowPrompt } from '../src/command-builder.js'; +import { AGENTS, buildAgentLaunchSpec, composeWorkflowPrompt } from '../src/command-builder.js'; test('buildAgentLaunchSpec builds codex plan profile with defaults', () => { const result = buildAgentLaunchSpec({ @@ -85,6 +85,32 @@ test('buildAgentLaunchSpec builds Claude CCR terminal command', () => { assert.equal(result.claudeRuntime, 'ccr'); }); +test('buildAgentLaunchSpec builds Pi plan profile without Codex permission flags', () => { + assert.deepEqual(AGENTS, ['codex', 'claude', 'pi']); + const result = buildAgentLaunchSpec({ + agent: 'pi', + profile: 'plan', + prompt: 'Inspect the repository', + extraArgs: ['--model', 'openai/gpt-5'], + }); + + assert.equal(result.command, 'pi --model openai/gpt-5'); + assert.match(result.prompt, /planning mode/i); + assert.doesNotMatch(result.command, /sandbox|approval|dangerously/i); +}); + +test('buildAgentLaunchSpec rejects structured permissions for Pi', () => { + assert.throws( + () => + buildAgentLaunchSpec({ + agent: 'pi', + prompt: 'Hello', + permissions: { sandbox: 'workspace-write' }, + }), + /structured permissions are not supported for pi/i, + ); +}); + test('buildAgentLaunchSpec rejects invalid Claude runtime', () => { assert.throws( () => diff --git a/packages/node-sdk/test/web-session-client.test.js b/packages/node-sdk/test/web-session-client.test.js index 2670493d..5782de25 100644 --- a/packages/node-sdk/test/web-session-client.test.js +++ b/packages/node-sdk/test/web-session-client.test.js @@ -5,8 +5,72 @@ import os from 'node:os'; import path from 'node:path'; import { CodeKanbanClient } from '../src/client.js'; +import { normalizeWebSessionRuntimeConfig } from '../src/web-session-shared.js'; import { createFetchMock, createJsonResponse, FakeWebSocket } from './helpers.js'; +test('runtime config normalizes Pi legacy fields without enabling unsupported sessions', () => { + const config = normalizeWebSessionRuntimeConfig({ + hasCodex: false, + hasClaudeCode: false, + hasPi: true, + piVersion: '0.84.1', + piRpcCompatible: true, + supportsPiWebSession: false, + piModels: [ + { + provider: 'anthropic', + id: 'claude-sonnet-4', + name: 'Claude Sonnet 4', + reasoning: true, + input: ['text', 'image'], + contextWindow: 200000, + }, + ], + }); + assert.deepEqual(config.piModels, [ + { + provider: 'anthropic', + id: 'claude-sonnet-4', + name: 'Claude Sonnet 4', + reasoning: true, + input: ['text', 'image'], + contextWindow: 200000, + }, + ]); + assert.deepEqual(config.agents.pi, { + installed: true, + version: '0.84.1', + supportsWebSession: false, + supportsTree: false, + supportsImages: false, + supportsCompaction: false, + supportsSteer: false, + supportsFollowUp: false, + supportsGoal: false, + supportsSubAgentRegistry: false, + permissionModes: [], + }); + + const enabled = normalizeWebSessionRuntimeConfig({ + hasPi: true, + piVersion: '0.84.1', + supportsPiWebSession: true, + }); + assert.deepEqual(enabled.agents.pi, { + installed: true, + version: '0.84.1', + supportsWebSession: true, + supportsTree: true, + supportsImages: true, + supportsCompaction: true, + supportsSteer: true, + supportsFollowUp: true, + supportsGoal: false, + supportsSubAgentRegistry: false, + permissionModes: [], + }); +}); + function createWebSessionSnapshot({ session = {}, items = [], @@ -200,12 +264,32 @@ test('CodeKanbanClient web session HTTP methods call the expected endpoints', as hasCodex: true, hasClaudeCode: false, codexVersion: '0.146.0', + hasPi: true, + piVersion: '0.84.1', + piRpcCompatible: true, + supportsPiWebSession: false, + piMinVersion: '0.84.1', supportsWebSession: true, webSessionMinCodexVersion: '', supportsMultiAgentV2: true, multiAgentV2MinCodexVersion: '0.146.0', supportsGoalMode: true, goalModeMinCodexVersion: '0.133.0', + agents: { + pi: { + installed: true, + version: '0.84.1', + supportsWebSession: false, + supportsTree: true, + supportsImages: false, + supportsCompaction: false, + supportsSteer: false, + supportsFollowUp: false, + supportsGoal: false, + supportsSubAgentRegistry: false, + permissionModes: [{ id: 'unrestricted', available: true }], + }, + }, }, })], ]); @@ -303,6 +387,17 @@ test('CodeKanbanClient web session HTTP methods call the expected endpoints', as assert.equal(runtimeConfig.supportsMultiAgentV2, true); assert.equal(runtimeConfig.multiAgentV2MinCodexVersion, '0.146.0'); assert.equal(runtimeConfig.supportsGoalMode, true); + assert.equal(runtimeConfig.agents.codex.supportsWebSession, true); + assert.equal(runtimeConfig.agents.claude.supportsWebSession, false); + assert.equal(runtimeConfig.hasPi, true); + assert.equal(runtimeConfig.piRpcCompatible, true); + assert.equal(runtimeConfig.agents.pi.installed, true); + assert.equal(runtimeConfig.agents.pi.version, '0.84.1'); + assert.equal(runtimeConfig.agents.pi.supportsWebSession, false); + assert.equal(runtimeConfig.agents.pi.supportsTree, true); + assert.deepEqual(runtimeConfig.agents.pi.permissionModes, [ + { id: 'unrestricted', available: true }, + ]); }); diff --git a/packages/node-sdk/test/web-session-command-channel.test.js b/packages/node-sdk/test/web-session-command-channel.test.js index 0b7e8118..9a1179d5 100644 --- a/packages/node-sdk/test/web-session-command-channel.test.js +++ b/packages/node-sdk/test/web-session-command-channel.test.js @@ -86,6 +86,7 @@ test('WebSessionCommandChannel connect returns a normalized snapshot', async () txt: 'adjust course', ra: 1710000005000, ps: true, + nq: true, ca: 1710000000004, }, ], @@ -103,6 +104,7 @@ test('WebSessionCommandChannel connect returns a normalized snapshot', async () attachmentIds: [], readyAt: '2024-03-09T16:00:05.000Z', paused: true, + nativeQueued: true, createdAt: '2024-03-09T16:00:00.004Z', }); channel.close(); @@ -204,6 +206,39 @@ test('WebSessionCommandChannel history sends the compact history payload and nor channel.close(); }); +test('WebSessionCommandChannel sends native compact without a message payload', async () => { + FakeWebSocket.reset(); + FakeWebSocket.setFactory(socket => { + queueMicrotask(() => socket.open()); + }); + + const channel = new WebSessionCommandChannel({ + url: 'ws://127.0.0.1:3000/api/v1/web-sessions/ws', + WebSocketImpl: FakeWebSocket, + }); + + await channel.waitForOpen(); + const promise = channel.compact('ws1'); + await new Promise(resolve => setTimeout(resolve, 0)); + + const socket = FakeWebSocket.instances[0]; + assert.equal(socket.sent[0].op, 'compact'); + assert.equal(socket.sent[0].sid, 'ws1'); + assert.deepEqual(socket.sent[0].p, {}); + socket.emitJson({ + v: 1, + k: 'ack', + rid: socket.sent[0].rid, + sid: 'ws1', + ts: 1710000000015, + op: 'compact', + ok: 1, + }); + + await promise; + channel.close(); +}); + test('WebSessionCommandChannel rejects sendMessage when the server follows the ack with an error frame', async () => { FakeWebSocket.reset(); FakeWebSocket.setFactory(socket => { @@ -301,6 +336,83 @@ test('WebSessionCommandChannel replies to heartbeat ping without disturbing pend channel.close(); }); +test('WebSessionCommandChannel supports Pi tree read and mutations', async () => { + FakeWebSocket.reset(); + FakeWebSocket.setFactory(socket => { + queueMicrotask(() => socket.open()); + }); + + const channel = new WebSessionCommandChannel({ + url: 'ws://127.0.0.1:3000/api/v1/web-sessions/ws', + WebSocketImpl: FakeWebSocket, + }); + await channel.waitForOpen(); + const socket = FakeWebSocket.instances[0]; + const tree = { + sessionId: 'ws1', + leafId: 'a1', + revision: 'rev-1', + nodes: [ + { id: 'u1', parentId: null, type: 'message', role: 'user', preview: 'Start', active: true, children: ['a1'] }, + { id: 'a1', parentId: 'u1', type: 'message', role: 'assistant', preview: 'Answer', active: true, children: [] }, + ], + }; + const resolveAck = async (promise, operation, payload) => { + await new Promise(resolve => setTimeout(resolve, 0)); + const frame = socket.sent.at(-1); + assert.equal(frame.op, operation); + socket.emitJson({ + v: 1, + k: 'ack', + rid: frame.rid, + sid: 'ws1', + ts: 1710000000040, + op: operation, + ok: 1, + p: payload, + }); + return await promise; + }; + + const readPromise = channel.getTree('ws1'); + const read = await resolveAck(readPromise, 'tree_get', tree); + assert.equal(read.leafId, 'a1'); + assert.equal(read.nodes[1].preview, 'Answer'); + + const navigatePromise = channel.navigateTree('ws1', { + targetId: 'u1', + revision: 'rev-1', + summarize: true, + }); + await new Promise(resolve => setTimeout(resolve, 0)); + assert.deepEqual(socket.sent.at(-1).p, { tid: 'u1', rev: 'rev-1', sum: true }); + const navigated = await resolveAck(navigatePromise, 'tree_nav', { tree, editorText: 'Start' }); + assert.equal(navigated.editorText, 'Start'); + + const forkPromise = channel.forkTree('ws1', { targetId: 'u1', revision: 'rev-1' }); + await new Promise(resolve => setTimeout(resolve, 0)); + assert.deepEqual(socket.sent.at(-1).p, { tid: 'u1', rev: 'rev-1' }); + const forked = await resolveAck(forkPromise, 'tree_fork', { + s: sampleWireSession({ id: 'ws2', ag: 'pi', nsid: 'pi-fork' }), + tree: { ...tree, sessionId: 'ws2' }, + editorText: 'Start', + }); + assert.equal(forked.session.id, 'ws2'); + assert.equal(forked.session.nativeSessionId, 'pi-fork'); + assert.equal(forked.tree.sessionId, 'ws2'); + + const clonePromise = channel.cloneTree('ws1', { revision: 'rev-1' }); + await new Promise(resolve => setTimeout(resolve, 0)); + assert.deepEqual(socket.sent.at(-1).p, { rev: 'rev-1' }); + const cloned = await resolveAck(clonePromise, 'tree_clone', { + s: sampleWireSession({ id: 'ws3', ag: 'pi', nsid: 'pi-clone' }), + tree: { ...tree, sessionId: 'ws3' }, + }); + assert.equal(cloned.session.id, 'ws3'); + assert.equal(cloned.editorText, ''); + channel.close(); +}); + test('WebSessionCommandChannel setGoal sends the expected payload', async () => { FakeWebSocket.reset(); FakeWebSocket.setFactory(socket => { diff --git a/service/ai_session_pi.go b/service/ai_session_pi.go new file mode 100644 index 00000000..a2d74b82 --- /dev/null +++ b/service/ai_session_pi.go @@ -0,0 +1,844 @@ +package service + +import ( + "bufio" + "context" + "encoding/base64" + "encoding/json" + "errors" + "os" + "path/filepath" + "runtime" + "sort" + "strings" + "time" + + "code-kanban/model" + "code-kanban/model/tables" + "code-kanban/utils" + "code-kanban/utils/ai_assistant2/log_watcher" + + "go.uber.org/zap" + "gorm.io/gorm" +) + +const piSessionDiscoveryBatchSize = 256 + +var errPiSessionDiscoveryBatchFull = errors.New("Pi session discovery batch is full") + +type piSessionEntry struct { + Type string `json:"type"` + ID string `json:"id"` + ParentID *string `json:"parentId"` + Timestamp string `json:"timestamp"` + Provider string `json:"provider"` + ModelID string `json:"modelId"` + Name *string `json:"name"` + Message json.RawMessage `json:"message"` +} + +type piSessionMessage struct { + Role string `json:"role"` + Content json.RawMessage `json:"content"` + Provider string `json:"provider"` + Model string `json:"model"` + Timestamp int64 `json:"timestamp"` +} + +type piSessionData struct { + SessionID string + Cwd string + Model string + Title string + StartedAt time.Time + LastMessageAt *time.Time + MessageCount int + AssistantMessageCount int + hasModelChange bool + activePath []*piSessionEntry +} + +type piSessionFileCandidate struct { + path string + info os.FileInfo + header log_watcher.PiSessionHeader +} + +func (s *AISessionService) getPiSessionsPhased( + ctx context.Context, + projectPath string, +) ([]*AISessionSummary, string, error) { + root, err := log_watcher.ResolvePiSessionDir() + if err != nil { + return nil, ScanPhaseComplete, err + } + if _, err := os.Stat(root); errors.Is(err, os.ErrNotExist) { + return nil, ScanPhaseComplete, nil + } else if err != nil { + return nil, ScanPhaseComplete, err + } + + if sessions, phase, scanning := piScanSnapshot(projectPath, root); scanning { + s.queuePiBackgroundScan(ctx, projectPath, root) + return dedupeAISessionSummariesBySessionID(sessions), phase, nil + } + + cacheKey := piSessionCacheKey(projectPath, root) + dirCacheMu.RLock() + cached, hasCached := dirCache[cacheKey] + dirCacheMu.RUnlock() + if hasCached && time.Since(cached.cachedAt) < dirCacheTTL { + cachedSessions := append([]*AISessionSummary(nil), cached.sessions...) + phase := s.getPiScanPhase(projectPath, root) + if phase != ScanPhaseComplete { + s.queuePiBackgroundScan(ctx, projectPath, root) + } + return dedupeAISessionSummariesBySessionID(cachedSessions), phase, nil + } + + db := model.GetDB() + if db == nil { + return nil, ScanPhaseComplete, model.ErrDBNotInitialized + } + var cachedRows []tables.AISessionTable + if err := db.WithContext(ctx). + Where("type = ?", tables.AISessionTypePi). + Find(&cachedRows).Error; err != nil { + return nil, ScanPhaseComplete, err + } + cachedByID := make(map[string]tables.AISessionTable, len(cachedRows)) + for _, row := range cachedRows { + cachedByID[row.SessionID] = row + } + + projectComparable := canonicalAIProjectPath(projectPath) + candidates, nextCursor, hasMore, err := listPiSessionCandidateBatch(ctx, root, projectComparable, "") + if err != nil { + return nil, ScanPhaseComplete, err + } + + now := time.Now() + sessions := make([]*AISessionSummary, 0, len(candidates)) + pendingFiles := make([]string, 0) + for _, candidate := range candidates { + if row, ok := cachedByID[candidate.header.ID]; ok && + filepath.Clean(row.FilePath) == filepath.Clean(candidate.path) && + row.FileModTime.Equal(candidate.info.ModTime()) && row.FileSize == candidate.info.Size() { + sessions = append(sessions, aiSessionSummaryFromRecord(row)) + continue + } + if now.Sub(candidate.info.ModTime()) > recentThreshold { + pendingFiles = append(pendingFiles, candidate.path) + continue + } + data, err := s.parsePiSessionFile(candidate.path) + if err != nil { + s.logger(ctx).Debug("failed to parse recent Pi session", + zap.String("file", candidate.path), + zap.Error(err)) + continue + } + session, err := s.savePiSession(ctx, db, candidate.path, candidate.info, data) + if err == nil { + sessions = append(sessions, session) + } + } + sessions = dedupeAISessionSummariesBySessionID(sessions) + sortAISessionSummaries(sessions) + + phase := ScanPhaseComplete + stateKey := piScanStateKey(projectPath, root) + if len(pendingFiles) > 0 || hasMore { + phase = ScanPhaseRecent + scanStatesMu.Lock() + state, exists := scanStates[stateKey] + if !exists { + state = &scanState{} + scanStates[stateKey] = state + } + state.mu.Lock() + state.phase = ScanPhaseRecent + state.sessions = append([]*AISessionSummary(nil), sessions...) + state.pendingDirs = pendingFiles + state.cursor = nextCursor + state.hasMore = hasMore + state.mu.Unlock() + scanStatesMu.Unlock() + s.queuePiBackgroundScan(ctx, projectPath, root) + } else { + scanStatesMu.Lock() + if state, exists := scanStates[stateKey]; exists { + state.mu.Lock() + state.phase = ScanPhaseComplete + state.pendingDirs = nil + state.sessions = sessions + state.cursor = "" + state.hasMore = false + state.mu.Unlock() + } + scanStatesMu.Unlock() + } + + dirCacheMu.Lock() + dirCache[cacheKey] = &dirCacheEntry{ + sessions: append([]*AISessionSummary(nil), sessions...), + cachedAt: time.Now(), + } + dirCacheMu.Unlock() + return sessions, phase, nil +} + +func piScanSnapshot(projectPath, root string) ([]*AISessionSummary, string, bool) { + scanStatesMu.RLock() + state := scanStates[piScanStateKey(projectPath, root)] + scanStatesMu.RUnlock() + if state == nil { + return nil, ScanPhaseComplete, false + } + state.mu.RLock() + defer state.mu.RUnlock() + if state.phase == ScanPhaseComplete { + return nil, ScanPhaseComplete, false + } + return append([]*AISessionSummary(nil), state.sessions...), state.phase, true +} + +func (s *AISessionService) queuePiBackgroundScan(ctx context.Context, projectPath, root string) { + stateKey := piScanStateKey(projectPath, root) + scanStatesMu.RLock() + state := scanStates[stateKey] + scanStatesMu.RUnlock() + if state == nil { + return + } + state.mu.Lock() + if state.phase == ScanPhaseComplete || state.backgroundActive { + state.mu.Unlock() + return + } + state.backgroundActive = true + state.mu.Unlock() + + select { + case bgScanQueue <- &bgScanTask{projectPath: projectPath, scanType: "pi", projectDir: root}: + default: + state.mu.Lock() + state.backgroundActive = false + state.mu.Unlock() + s.logger(ctx).Debug("background scan queue full, deferring Pi scan") + } +} + +func (s *AISessionService) scanPiExtendedPhase(ctx context.Context, projectPath, root string) error { + stateKey := piScanStateKey(projectPath, root) + scanStatesMu.RLock() + state, exists := scanStates[stateKey] + scanStatesMu.RUnlock() + if !exists { + return nil + } + defer func() { + state.mu.Lock() + state.backgroundActive = false + state.mu.Unlock() + }() + + db := model.GetDB() + if db == nil { + return model.ErrDBNotInitialized + } + projectComparable := canonicalAIProjectPath(projectPath) + + state.mu.RLock() + pendingFiles := append([]string(nil), state.pendingDirs...) + cursor := state.cursor + hasMore := state.hasMore + state.mu.RUnlock() + + newSessions := make([]*AISessionSummary, 0, len(pendingFiles)) + for _, filePath := range pendingFiles { + if ctx.Err() != nil { + return ctx.Err() + } + session, err := s.loadPiSessionCandidate(ctx, db, filePath, projectComparable) + if err != nil { + s.logger(ctx).Debug("failed to parse older Pi session", + zap.String("file", filePath), + zap.Error(err)) + continue + } + if session != nil { + newSessions = append(newSessions, session) + } + } + allSessions := updatePiScanProgress(state, newSessions, nil, cursor, hasMore) + updatePiDirectoryCache(projectPath, root, allSessions) + + for hasMore { + if ctx.Err() != nil { + return ctx.Err() + } + candidates, nextCursor, nextHasMore, err := listPiSessionCandidateBatch( + ctx, + root, + projectComparable, + cursor, + ) + if err != nil { + return err + } + batchSessions := make([]*AISessionSummary, 0, len(candidates)) + for _, candidate := range candidates { + session, err := s.loadPiSessionCandidate(ctx, db, candidate.path, projectComparable) + if err != nil { + s.logger(ctx).Debug("failed to parse Pi discovery batch entry", + zap.String("file", candidate.path), + zap.Error(err)) + continue + } + if session != nil { + batchSessions = append(batchSessions, session) + } + } + cursor = nextCursor + hasMore = nextHasMore + allSessions = updatePiScanProgress(state, batchSessions, nil, cursor, hasMore) + updatePiDirectoryCache(projectPath, root, allSessions) + runtime.Gosched() + } + + state.mu.Lock() + state.phase = ScanPhaseComplete + state.cursor = "" + state.hasMore = false + state.pendingDirs = nil + allSessions = append([]*AISessionSummary(nil), state.sessions...) + state.mu.Unlock() + updatePiDirectoryCache(projectPath, root, allSessions) + return nil +} + +func (s *AISessionService) loadPiSessionCandidate( + ctx context.Context, + db *gorm.DB, + filePath string, + projectComparable string, +) (*AISessionSummary, error) { + header, err := log_watcher.ReadPiSessionHeader(filePath) + if err != nil || canonicalAIProjectPath(header.Cwd) != projectComparable { + return nil, err + } + info, err := os.Stat(filePath) + if err != nil { + return nil, err + } + var cached tables.AISessionTable + if err := db.WithContext(ctx). + Where("session_id = ? AND type = ?", header.ID, tables.AISessionTypePi). + First(&cached).Error; err == nil && + filepath.Clean(cached.FilePath) == filepath.Clean(filePath) && + cached.FileModTime.Equal(info.ModTime()) && cached.FileSize == info.Size() { + return aiSessionSummaryFromRecord(cached), nil + } + data, err := s.parsePiSessionFile(filePath) + if err != nil { + return nil, err + } + return s.savePiSession(ctx, db, filePath, info, data) +} + +func updatePiScanProgress( + state *scanState, + newSessions []*AISessionSummary, + pendingFiles []string, + cursor string, + hasMore bool, +) []*AISessionSummary { + state.mu.Lock() + state.sessions = dedupeAISessionSummariesBySessionID(append(state.sessions, newSessions...)) + sortAISessionSummaries(state.sessions) + state.pendingDirs = pendingFiles + state.cursor = cursor + state.hasMore = hasMore + if hasMore || len(pendingFiles) > 0 { + state.phase = ScanPhaseExtended + } else { + state.phase = ScanPhaseComplete + } + allSessions := append([]*AISessionSummary(nil), state.sessions...) + state.mu.Unlock() + return allSessions +} + +func updatePiDirectoryCache(projectPath, root string, sessions []*AISessionSummary) { + dirCacheMu.Lock() + dirCache[piSessionCacheKey(projectPath, root)] = &dirCacheEntry{ + sessions: append([]*AISessionSummary(nil), sessions...), + cachedAt: time.Now(), + } + dirCacheMu.Unlock() +} + +func (s *AISessionService) getPiScanPhase(projectPath, root string) string { + scanStatesMu.RLock() + state := scanStates[piScanStateKey(projectPath, root)] + scanStatesMu.RUnlock() + if state == nil { + return ScanPhaseComplete + } + state.mu.RLock() + defer state.mu.RUnlock() + return state.phase +} + +func (s *AISessionService) currentPiScanProgress(projectPath, fallbackPhase string) (string, string) { + root, err := log_watcher.ResolvePiSessionDir() + if err != nil { + return fallbackPhase, "" + } + scanStatesMu.RLock() + state := scanStates[piScanStateKey(projectPath, root)] + scanStatesMu.RUnlock() + if state == nil { + return fallbackPhase, "" + } + state.mu.RLock() + defer state.mu.RUnlock() + if state.phase == ScanPhaseComplete || strings.TrimSpace(state.cursor) == "" { + return state.phase, "" + } + return state.phase, base64.RawURLEncoding.EncodeToString([]byte(state.cursor)) +} + +func piSessionCacheKey(projectPath, root string) string { + return canonicalAIProjectPath(projectPath) + ":pi:" + canonicalAIProjectPath(root) +} + +func piScanStateKey(projectPath, root string) string { + return piSessionCacheKey(projectPath, root) + ":scan" +} + +func canonicalAIProjectPath(value string) string { + value = filepath.Clean(filepath.FromSlash(strings.TrimSpace(value))) + if absolute, err := filepath.Abs(value); err == nil { + value = absolute + } + if resolved, err := filepath.EvalSymlinks(value); err == nil && strings.TrimSpace(resolved) != "" { + value = filepath.Clean(resolved) + } + if runtime.GOOS == "windows" { + value = strings.ToLower(value) + } + return value +} + +func listPiSessionCandidateBatch( + ctx context.Context, + root string, + projectComparable string, + afterCursor string, +) ([]piSessionFileCandidate, string, bool, error) { + candidates := make([]piSessionFileCandidate, 0, piSessionDiscoveryBatchSize) + visited := 0 + nextCursor := afterCursor + root = filepath.Clean(root) + afterCursor = filepath.Clean(afterCursor) + if afterCursor == "." { + afterCursor = "" + } + + err := filepath.WalkDir(root, func(path string, entry os.DirEntry, walkErr error) error { + if ctx.Err() != nil { + return ctx.Err() + } + if walkErr != nil || entry == nil { + return nil + } + if afterCursor != "" && path <= afterCursor { + if entry.IsDir() && path != root && + !strings.HasPrefix(afterCursor, path+string(os.PathSeparator)) { + return filepath.SkipDir + } + return nil + } + if entry.IsDir() || !strings.HasSuffix(strings.ToLower(entry.Name()), ".jsonl") { + return nil + } + + visited++ + nextCursor = path + header, err := log_watcher.ReadPiSessionHeader(path) + if err == nil && canonicalAIProjectPath(header.Cwd) == projectComparable { + if info, infoErr := entry.Info(); infoErr == nil { + candidates = append(candidates, piSessionFileCandidate{path: path, info: info, header: header}) + } + } + if visited >= piSessionDiscoveryBatchSize { + return errPiSessionDiscoveryBatchFull + } + return nil + }) + switch { + case errors.Is(err, errPiSessionDiscoveryBatchFull): + return candidates, nextCursor, true, nil + case errors.Is(err, os.ErrNotExist): + return candidates, "", false, nil + case err != nil: + return candidates, nextCursor, false, err + default: + return candidates, "", false, nil + } +} + +func (s *AISessionService) parsePiSessionFile(filePath string) (*piSessionData, error) { + header, err := log_watcher.ReadPiSessionHeader(filePath) + if err != nil { + return nil, err + } + file, err := os.Open(filePath) + if err != nil { + return nil, err + } + defer file.Close() + + entries := make([]*piSessionEntry, 0, 128) + byID := make(map[string]*piSessionEntry) + var leaf *piSessionEntry + latestName := "" + scanner := bufio.NewScanner(file) + scanner.Buffer(make([]byte, 0, 64*1024), 4*1024*1024) + lineIndex := 0 + for scanner.Scan() { + lineIndex++ + if lineIndex == 1 || strings.TrimSpace(scanner.Text()) == "" { + continue + } + entry := &piSessionEntry{} + if json.Unmarshal(scanner.Bytes(), entry) != nil || strings.TrimSpace(entry.ID) == "" { + continue + } + entries = append(entries, entry) + byID[entry.ID] = entry + leaf = entry + if entry.Type == "session_info" && entry.Name != nil { + latestName = strings.TrimSpace(*entry.Name) + } + } + if err := scanner.Err(); err != nil { + return nil, err + } + + activePath := piActivePath(entries, byID, leaf) + startedAt, _ := time.Parse(time.RFC3339Nano, header.Timestamp) + if startedAt.IsZero() { + if info, statErr := os.Stat(filePath); statErr == nil { + startedAt = info.ModTime() + } + } + data := &piSessionData{ + SessionID: header.ID, + Cwd: filepath.Clean(header.Cwd), + Title: latestName, + StartedAt: startedAt, + activePath: activePath, + } + for _, entry := range activePath { + switch entry.Type { + case "model_change": + if model := canonicalPiModel(entry.Provider, entry.ModelID); model != "" { + data.Model = model + data.hasModelChange = true + } + case "message": + message, ok := decodePiSessionMessage(entry.Message) + if !ok { + continue + } + ts := piEntryMessageTime(entry.Timestamp, message.Timestamp) + switch message.Role { + case "user": + data.MessageCount++ + data.LastMessageAt = copyTimePointer(ts) + if data.Title == "" { + data.Title = truncateRunes(piSessionContentText(message.Content), 100) + } + case "assistant": + data.AssistantMessageCount++ + data.LastMessageAt = copyTimePointer(ts) + if model := canonicalPiModel(message.Provider, message.Model); !data.hasModelChange && model != "" { + data.Model = model + } + } + } + } + return data, nil +} + +func piActivePath( + entries []*piSessionEntry, + byID map[string]*piSessionEntry, + leaf *piSessionEntry, +) []*piSessionEntry { + if leaf == nil { + return nil + } + path := make([]*piSessionEntry, 0, len(entries)) + seen := make(map[string]struct{}, len(entries)) + for current := leaf; current != nil; { + if _, exists := seen[current.ID]; exists { + break + } + seen[current.ID] = struct{}{} + path = append(path, current) + if current.ParentID == nil || strings.TrimSpace(*current.ParentID) == "" { + break + } + current = byID[*current.ParentID] + } + for left, right := 0, len(path)-1; left < right; left, right = left+1, right-1 { + path[left], path[right] = path[right], path[left] + } + return path +} + +func decodePiSessionMessage(raw json.RawMessage) (piSessionMessage, bool) { + var message piSessionMessage + if len(raw) == 0 || json.Unmarshal(raw, &message) != nil || strings.TrimSpace(message.Role) == "" { + return piSessionMessage{}, false + } + return message, true +} + +func piSessionContentText(raw json.RawMessage) string { + var text string + if json.Unmarshal(raw, &text) == nil { + return strings.TrimSpace(text) + } + var blocks []struct { + Type string `json:"type"` + Text string `json:"text"` + Thinking string `json:"thinking"` + } + if json.Unmarshal(raw, &blocks) != nil { + return "" + } + parts := make([]string, 0, len(blocks)) + for _, block := range blocks { + if block.Type == "text" && strings.TrimSpace(block.Text) != "" { + parts = append(parts, strings.TrimSpace(block.Text)) + } + } + return strings.Join(parts, "\n") +} + +func canonicalPiModel(provider, modelID string) string { + provider = strings.TrimSpace(provider) + modelID = strings.TrimSpace(modelID) + if provider == "" { + return modelID + } + if modelID == "" { + return provider + } + return provider + "/" + modelID +} + +func piEntryMessageTime(entryTimestamp string, messageTimestamp int64) time.Time { + if messageTimestamp > 0 { + return time.UnixMilli(messageTimestamp) + } + ts, _ := time.Parse(time.RFC3339Nano, entryTimestamp) + return ts +} + +func copyTimePointer(value time.Time) *time.Time { + if value.IsZero() { + return nil + } + copy := value + return © +} + +func truncateRunes(value string, limit int) string { + value = strings.TrimSpace(value) + runes := []rune(value) + if limit <= 0 || len(runes) <= limit { + return value + } + return string(runes[:limit]) + "..." +} + +func sortAISessionSummaries(sessions []*AISessionSummary) { + sort.Slice(sessions, func(i, j int) bool { + left := sessions[i].LastMessageAt + right := sessions[j].LastMessageAt + if left != nil && right != nil && !left.Equal(*right) { + return left.After(*right) + } + if left != nil && right == nil { + return true + } + if left == nil && right != nil { + return false + } + return sessions[i].SessionStartedAt.After(sessions[j].SessionStartedAt) + }) +} + +func aiSessionSummaryFromRecord(record tables.AISessionTable) *AISessionSummary { + return &AISessionSummary{ + ID: record.ID, + SessionID: record.SessionID, + Type: string(record.Type), + Model: record.Model, + Title: record.Title, + SessionStartedAt: record.SessionStartedAt, + LastMessageAt: record.LastMessageAt, + MessageCount: record.MessageCount, + AssistantMessageCount: record.AssistantMessageCount, + FilePath: record.FilePath, + } +} + +func (s *AISessionService) savePiSession( + ctx context.Context, + db *gorm.DB, + filePath string, + fileInfo os.FileInfo, + data *piSessionData, +) (*AISessionSummary, error) { + var existing tables.AISessionTable + err := db.WithContext(ctx). + Where("session_id = ? AND type = ?", data.SessionID, tables.AISessionTypePi). + First(&existing).Error + now := time.Now() + record := tables.AISessionTable{ + SessionID: data.SessionID, + Type: tables.AISessionTypePi, + ProjectPath: data.Cwd, + FilePath: filePath, + Model: data.Model, + Title: data.Title, + SessionStartedAt: data.StartedAt, + LastMessageAt: data.LastMessageAt, + MessageCount: data.MessageCount, + AssistantMessageCount: data.AssistantMessageCount, + FileModTime: fileInfo.ModTime(), + FileSize: fileInfo.Size(), + } + if err == nil { + record.ID = existing.ID + record.CreatedAt = existing.CreatedAt + record.UpdatedAt = now + if err := db.WithContext(ctx).Save(&record).Error; err != nil { + return nil, err + } + } else if errors.Is(err, gorm.ErrRecordNotFound) { + record.ID = utils.NewID() + record.CreatedAt = now + record.UpdatedAt = now + if err := db.WithContext(ctx).Create(&record).Error; err != nil { + return nil, err + } + } else { + return nil, err + } + return aiSessionSummaryFromRecord(record), nil +} + +func (s *AISessionService) ResolvePiSessionBySessionID( + ctx context.Context, + sessionID string, +) (*tables.AISessionTable, error) { + ctx = ensureContext(ctx) + db := model.GetDB() + if db == nil { + return nil, model.ErrDBNotInitialized + } + sessionID = strings.TrimSpace(sessionID) + if sessionID == "" { + return nil, gorm.ErrRecordNotFound + } + + var cached tables.AISessionTable + if err := db.WithContext(ctx). + Where("session_id = ? AND type = ?", sessionID, tables.AISessionTypePi). + First(&cached).Error; err == nil { + if info, statErr := os.Stat(cached.FilePath); statErr == nil && + cached.FileModTime.Equal(info.ModTime()) && cached.FileSize == info.Size() { + return &cached, nil + } + } + + root, err := log_watcher.ResolvePiSessionDir() + if err != nil { + return nil, err + } + searcher := log_watcher.NewPiFileSearcherWithSessionDir(root, "") + filePath, err := searcher.FindBySessionID(ctx, sessionID) + if err != nil || filePath == "" { + if err == nil { + err = gorm.ErrRecordNotFound + } + return nil, err + } + info, err := os.Stat(filePath) + if err != nil { + return nil, err + } + data, err := s.parsePiSessionFile(filePath) + if err != nil { + return nil, err + } + if _, err := s.savePiSession(ctx, db, filePath, info, data); err != nil { + return nil, err + } + var record tables.AISessionTable + if err := db.WithContext(ctx). + Where("session_id = ? AND type = ?", sessionID, tables.AISessionTypePi). + First(&record).Error; err != nil { + return nil, err + } + return &record, nil +} + +func (s *AISessionService) ResolvePiSessionByID(ctx context.Context, dbID string) (*tables.AISessionTable, error) { + db := model.GetDB() + if db == nil { + return nil, model.ErrDBNotInitialized + } + var record tables.AISessionTable + if err := db.WithContext(ensureContext(ctx)). + Where("id = ? AND type = ?", strings.TrimSpace(dbID), tables.AISessionTypePi). + First(&record).Error; err != nil { + return nil, err + } + return s.ResolvePiSessionBySessionID(ctx, record.SessionID) +} + +func (s *AISessionService) parsePiConversation(filePath string) ([]*ConversationMessage, error) { + data, err := s.parsePiSessionFile(filePath) + if err != nil { + return nil, err + } + messages := make([]*ConversationMessage, 0, data.MessageCount+data.AssistantMessageCount) + for _, entry := range data.activePath { + if entry.Type != "message" { + continue + } + message, ok := decodePiSessionMessage(entry.Message) + if !ok || (message.Role != "user" && message.Role != "assistant") { + continue + } + content := piSessionContentText(message.Content) + if content == "" { + continue + } + messages = append(messages, &ConversationMessage{ + Role: message.Role, + Content: content, + Timestamp: piEntryMessageTime(entry.Timestamp, message.Timestamp), + }) + } + return messages, nil +} diff --git a/service/ai_session_pi_test.go b/service/ai_session_pi_test.go new file mode 100644 index 00000000..5c119dd4 --- /dev/null +++ b/service/ai_session_pi_test.go @@ -0,0 +1,232 @@ +package service + +import ( + "context" + "fmt" + "os" + "path/filepath" + "strings" + "testing" + "time" + + "code-kanban/model" + "code-kanban/model/tables" +) + +func writeServicePiSession(t *testing.T, root, dirName, id, cwd string, started time.Time, branched bool) string { + t.Helper() + dir := filepath.Join(root, dirName) + if err := os.MkdirAll(dir, 0o755); err != nil { + t.Fatalf("create fixture dir: %v", err) + } + lines := []string{ + fmt.Sprintf(`{"type":"session","version":3,"id":%q,"timestamp":%q,"cwd":%q}`, id, started.UTC().Format(time.RFC3339Nano), cwd), + fmt.Sprintf(`{"type":"model_change","id":"model001","parentId":null,"timestamp":%q,"provider":"openai","modelId":"gpt-5"}`, started.Add(time.Second).UTC().Format(time.RFC3339Nano)), + fmt.Sprintf(`{"type":"message","id":"user0001","parentId":"model001","timestamp":%q,"message":{"role":"user","content":"root request","timestamp":%d}}`, started.Add(2*time.Second).UTC().Format(time.RFC3339Nano), started.Add(2*time.Second).UnixMilli()), + fmt.Sprintf(`{"type":"message","id":"assist01","parentId":"user0001","timestamp":%q,"message":{"role":"assistant","content":[{"type":"text","text":"root answer"}],"provider":"openai","model":"gpt-5","timestamp":%d}}`, started.Add(3*time.Second).UTC().Format(time.RFC3339Nano), started.Add(3*time.Second).UnixMilli()), + } + if branched { + lines = append(lines, + fmt.Sprintf(`{"type":"message","id":"olduser1","parentId":"assist01","timestamp":%q,"message":{"role":"user","content":"abandoned request","timestamp":%d}}`, started.Add(4*time.Second).UTC().Format(time.RFC3339Nano), started.Add(4*time.Second).UnixMilli()), + fmt.Sprintf(`{"type":"message","id":"oldasst1","parentId":"olduser1","timestamp":%q,"message":{"role":"assistant","content":[{"type":"text","text":"abandoned answer"}],"provider":"other","model":"old","timestamp":%d}}`, started.Add(5*time.Second).UTC().Format(time.RFC3339Nano), started.Add(5*time.Second).UnixMilli()), + fmt.Sprintf(`{"type":"message","id":"activeu1","parentId":"assist01","timestamp":%q,"message":{"role":"user","content":"active request","timestamp":%d}}`, started.Add(6*time.Second).UTC().Format(time.RFC3339Nano), started.Add(6*time.Second).UnixMilli()), + fmt.Sprintf(`{"type":"message","id":"activea1","parentId":"activeu1","timestamp":%q,"message":{"role":"assistant","content":[{"type":"text","text":"active answer"}],"provider":"openai","model":"gpt-5","timestamp":%d}}`, started.Add(7*time.Second).UTC().Format(time.RFC3339Nano), started.Add(7*time.Second).UnixMilli()), + fmt.Sprintf(`{"type":"session_info","id":"info0001","parentId":"activea1","timestamp":%q,"name":"Named Pi Session"}`, started.Add(8*time.Second).UTC().Format(time.RFC3339Nano)), + fmt.Sprintf(`{"type":"custom","id":"marker01","parentId":"info0001","timestamp":%q,"customType":"codekanban.active-leaf.v1"}`, started.Add(9*time.Second).UTC().Format(time.RFC3339Nano)), + ) + } + filePath := filepath.Join(dir, started.UTC().Format("20060102T150405")+"_"+id+".jsonl") + if err := os.WriteFile(filePath, []byte(strings.Join(lines, "\n")+"\n"), 0o644); err != nil { + t.Fatalf("write fixture: %v", err) + } + return filePath +} + +func resetPiHistoryTestCaches() { + dirCacheMu.Lock() + dirCache = make(map[string]*dirCacheEntry) + dirCacheMu.Unlock() + scanStatesMu.Lock() + scanStates = make(map[string]*scanState) + scanStatesMu.Unlock() +} + +func TestParsePiSessionUsesActiveBranch(t *testing.T) { + root := t.TempDir() + project := filepath.Join(root, "project") + started := time.Date(2026, 8, 1, 10, 0, 0, 0, time.UTC) + filePath := writeServicePiSession(t, root, "encoded", "pi-session", project, started, true) + + svc := NewAISessionService() + data, err := svc.parsePiSessionFile(filePath) + if err != nil { + t.Fatalf("parsePiSessionFile: %v", err) + } + if data.SessionID != "pi-session" || data.Cwd != filepath.Clean(project) { + t.Fatalf("unexpected identity: id=%q cwd=%q", data.SessionID, data.Cwd) + } + if data.Title != "Named Pi Session" || data.Model != "openai/gpt-5" { + t.Fatalf("unexpected metadata: title=%q model=%q", data.Title, data.Model) + } + if data.MessageCount != 2 || data.AssistantMessageCount != 2 { + t.Fatalf("unexpected active counts: user=%d assistant=%d", data.MessageCount, data.AssistantMessageCount) + } + + messages, err := svc.parsePiConversation(filePath) + if err != nil { + t.Fatalf("parsePiConversation: %v", err) + } + if len(messages) != 4 { + t.Fatalf("active conversation length = %d, want 4", len(messages)) + } + joined := "" + for _, message := range messages { + joined += message.Content + "\n" + } + if strings.Contains(joined, "abandoned") || !strings.Contains(joined, "active request") { + t.Fatalf("conversation did not follow active branch: %q", joined) + } +} + +func TestGetPiSessionsKeepsInProgressScanStatePastCacheTTL(t *testing.T) { + resetPiHistoryTestCaches() + root := t.TempDir() + project := filepath.Join(root, "project") + t.Setenv("PI_CODING_AGENT_SESSION_DIR", root) + stateKey := piScanStateKey(project, root) + expected := &AISessionSummary{SessionID: "already-discovered", Type: string(tables.AISessionTypePi)} + scanStatesMu.Lock() + scanStates[stateKey] = &scanState{ + phase: ScanPhaseExtended, + sessions: []*AISessionSummary{expected}, + cursor: filepath.Join(root, "cursor.jsonl"), + hasMore: true, + backgroundActive: true, + } + scanStatesMu.Unlock() + dirCacheMu.Lock() + dirCache[piSessionCacheKey(project, root)] = &dirCacheEntry{ + sessions: []*AISessionSummary{}, + cachedAt: time.Now().Add(-2 * dirCacheTTL), + } + dirCacheMu.Unlock() + + sessions, phase, err := NewAISessionService().getPiSessionsPhased(context.Background(), project) + if err != nil { + t.Fatalf("getPiSessionsPhased returned error: %v", err) + } + if phase != ScanPhaseExtended || len(sessions) != 1 || sessions[0].SessionID != expected.SessionID { + t.Fatalf("in-progress snapshot = %#v, phase=%q", sessions, phase) + } + state := scanStates[stateKey] + state.mu.RLock() + defer state.mu.RUnlock() + if state.cursor == "" || !state.hasMore || len(state.sessions) != 1 { + t.Fatalf("in-progress state was reset: %#v", state) + } +} + +func TestPiSessionDiscoveryUsesBoundedCursorBatches(t *testing.T) { + root := t.TempDir() + project := filepath.Join(root, "project") + if err := os.MkdirAll(project, 0o755); err != nil { + t.Fatalf("create project: %v", err) + } + started := time.Date(2026, 8, 1, 10, 0, 0, 0, time.UTC) + for index := 0; index < piSessionDiscoveryBatchSize+44; index++ { + id := fmt.Sprintf("session-%03d", index) + writeServicePiSession(t, root, "sessions", id, project, started.Add(time.Duration(index)*time.Second), false) + } + + first, cursor, hasMore, err := listPiSessionCandidateBatch( + context.Background(), + root, + canonicalAIProjectPath(project), + "", + ) + if err != nil { + t.Fatalf("first discovery batch: %v", err) + } + if len(first) != piSessionDiscoveryBatchSize || !hasMore || cursor == "" { + t.Fatalf("first batch count=%d cursor=%q hasMore=%v", len(first), cursor, hasMore) + } + second, secondCursor, secondHasMore, err := listPiSessionCandidateBatch( + context.Background(), + root, + canonicalAIProjectPath(project), + cursor, + ) + if err != nil { + t.Fatalf("second discovery batch: %v", err) + } + if len(second) != 44 || secondHasMore || secondCursor != "" { + t.Fatalf("second batch count=%d cursor=%q hasMore=%v", len(second), secondCursor, secondHasMore) + } +} + +func TestProjectAISessionsIndexesRecentAndOlderPiHistory(t *testing.T) { + cleanup := initTestDB(t) + defer cleanup() + resetPiHistoryTestCaches() + + home := t.TempDir() + root := filepath.Join(home, "pi-sessions") + project := filepath.Join(home, "project") + otherProject := filepath.Join(home, "other") + if err := os.MkdirAll(project, 0o755); err != nil { + t.Fatalf("create project: %v", err) + } + t.Setenv("HOME", home) + t.Setenv("USERPROFILE", home) + t.Setenv("PI_CODING_AGENT_SESSION_DIR", root) + + now := time.Now().UTC() + writeServicePiSession(t, root, "recent-dir", "recent-session", project, now.Add(-time.Hour), true) + oldPath := writeServicePiSession(t, root, "old-dir", "old-session", project, now.Add(-30*24*time.Hour), false) + oldTime := now.Add(-30 * 24 * time.Hour) + if err := os.Chtimes(oldPath, oldTime, oldTime); err != nil { + t.Fatalf("age old fixture: %v", err) + } + writeServicePiSession(t, root, "misleading-project-dir", "other-session", otherProject, now, false) + if err := os.WriteFile(filepath.Join(root, "broken.jsonl"), []byte("not-json\n"), 0o644); err != nil { + t.Fatalf("write broken fixture: %v", err) + } + + svc := NewAISessionService() + var result *ProjectAISessions + deadline := time.Now().Add(4 * time.Second) + for time.Now().Before(deadline) { + var err error + result, err = svc.GetProjectAISessions(context.Background(), project) + if err != nil { + t.Fatalf("GetProjectAISessions: %v", err) + } + if result.PiScanPhase == ScanPhaseComplete && len(result.PiSessions) == 2 { + break + } + time.Sleep(20 * time.Millisecond) + } + if result == nil || !result.HasPi || len(result.PiSessions) != 2 { + t.Fatalf("unexpected Pi history result: %#v", result) + } + ids := map[string]bool{} + for _, session := range result.PiSessions { + ids[session.SessionID] = true + if session.Type != string(tables.AISessionTypePi) { + t.Fatalf("session type = %q, want pi", session.Type) + } + } + if !ids["recent-session"] || !ids["old-session"] || ids["other-session"] { + t.Fatalf("unexpected indexed session IDs: %#v", ids) + } + + var count int64 + if err := model.GetDB().Model(&tables.AISessionTable{}). + Where("type = ?", tables.AISessionTypePi).Count(&count).Error; err != nil || count != 2 { + t.Fatalf("Pi cache count = %d, err=%v", count, err) + } + conversation, err := svc.GetSessionConversationBySessionID(context.Background(), "recent-session") + if err != nil || len(conversation.Messages) != 4 { + t.Fatalf("Pi conversation by session ID = %#v, %v", conversation, err) + } +} diff --git a/service/ai_session_service.go b/service/ai_session_service.go index ef19675e..bf7826cc 100644 --- a/service/ai_session_service.go +++ b/service/ai_session_service.go @@ -52,10 +52,13 @@ const ( // scanState tracks the scanning progress for a project type scanState struct { - phase string // Current scan phase - sessions []*AISessionSummary // Accumulated sessions - pendingDirs []string // Directories pending for extended scan - mu sync.RWMutex + phase string // Current scan phase + sessions []*AISessionSummary // Accumulated sessions + pendingDirs []string // Directories pending for extended scan + cursor string // Provider-specific discovery cursor + hasMore bool // Whether discovery has another batch + backgroundActive bool // Whether a provider background scan is queued or running + mu sync.RWMutex } // Global directory cache (projectDir -> cache entry) @@ -76,7 +79,7 @@ var ( // bgScanTask represents a background scan task type bgScanTask struct { projectPath string - scanType string // "claude" or "codex" + scanType string // "claude", "codex", or "pi" projectDir string // For claude: the specific project directory } @@ -112,6 +115,12 @@ func startBackgroundScanner() { zap.String("projectPath", task.projectPath), zap.Error(err)) } + case "pi": + if err := service.scanPiExtendedPhase(ctx, task.projectPath, task.projectDir); err != nil { + logger.Debug("background Pi scan failed", + zap.String("projectPath", task.projectPath), + zap.Error(err)) + } } } }() @@ -126,10 +135,14 @@ func NewAISessionService() *AISessionService { type ProjectAISessions struct { HasClaudeCode bool `json:"hasClaudeCode"` HasCodex bool `json:"hasCodex"` + HasPi bool `json:"hasPi"` ClaudeSessions []*AISessionSummary `json:"claudeSessions,omitempty"` CodexSessions []*AISessionSummary `json:"codexSessions,omitempty"` + PiSessions []*AISessionSummary `json:"piSessions"` ClaudeScanPhase string `json:"claudeScanPhase,omitempty"` // "recent", "extended", "complete" CodexScanPhase string `json:"codexScanPhase,omitempty"` // "recent", "extended", "complete" + PiScanPhase string `json:"piScanPhase"` + PiBeforeCursor string `json:"piBeforeCursor,omitempty"` } // AISessionSummary contains summary information about an AI session. @@ -162,8 +175,10 @@ func (s *AISessionService) GetProjectAISessions(ctx context.Context, projectPath result := &ProjectAISessions{ ClaudeSessions: make([]*AISessionSummary, 0), CodexSessions: make([]*AISessionSummary, 0), + PiSessions: make([]*AISessionSummary, 0), ClaudeScanPhase: ScanPhaseComplete, CodexScanPhase: ScanPhaseComplete, + PiScanPhase: ScanPhaseComplete, } // Normalize the project path @@ -205,11 +220,28 @@ func (s *AISessionService) GetProjectAISessions(ctx context.Context, projectPath zap.String("phase", codexPhase)) } + piSessions, piPhase, err := s.getPiSessionsPhased(ctx, projectPath) + if err != nil { + logger.Warn("failed to get Pi sessions", zap.Error(err), zap.String("path", projectPath)) + } else { + if piSessions == nil { + piSessions = []*AISessionSummary{} + } + result.PiSessions = piSessions + result.HasPi = len(piSessions) > 0 + result.PiScanPhase, result.PiBeforeCursor = s.currentPiScanProgress(projectPath, piPhase) + logger.Debug("Pi sessions found", + zap.Int("count", len(piSessions)), + zap.String("phase", piPhase)) + } + logger.Info("GetProjectAISessions returning", zap.Bool("hasClaudeCode", result.HasClaudeCode), zap.Bool("hasCodex", result.HasCodex), + zap.Bool("hasPi", result.HasPi), zap.Int("claudeCount", len(result.ClaudeSessions)), - zap.Int("codexCount", len(result.CodexSessions))) + zap.Int("codexCount", len(result.CodexSessions)), + zap.Int("piCount", len(result.PiSessions))) return result, nil } @@ -1846,24 +1878,36 @@ func (s *AISessionService) GetSessionConversationBySessionID(ctx context.Context var session tables.AISessionTable err := db.WithContext(ctx).Where("session_id = ?", sessionID).First(&session).Error if err == nil { - if session.Type == tables.AISessionTypeCodex { + switch session.Type { + case tables.AISessionTypeCodex: resolved, resolveErr := s.ResolveCodexSessionBySessionID(ctx, sessionID) if resolveErr != nil { return nil, resolveErr } return s.getConversationFromSession(ctx, *resolved) + case tables.AISessionTypePi: + resolved, resolveErr := s.ResolvePiSessionBySessionID(ctx, sessionID) + if resolveErr != nil { + return nil, resolveErr + } + return s.getConversationFromSession(ctx, *resolved) + default: + return s.getConversationFromSession(ctx, session) } - return s.getConversationFromSession(ctx, session) } if !errors.Is(err, gorm.ErrRecordNotFound) { return nil, err } resolved, resolveErr := s.ResolveCodexSessionBySessionID(ctx, sessionID) - if resolveErr != nil { - return nil, resolveErr + if resolveErr == nil { + return s.getConversationFromSession(ctx, *resolved) + } + resolvedPi, piErr := s.ResolvePiSessionBySessionID(ctx, sessionID) + if piErr == nil { + return s.getConversationFromSession(ctx, *resolvedPi) } - return s.getConversationFromSession(ctx, *resolved) + return nil, resolveErr } func (s *AISessionService) GetSessionConversationWindow( @@ -1896,24 +1940,36 @@ func (s *AISessionService) GetSessionConversationWindowBySessionID( var session tables.AISessionTable err := db.WithContext(ctx).Where("session_id = ?", sessionID).First(&session).Error if err == nil { - if session.Type == tables.AISessionTypeCodex { + switch session.Type { + case tables.AISessionTypeCodex: resolved, resolveErr := s.ResolveCodexSessionBySessionID(ctx, sessionID) if resolveErr != nil { return nil, resolveErr } return s.getConversationWindowFromSession(ctx, *resolved, beforeCursor, limit) + case tables.AISessionTypePi: + resolved, resolveErr := s.ResolvePiSessionBySessionID(ctx, sessionID) + if resolveErr != nil { + return nil, resolveErr + } + return s.getConversationWindowFromSession(ctx, *resolved, beforeCursor, limit) + default: + return s.getConversationWindowFromSession(ctx, session, beforeCursor, limit) } - return s.getConversationWindowFromSession(ctx, session, beforeCursor, limit) } if !errors.Is(err, gorm.ErrRecordNotFound) { return nil, err } resolved, resolveErr := s.ResolveCodexSessionBySessionID(ctx, sessionID) - if resolveErr != nil { - return nil, resolveErr + if resolveErr == nil { + return s.getConversationWindowFromSession(ctx, *resolved, beforeCursor, limit) + } + resolvedPi, piErr := s.ResolvePiSessionBySessionID(ctx, sessionID) + if piErr == nil { + return s.getConversationWindowFromSession(ctx, *resolvedPi, beforeCursor, limit) } - return s.getConversationWindowFromSession(ctx, *resolved, beforeCursor, limit) + return nil, resolveErr } // RefreshSessionAndGetConversation clears the cached session data in DB and re-parses the file. @@ -2008,6 +2064,11 @@ func (s *AISessionService) RefreshSessionAndGetConversation(ctx context.Context, } } } + case tables.AISessionTypePi: + messages, err = s.parsePiConversation(filePath) + if err == nil { + title = conversationTitleFromMessages(messages, "") + } default: return nil, errors.New("unknown session type") } @@ -2121,6 +2182,8 @@ func (s *AISessionService) getConversationFromSession(ctx context.Context, sessi messages, err = s.parseClaudeCodeConversation(session.FilePath) case tables.AISessionTypeCodex: messages, err = s.parseCodexConversation(session.FilePath) + case tables.AISessionTypePi: + messages, err = s.parsePiConversation(session.FilePath) default: return nil, errors.New("unknown session type") } @@ -2318,6 +2381,8 @@ func (s *AISessionService) getConversationWindowFromSession( messages, err = s.parseClaudeCodeConversation(session.FilePath) case tables.AISessionTypeCodex: messages, err = s.parseCodexConversation(session.FilePath) + case tables.AISessionTypePi: + messages, err = s.parsePiConversation(session.FilePath) default: return nil, errors.New("unknown session type") } diff --git a/service/project_agent_trust.go b/service/project_agent_trust.go new file mode 100644 index 00000000..bba5da24 --- /dev/null +++ b/service/project_agent_trust.go @@ -0,0 +1,252 @@ +package service + +import ( + "context" + "errors" + "fmt" + "os" + "path/filepath" + "runtime" + "strings" + "time" + + "code-kanban/model" + "code-kanban/model/tables" + + "gorm.io/gorm" +) + +type ProjectAgent string + +const ProjectAgentPi ProjectAgent = "pi" + +var ( + ErrUnsupportedProjectAgent = errors.New("unsupported project agent") + ErrProjectAgentTrustRequired = errors.New("project agent trust is required") + ErrProjectAgentPathNotAllowed = errors.New("path is not managed by the project") +) + +type ProjectAgentTrustStatus struct { + ProjectID string `json:"projectId"` + Agent string `json:"agent"` + ProjectPath string `json:"projectPath"` + TrustedPath string `json:"trustedPath,omitempty"` + Trusted bool `json:"trusted"` + TrustedAt *time.Time `json:"trustedAt,omitempty"` + RevokedAt *time.Time `json:"revokedAt,omitempty"` +} + +type ProjectAgentTrustService struct { + projectSvc *model.ProjectService + worktreeSvc *WorktreeService +} + +func NewProjectAgentTrustService() *ProjectAgentTrustService { + return &ProjectAgentTrustService{ + projectSvc: model.NewProjectService(), + worktreeSvc: NewWorktreeService(), + } +} + +func normalizeProjectAgent(agent ProjectAgent) (ProjectAgent, error) { + normalized := ProjectAgent(strings.ToLower(strings.TrimSpace(string(agent)))) + if normalized != ProjectAgentPi { + return "", ErrUnsupportedProjectAgent + } + return normalized, nil +} + +func CanonicalAgentTrustPath(value string) (string, error) { + value = strings.TrimSpace(value) + if value == "" { + return "", fmt.Errorf("project path is empty") + } + absolute, err := filepath.Abs(filepath.FromSlash(value)) + if err != nil { + return "", fmt.Errorf("resolve absolute project path: %w", err) + } + canonical := filepath.Clean(absolute) + if resolved, resolveErr := filepath.EvalSymlinks(canonical); resolveErr == nil && strings.TrimSpace(resolved) != "" { + canonical = filepath.Clean(resolved) + } + if runtime.GOOS == "windows" { + canonical = strings.ToLower(canonical) + } + return canonical, nil +} + +func (s *ProjectAgentTrustService) GetStatus( + ctx context.Context, + projectID string, + agent ProjectAgent, +) (ProjectAgentTrustStatus, error) { + if ctx == nil { + ctx = context.Background() + } + normalizedAgent, err := normalizeProjectAgent(agent) + if err != nil { + return ProjectAgentTrustStatus{}, err + } + project, err := s.projectSvc.GetProject(ctx, strings.TrimSpace(projectID)) + if err != nil { + return ProjectAgentTrustStatus{}, err + } + projectPath, err := CanonicalAgentTrustPath(project.Path) + if err != nil { + return ProjectAgentTrustStatus{}, err + } + status := ProjectAgentTrustStatus{ + ProjectID: project.Id, + Agent: string(normalizedAgent), + ProjectPath: projectPath, + } + + db := model.GetDB() + if db == nil { + return ProjectAgentTrustStatus{}, model.ErrDBNotInitialized + } + var record tables.ProjectAgentTrustTable + err = db.WithContext(ctx). + Where("project_id = ? AND agent = ?", project.Id, string(normalizedAgent)). + First(&record).Error + if errors.Is(err, gorm.ErrRecordNotFound) { + return status, nil + } + if err != nil { + return ProjectAgentTrustStatus{}, err + } + status.TrustedPath = record.TrustedPath + trustedAt := record.TrustedAt + status.TrustedAt = &trustedAt + status.RevokedAt = record.RevokedAt + trustedPath, pathErr := CanonicalAgentTrustPath(record.TrustedPath) + status.Trusted = pathErr == nil && record.RevokedAt == nil && trustedPath == projectPath + return status, nil +} + +func (s *ProjectAgentTrustService) Trust( + ctx context.Context, + projectID string, + agent ProjectAgent, +) (ProjectAgentTrustStatus, error) { + if ctx == nil { + ctx = context.Background() + } + normalizedAgent, err := normalizeProjectAgent(agent) + if err != nil { + return ProjectAgentTrustStatus{}, err + } + project, err := s.projectSvc.GetProject(ctx, strings.TrimSpace(projectID)) + if err != nil { + return ProjectAgentTrustStatus{}, err + } + info, err := os.Stat(project.Path) + if err != nil || !info.IsDir() { + if err == nil { + err = fmt.Errorf("project path is not a directory") + } + return ProjectAgentTrustStatus{}, fmt.Errorf("project path is unavailable: %w", err) + } + trustedPath, err := CanonicalAgentTrustPath(project.Path) + if err != nil { + return ProjectAgentTrustStatus{}, err + } + db := model.GetDB() + if db == nil { + return ProjectAgentTrustStatus{}, model.ErrDBNotInitialized + } + now := time.Now() + var record tables.ProjectAgentTrustTable + err = db.WithContext(ctx).Unscoped(). + Where("project_id = ? AND agent = ?", project.Id, string(normalizedAgent)). + First(&record).Error + switch { + case errors.Is(err, gorm.ErrRecordNotFound): + record = tables.ProjectAgentTrustTable{ + ProjectID: project.Id, + Agent: string(normalizedAgent), + TrustedPath: trustedPath, + TrustedAt: now, + } + if err := db.WithContext(ctx).Create(&record).Error; err != nil { + return ProjectAgentTrustStatus{}, err + } + case err != nil: + return ProjectAgentTrustStatus{}, err + default: + if err := db.WithContext(ctx).Unscoped().Model(&record).Updates(map[string]any{ + "trusted_path": trustedPath, + "trusted_at": now, + "revoked_at": nil, + "deleted_at": nil, + }).Error; err != nil { + return ProjectAgentTrustStatus{}, err + } + } + return s.GetStatus(ctx, project.Id, normalizedAgent) +} + +func (s *ProjectAgentTrustService) Revoke( + ctx context.Context, + projectID string, + agent ProjectAgent, +) (ProjectAgentTrustStatus, error) { + if ctx == nil { + ctx = context.Background() + } + normalizedAgent, err := normalizeProjectAgent(agent) + if err != nil { + return ProjectAgentTrustStatus{}, err + } + project, err := s.projectSvc.GetProject(ctx, strings.TrimSpace(projectID)) + if err != nil { + return ProjectAgentTrustStatus{}, err + } + db := model.GetDB() + if db == nil { + return ProjectAgentTrustStatus{}, model.ErrDBNotInitialized + } + now := time.Now() + if err := db.WithContext(ctx).Model(&tables.ProjectAgentTrustTable{}). + Where("project_id = ? AND agent = ?", project.Id, string(normalizedAgent)). + Updates(map[string]any{"revoked_at": now, "updated_at": now}).Error; err != nil { + return ProjectAgentTrustStatus{}, err + } + return s.GetStatus(ctx, project.Id, normalizedAgent) +} + +func (s *ProjectAgentTrustService) EnsureTrustedPath( + ctx context.Context, + projectID string, + agent ProjectAgent, + cwd string, +) error { + status, err := s.GetStatus(ctx, projectID, agent) + if err != nil { + return err + } + if !status.Trusted { + return ErrProjectAgentTrustRequired + } + candidate, err := CanonicalAgentTrustPath(cwd) + if err != nil { + return ErrProjectAgentPathNotAllowed + } + if candidate == status.ProjectPath { + return nil + } + worktrees, err := s.worktreeSvc.ListWorktrees(ctx, status.ProjectID) + if err != nil { + return err + } + for _, worktree := range worktrees { + if worktree == nil || worktree.ProjectId != status.ProjectID { + continue + } + worktreePath, pathErr := CanonicalAgentTrustPath(worktree.Path) + if pathErr == nil && candidate == worktreePath { + return nil + } + } + return ErrProjectAgentPathNotAllowed +} diff --git a/service/project_agent_trust_test.go b/service/project_agent_trust_test.go new file mode 100644 index 00000000..1f11ae80 --- /dev/null +++ b/service/project_agent_trust_test.go @@ -0,0 +1,144 @@ +package service + +import ( + "context" + "errors" + "path/filepath" + "runtime" + "strings" + "testing" + + "code-kanban/model" + "code-kanban/model/tables" +) + +func seedAgentTrustProject(t *testing.T, path string) *tables.ProjectTable { + t.Helper() + project := &tables.ProjectTable{Name: "Agent Trust Test", Path: path} + project.Init() + if err := model.GetDB().Create(project).Error; err != nil { + t.Fatalf("seed project: %v", err) + } + return project +} + +func TestProjectAgentTrustLifecycleAndPathChange(t *testing.T) { + cleanup := initTestDB(t) + defer cleanup() + + projectPath := t.TempDir() + project := seedAgentTrustProject(t, projectPath) + svc := NewProjectAgentTrustService() + ctx := context.Background() + + status, err := svc.GetStatus(ctx, project.ID, ProjectAgentPi) + if err != nil || status.Trusted { + t.Fatalf("initial status = %#v, %v", status, err) + } + status, err = svc.Trust(ctx, project.ID, ProjectAgentPi) + if err != nil || !status.Trusted || status.TrustedAt == nil { + t.Fatalf("trusted status = %#v, %v", status, err) + } + expectedPath, err := CanonicalAgentTrustPath(projectPath) + if err != nil || status.TrustedPath != expectedPath { + t.Fatalf("trusted path = %q, want %q (%v)", status.TrustedPath, expectedPath, err) + } + if err := svc.EnsureTrustedPath(ctx, project.ID, ProjectAgentPi, projectPath); err != nil { + t.Fatalf("EnsureTrustedPath(project) returned error: %v", err) + } + + movedPath := t.TempDir() + if err := model.GetDB().Model(&tables.ProjectTable{}). + Where("id = ?", project.ID). + Update("path", movedPath).Error; err != nil { + t.Fatalf("move project path in database: %v", err) + } + status, err = svc.GetStatus(ctx, project.ID, ProjectAgentPi) + if err != nil || status.Trusted { + t.Fatalf("status after path change = %#v, %v", status, err) + } + if err := svc.EnsureTrustedPath(ctx, project.ID, ProjectAgentPi, movedPath); !errors.Is(err, ErrProjectAgentTrustRequired) { + t.Fatalf("EnsureTrustedPath after path change = %v, want trust required", err) + } + + status, err = svc.Trust(ctx, project.ID, ProjectAgentPi) + if err != nil || !status.Trusted { + t.Fatalf("re-trust after path change = %#v, %v", status, err) + } + status, err = svc.Revoke(ctx, project.ID, ProjectAgentPi) + if err != nil || status.Trusted || status.RevokedAt == nil { + t.Fatalf("revoked status = %#v, %v", status, err) + } +} + +func TestProjectAgentTrustAllowsOnlyManagedWorktrees(t *testing.T) { + cleanup := initTestDB(t) + defer cleanup() + + projectPath := t.TempDir() + worktreePath := t.TempDir() + unrelatedPath := t.TempDir() + project := seedAgentTrustProject(t, projectPath) + worktree := &tables.WorktreeTable{ + ProjectID: project.ID, + BranchName: "feature", + Path: worktreePath, + } + worktree.Init() + if err := model.GetDB().Create(worktree).Error; err != nil { + t.Fatalf("seed worktree: %v", err) + } + + svc := NewProjectAgentTrustService() + if _, err := svc.Trust(context.Background(), project.ID, ProjectAgentPi); err != nil { + t.Fatalf("Trust returned error: %v", err) + } + if err := svc.EnsureTrustedPath(context.Background(), project.ID, ProjectAgentPi, worktreePath); err != nil { + t.Fatalf("managed worktree rejected: %v", err) + } + if err := svc.EnsureTrustedPath(context.Background(), project.ID, ProjectAgentPi, unrelatedPath); !errors.Is(err, ErrProjectAgentPathNotAllowed) { + t.Fatalf("unrelated path result = %v, want path not allowed", err) + } +} + +func TestDeleteProjectHardDeletesAgentTrust(t *testing.T) { + cleanup := initTestDB(t) + defer cleanup() + + project := seedAgentTrustProject(t, t.TempDir()) + svc := NewProjectAgentTrustService() + if _, err := svc.Trust(context.Background(), project.ID, ProjectAgentPi); err != nil { + t.Fatalf("Trust returned error: %v", err) + } + if err := model.NewProjectService().DeleteProject(context.Background(), project.ID); err != nil { + t.Fatalf("DeleteProject returned error: %v", err) + } + var count int64 + if err := model.GetDB().Unscoped().Model(&tables.ProjectAgentTrustTable{}). + Where("project_id = ?", project.ID). + Count(&count).Error; err != nil { + t.Fatalf("count trust rows: %v", err) + } + if count != 0 { + t.Fatalf("trust row count = %d, want 0", count) + } +} + +func TestCanonicalAgentTrustPathNormalizesRelativeSegments(t *testing.T) { + root := t.TempDir() + value := filepath.Join(root, "child", "..") + got, err := CanonicalAgentTrustPath(value) + if err != nil { + t.Fatalf("CanonicalAgentTrustPath returned error: %v", err) + } + want, err := CanonicalAgentTrustPath(root) + if err != nil || got != want { + t.Fatalf("canonical path = %q, want %q (%v)", got, want, err) + } + if runtime.GOOS == "windows" { + upper, err := CanonicalAgentTrustPath(strings.ToUpper(root)) + if err != nil || upper != want { + t.Fatalf("case-normalized path = %q, want %q (%v)", upper, want, err) + } + } +} diff --git a/service/websession/codex_app_server.go b/service/websession/codex_app_server.go index 009ef0b4..ebd2a5fe 100644 --- a/service/websession/codex_app_server.go +++ b/service/websession/codex_app_server.go @@ -99,17 +99,21 @@ type pendingServerRequest struct { RawID json.RawMessage // Claude's stream-json control protocol has its own request id in addition // to the tool_use id stored in ItemID. Codex continues to use RawID. - ControlRequestID string - Kind pendingServerRequestKind - ThreadID string - TurnID string - ItemID string - Prompt string - Command string - Input any - RequestedAt *time.Time - Questions []toolRequestQuestion - Permissions map[string]any + ControlRequestID string + Kind pendingServerRequestKind + ThreadID string + TurnID string + ItemID string + Prompt string + Command string + Input any + RequestedAt *time.Time + Questions []toolRequestQuestion + Permissions map[string]any + PiRuntime *piSessionRuntime + PiRequestID string + PiMethod string + PiResponseGeneration uint64 } func (r *pendingServerRequest) clone() *pendingServerRequest { @@ -117,16 +121,20 @@ func (r *pendingServerRequest) clone() *pendingServerRequest { return nil } clone := &pendingServerRequest{ - RawID: append(json.RawMessage(nil), r.RawID...), - ControlRequestID: r.ControlRequestID, - Kind: r.Kind, - ThreadID: r.ThreadID, - TurnID: r.TurnID, - ItemID: r.ItemID, - Prompt: r.Prompt, - Command: r.Command, - Input: r.Input, - Permissions: nil, + RawID: append(json.RawMessage(nil), r.RawID...), + ControlRequestID: r.ControlRequestID, + Kind: r.Kind, + ThreadID: r.ThreadID, + TurnID: r.TurnID, + ItemID: r.ItemID, + Prompt: r.Prompt, + Command: r.Command, + Input: r.Input, + Permissions: nil, + PiRuntime: r.PiRuntime, + PiRequestID: r.PiRequestID, + PiMethod: r.PiMethod, + PiResponseGeneration: r.PiResponseGeneration, } if r.RequestedAt != nil { requestedAt := *r.RequestedAt diff --git a/service/websession/codex_context.go b/service/websession/codex_context.go index 9bab4e8c..2118d925 100644 --- a/service/websession/codex_context.go +++ b/service/websession/codex_context.go @@ -66,15 +66,53 @@ type CodexModelInfo struct { SupportedReasoningEfforts []ReasoningEffort `json:"supportedReasoningEfforts"` } -type CodexRuntimeConfig struct { - Model string `json:"model,omitempty"` - ContextWindowTokens int64 `json:"contextWindowTokens"` - CompactLimitTokens int64 `json:"compactLimitTokens"` - Source ContextWindowSource `json:"source"` - Models []CodexModelInfo `json:"models"` - HasCodex bool `json:"hasCodex"` - HasClaudeCode bool `json:"hasClaudeCode"` - CodexVersion *string `json:"codexVersion,omitempty"` +type AgentPermissionModeCapability struct { + ID string `json:"id"` + Available bool `json:"available"` +} + +type AgentCapability struct { + Installed bool `json:"installed"` + Version *string `json:"version,omitempty"` + SupportsWebSession bool `json:"supportsWebSession"` + SupportsTree bool `json:"supportsTree"` + SupportsImages bool `json:"supportsImages"` + SupportsCompaction bool `json:"supportsCompaction"` + SupportsSteer bool `json:"supportsSteer"` + SupportsFollowUp bool `json:"supportsFollowUp"` + SupportsGoal bool `json:"supportsGoal"` + SupportsSubAgentRegistry bool `json:"supportsSubAgentRegistry"` + PermissionModes []AgentPermissionModeCapability `json:"permissionModes"` +} + +type PiModelInfo struct { + Provider string `json:"provider"` + ID string `json:"id"` + Name string `json:"name"` + Reasoning bool `json:"reasoning"` + Input []string `json:"input"` + ContextWindow int64 `json:"contextWindow"` + MaxTokens int64 `json:"maxTokens"` +} + +type WebSessionRuntimeConfig struct { + Agents map[Agent]AgentCapability `json:"agents"` + Model string `json:"model,omitempty"` + ContextWindowTokens int64 `json:"contextWindowTokens"` + CompactLimitTokens int64 `json:"compactLimitTokens"` + Source ContextWindowSource `json:"source"` + Models []CodexModelInfo `json:"models"` + PiModels []PiModelInfo `json:"piModels"` + // Legacy top-level fields are retained for one compatibility cycle. + HasCodex bool `json:"hasCodex"` + HasClaudeCode bool `json:"hasClaudeCode"` + CodexVersion *string `json:"codexVersion,omitempty"` + HasPi bool `json:"hasPi"` + PiVersion *string `json:"piVersion,omitempty"` + SupportsPiWebSession bool `json:"supportsPiWebSession"` + PiRPCCompatible bool `json:"piRpcCompatible"` + PiMinVersion string `json:"piMinVersion"` + PiDiagnostics string `json:"piDiagnostics,omitempty"` // SupportsWebSession reports whether ordinary Codex web sessions can run. SupportsWebSession bool `json:"supportsWebSession"` WebSessionMinVersion string `json:"webSessionMinCodexVersion"` @@ -85,6 +123,9 @@ type CodexRuntimeConfig struct { GoalModeMinVersion string `json:"goalModeMinCodexVersion"` } +// CodexRuntimeConfig remains an alias while callers migrate to the provider-neutral name. +type CodexRuntimeConfig = WebSessionRuntimeConfig + type CodexSkillSource string const ( @@ -125,6 +166,16 @@ func (m *Manager) decorateSessionSummary(summary *SessionSummary) { summary.ContextWindowSource = ContextWindowSourceUnavailable return } + if normalizeAgent(summary.Agent) == AgentPi { + if summary.ContextWindowTokens != nil && + *summary.ContextWindowTokens > 0 && + summary.ContextWindowSource == ContextWindowSourceSessionUsage { + return + } + summary.ContextWindowTokens = nil + summary.ContextWindowSource = ContextWindowSourceUnavailable + return + } if normalizeAgent(summary.Agent) != AgentCodex { summary.ContextWindowTokens = nil summary.ContextWindowSource = ContextWindowSourceUnavailable @@ -158,6 +209,7 @@ func (m *Manager) GetCodexRuntimeConfig() CodexRuntimeConfig { SupportsGoalMode: false, GoalModeMinVersion: goalModeMinCodexVersion.String(), } + defaultConfig.Agents = runtimeAgentCapabilities(defaultConfig) if m == nil { return defaultConfig } @@ -211,9 +263,78 @@ func (m *Manager) applyCodexRuntimeCapabilities(config CodexRuntimeConfig) Codex if config.Models == nil { config.Models = []CodexModelInfo{} } + config.Agents = runtimeAgentCapabilities(config) + return config +} + +func availablePermissionModes(unrestricted, approval, sandbox bool) []AgentPermissionModeCapability { + return []AgentPermissionModeCapability{ + {ID: "unrestricted", Available: unrestricted}, + {ID: "approval", Available: approval}, + {ID: "sandbox", Available: sandbox}, + } +} + +func runtimeAgentCapabilities(config WebSessionRuntimeConfig) map[Agent]AgentCapability { + return map[Agent]AgentCapability{ + AgentClaude: { + Installed: config.HasClaudeCode, + SupportsWebSession: config.HasClaudeCode, + SupportsImages: true, + SupportsCompaction: true, + SupportsSteer: true, + SupportsFollowUp: true, + SupportsSubAgentRegistry: false, + PermissionModes: availablePermissionModes(true, true, false), + }, + AgentCodex: { + Installed: config.HasCodex, + Version: config.CodexVersion, + SupportsWebSession: config.SupportsWebSession, + SupportsImages: true, + SupportsCompaction: true, + SupportsSteer: true, + SupportsFollowUp: true, + SupportsGoal: config.SupportsGoalMode, + SupportsSubAgentRegistry: config.SupportsMultiAgentV2, + PermissionModes: availablePermissionModes(true, true, true), + }, + AgentPi: { + Installed: config.HasPi, + Version: config.PiVersion, + SupportsWebSession: config.SupportsPiWebSession, + SupportsTree: config.SupportsPiWebSession, + SupportsImages: true, + SupportsCompaction: config.SupportsPiWebSession, + SupportsSteer: config.SupportsPiWebSession, + SupportsFollowUp: config.SupportsPiWebSession, + SupportsGoal: false, + SupportsSubAgentRegistry: false, + PermissionModes: availablePermissionModes(true, false, false), + }, + } +} + +func (m *Manager) GetWebSessionRuntimeConfig() WebSessionRuntimeConfig { + return m.applyPiRuntimeCapabilities(m.GetCodexRuntimeConfig()) +} + +func (m *Manager) SupportsPiSessionTree() bool { + if m == nil { + return false + } + return m.GetWebSessionRuntimeConfig().Agents[AgentPi].SupportsTree +} + +func (m *Manager) GetWebSessionRuntimeConfigWithModels() WebSessionRuntimeConfig { + config := m.GetWebSessionRuntimeConfig() + if config.HasCodex { + config.Models = m.getCodexModelCatalog() + } return config } +// GetCodexRuntimeConfigWithModels is kept for callers using the previous API name. func (m *Manager) GetCodexRuntimeConfigWithModels() CodexRuntimeConfig { config := m.GetCodexRuntimeConfig() if config.HasCodex { @@ -467,20 +588,41 @@ func splitCommandParts(command string) []string { if trimmed == "" { return nil } - if trimmed[0] != '"' && trimmed[0] != '\'' { - return strings.Fields(trimmed) - } - quote := trimmed[0] - end := strings.IndexByte(trimmed[1:], quote) - if end < 0 { - return strings.Fields(trimmed) + + parts := make([]string, 0, 4) + var current strings.Builder + var quote byte + flush := func() { + if current.Len() == 0 { + return + } + parts = append(parts, current.String()) + current.Reset() + } + for i := 0; i < len(trimmed); i++ { + char := trimmed[i] + if quote != 0 { + if char == quote { + quote = 0 + } else { + current.WriteByte(char) + } + continue + } + switch char { + case '"', '\'': + quote = char + case ' ', '\t', '\r', '\n': + flush() + default: + current.WriteByte(char) + } } - commandPart := trimmed[1 : end+1] - remainder := strings.TrimSpace(trimmed[end+2:]) - if remainder == "" { - return []string{commandPart} + if quote != 0 { + return nil } - return append([]string{commandPart}, strings.Fields(remainder)...) + flush() + return parts } func codexVersionAtLeast(raw string, min *semver.Version) bool { diff --git a/service/websession/live_cache.go b/service/websession/live_cache.go index a3389682..94ca0a9e 100644 --- a/service/websession/live_cache.go +++ b/service/websession/live_cache.go @@ -404,11 +404,15 @@ func (m *Manager) applyEventToHistoryCacheDB( next.SourceTurnID = sourceTurnID next.Kind = "assistant" next.ItemType = "agent_message" + if text, authoritative := payload["txt"].(string); authoritative { + next.Text = text + } next.ObservedAt = ptr(event.Timestamp) if next.Timestamp == nil { next.Timestamp = ptr(event.Timestamp) } next.Done = true + next.Payload = payload }) if err != nil { return nil, err diff --git a/service/websession/manager.go b/service/websession/manager.go index e77175a2..33f5874e 100644 --- a/service/websession/manager.go +++ b/service/websession/manager.go @@ -42,6 +42,7 @@ const ( recoveryMessageProcessRestart = "Session runtime was interrupted because the app restarted. Send a new message to continue." errCodexNotInstalled = "Codex is not installed. Install Codex before sending messages in this session." errClaudeCodeNotInstalled = "Claude Code is not installed. Install Claude Code before sending messages in this session." + errPiWebSessionUnavailable = "Pi Web Sessions require a compatible Pi RPC runtime." ) var ( @@ -109,6 +110,8 @@ type Config struct { CCRPath string CCRConfigPath string CodexPath string + PiPath string + PiRuntimeIdleTTL time.Duration DefaultCodexModel func() string DefaultCodexReasoningEffort func() ReasoningEffort DefaultCodexPermissionLevel func() string @@ -117,12 +120,13 @@ type Config struct { } type Manager struct { - cfg Config - logger *zap.Logger - store *store - projectSvc *model.ProjectService - worktreeSvc *service.WorktreeService - aiSessionSvc *service.AISessionService + cfg Config + logger *zap.Logger + store *store + projectSvc *model.ProjectService + worktreeSvc *service.WorktreeService + aiSessionSvc *service.AISessionService + agentTrustSvc *service.ProjectAgentTrustService eventStatesMu sync.Mutex eventStates map[string]*sessionEventState @@ -145,9 +149,15 @@ type Manager struct { sessionDispatchLocks [64]sync.Mutex revisionBroadcastLocks [64]sync.Mutex pendingInputs map[string][]PendingInput + piNativeQueuedInputs map[string][]PendingInput pendingProcessing map[string]bool pendingDirty map[string]bool codexContextWindow codexContextWindowResolver + piProbeMu sync.Mutex + piProbe piRuntimeProbeCache + piRuntimeMu sync.Mutex + piRuntimeTerminators map[string]piRuntimeTerminator + piRuntimes map[string]*piSessionRuntime claudeHookOnce sync.Once claudeHookBaseURL string claudeHookToken string @@ -189,61 +199,68 @@ type wsConn interface { } type activeRun struct { - sessionID string - agent Agent - backend SessionBackend - runID string - fromAutoRetry bool - hiddenBootstrap bool - bootstrapGoalObjective string - bootstrapGoalState GoalStatus - bootstrapResult chan error - bootstrapOnce sync.Once - assistantMessageID string - currentToolMessage string - lastError string - lastErrorCode string - transportRetrySeen bool - transportRemoteURL string - cmd *exec.Cmd - cancel context.CancelFunc - done chan struct{} - mu sync.Mutex - stdin io.WriteCloser - recentRuntimeLines []string - pendingApproval string - pendingServerReq *pendingServerRequest - inputResponsePending bool - app *codexAppServerClient - codexThreadID string - codexTurnID string - assistantDeltaSeen map[string]bool - assistantMessagePhases map[string]string - assistantMessageText map[string]bool - completedReply bool - claudeResumeOnly bool - claudeStdioControl bool - claudeNativeSessionID string - claudeResolvedModel string - claudeCwd string - deferredUserInput bool - completedPlanTool bool - claudeCompactionToolID string - claudeCompactionStarted time.Time - commandGroupID string - commandGroupKind string - commandGroupKey string - commandGroupFirst int64 - commandGroupCount int - commandGroupTools map[string]struct{} - abortPayload map[string]any - activeCalls map[string]trackedActiveCall - activeCallPausedAt *time.Time - activeCallTimer *time.Timer - activeCallInFlight bool - codexCollaboration codexCollaborationTracker - codexRolloutMonitor *codexRolloutMonitor - syncSourceAfterRun bool + sessionID string + agent Agent + backend SessionBackend + runID string + fromAutoRetry bool + hiddenBootstrap bool + bootstrapGoalObjective string + bootstrapGoalState GoalStatus + bootstrapResult chan error + bootstrapOnce sync.Once + assistantMessageID string + currentToolMessage string + lastError string + lastErrorCode string + transportRetrySeen bool + transportRemoteURL string + cmd *exec.Cmd + cancel context.CancelFunc + done chan struct{} + mu sync.Mutex + stdin io.WriteCloser + recentRuntimeLines []string + pendingApproval string + pendingServerReq *pendingServerRequest + inputResponsePending bool + piResponseGeneration uint64 + piResponsePending uint64 + piResponseHistoryDone chan struct{} + piResponseHistoryComplete uint64 + piResponseHistoryErr error + piResponseRequest *pendingServerRequest + app *codexAppServerClient + codexThreadID string + codexTurnID string + assistantDeltaSeen map[string]bool + assistantMessagePhases map[string]string + assistantMessageText map[string]bool + completedReply bool + claudeResumeOnly bool + claudeStdioControl bool + claudeNativeSessionID string + claudeResolvedModel string + claudeCwd string + deferredUserInput bool + completedPlanTool bool + claudeCompactionToolID string + claudeCompactionStarted time.Time + commandGroupID string + commandGroupKind string + commandGroupKey string + commandGroupFirst int64 + commandGroupCount int + commandGroupTools map[string]struct{} + abortPayload map[string]any + activeCalls map[string]trackedActiveCall + activeCallPausedAt *time.Time + activeCallTimer *time.Timer + activeCallInFlight bool + codexCollaboration codexCollaborationTracker + codexRolloutMonitor *codexRolloutMonitor + syncSourceAfterRun bool + piCompaction bool } type attachmentMeta struct { @@ -399,6 +416,12 @@ func NewManager(cfg Config, logger *zap.Logger) (*Manager, error) { if cfg.CodexPath == "" { cfg.CodexPath = getenvDefault("CODEX_PATH", "codex") } + if cfg.PiPath == "" { + cfg.PiPath = getenvDefault("PI_PATH", "pi") + } + if cfg.PiRuntimeIdleTTL <= 0 { + cfg.PiRuntimeIdleTTL = 2 * time.Minute + } if logger == nil { logger = utils.Logger() } @@ -415,6 +438,9 @@ func NewManager(cfg Config, logger *zap.Logger) (*Manager, error) { projectSvc: model.NewProjectService(), worktreeSvc: service.NewWorktreeService(), aiSessionSvc: service.NewAISessionService(), + agentTrustSvc: service.NewProjectAgentTrustService(), + piRuntimeTerminators: make(map[string]piRuntimeTerminator), + piRuntimes: make(map[string]*piSessionRuntime), runs: make(map[string]*activeRun), clients: make(map[*client]struct{}), autoRetryTimers: make(map[string]*time.Timer), @@ -424,6 +450,7 @@ func NewManager(cfg Config, logger *zap.Logger) (*Manager, error) { pendingInputTimerDeadlines: make(map[string]time.Time), pendingSteerDelay: defaultPendingSteerDelay, pendingInputs: make(map[string][]PendingInput), + piNativeQueuedInputs: make(map[string][]PendingInput), pendingProcessing: make(map[string]bool), pendingDirty: make(map[string]bool), eventStates: make(map[string]*sessionEventState), @@ -843,7 +870,10 @@ func (m *Manager) ListArchivedSessions( } func (m *Manager) CreateSession(ctx context.Context, params CreateParams) (SessionSummary, error) { - agent := normalizeAgent(params.Agent) + agent, err := validateAgent(params.Agent) + if err != nil { + return SessionSummary{}, err + } permissionLevel := m.resolveSessionPermissionLevel(agent, params.PermissionLevel) if err := validateWebSessionPermissionLevel(agent, permissionLevel); err != nil { return SessionSummary{}, err @@ -855,7 +885,7 @@ func (m *Manager) CreateSession(ctx context.Context, params CreateParams) (Sessi title := strings.TrimSpace(params.Title) if title == "" { - title = defaultTitle(params.Agent, project.Name) + title = defaultTitle(agent, project.Name) } orderIndex, err := m.getNextSessionOrderIndex(ctx, project.Id) @@ -865,6 +895,16 @@ func (m *Manager) CreateSession(ctx context.Context, params CreateParams) (Sessi modelName := m.resolveSessionModel(agent, params.Model) reasoningEffort := m.resolveSessionReasoningEffort(agent, modelName, params.ReasoningEffort) + if agent == AgentPi { + if modelName != "" { + if _, _, err := splitPiModel(modelName); err != nil { + return SessionSummary{}, err + } + } + if err := validatePiReasoningEffort(reasoningEffort); err != nil { + return SessionSummary{}, err + } + } now := time.Now() record := tables.WebSessionTable{ ProjectID: project.Id, @@ -893,7 +933,7 @@ func (m *Manager) CreateSession(ctx context.Context, params CreateParams) (Sessi ActivityAt: now, StatusUpdatedAt: &now, AssistantStateUpdatedAt: nil, - SourceKind: defaultSourceKind(normalizeAgent(params.Agent)), + SourceKind: defaultSourceKind(agent), SyncState: string(SyncStateMissing), LastSyncMode: "", SourceCreatedAt: nil, @@ -1085,9 +1125,10 @@ func preferImportedCodexSession(current, candidate tables.WebSessionTable) table return current } -func (m *Manager) existingImportedCodexSessionsByNativeID( +func (m *Manager) existingImportedSessionsByNativeID( ctx context.Context, projectID string, + agent Agent, sessionIDs []string, ) (map[string]tables.WebSessionTable, error) { normalized := make([]string, 0, len(sessionIDs)) @@ -1114,7 +1155,7 @@ func (m *Manager) existingImportedCodexSessionsByNativeID( Where( "project_id = ? AND agent = ? AND native_session_id IN ?", projectID, - string(AgentCodex), + string(agent), normalized, ). Find(&existing).Error; err != nil { @@ -1134,6 +1175,14 @@ func (m *Manager) existingImportedCodexSessionsByNativeID( return result, nil } +func (m *Manager) existingImportedCodexSessionsByNativeID( + ctx context.Context, + projectID string, + sessionIDs []string, +) (map[string]tables.WebSessionTable, error) { + return m.existingImportedSessionsByNativeID(ctx, projectID, AgentCodex, sessionIDs) +} + func sortImportSourceItems(items []ImportSourceSummary) { sort.Slice(items, func(i, j int) bool { left := items[i] @@ -1200,6 +1249,8 @@ func (m *Manager) buildImportSourceItemFromThread( } item := ImportSourceSummary{ + Agent: AgentCodex, + Importable: true, AISessionID: aiSessionID, SessionID: strings.TrimSpace(thread.ID), Model: model, @@ -1220,9 +1271,13 @@ func (m *Manager) buildImportSourceItemFromThread( func (m *Manager) buildImportSourceItemFromAISession( source *service.AISessionSummary, + agent Agent, + importable bool, existingByNativeID map[string]tables.WebSessionTable, ) ImportSourceSummary { item := ImportSourceSummary{ + Agent: agent, + Importable: importable, AISessionID: strings.TrimSpace(source.ID), SessionID: strings.TrimSpace(source.SessionID), Model: strings.TrimSpace(source.Model), @@ -1355,7 +1410,7 @@ func (m *Manager) ListCodexImportSources( if source == nil { continue } - items = append(items, m.buildImportSourceItemFromAISession(source, existingByNativeID)) + items = append(items, m.buildImportSourceItemFromAISession(source, AgentCodex, true, existingByNativeID)) } sortImportSourceItems(items) return ImportSourceList{ @@ -1364,6 +1419,68 @@ func (m *Manager) ListCodexImportSources( }, nil } +func (m *Manager) ListImportSources( + ctx context.Context, + projectID string, +) (ImportSourceList, error) { + codexList, err := m.ListCodexImportSources(ctx, projectID) + if err != nil { + return ImportSourceList{}, err + } + if m.aiSessionSvc == nil { + return codexList, nil + } + project, err := m.projectSvc.GetProject(ctx, projectID) + if err != nil { + return ImportSourceList{}, err + } + aiSessions, err := m.aiSessionSvc.GetProjectAISessions(ctx, project.Path) + if err != nil { + // Codex thread/list is independent of the filesystem index. Preserve the + // existing result when optional Pi discovery cannot read its session root. + return codexList, nil + } + + piSessionIDs := make([]string, 0, len(aiSessions.PiSessions)) + for _, source := range aiSessions.PiSessions { + if source != nil && strings.TrimSpace(source.SessionID) != "" { + piSessionIDs = append(piSessionIDs, strings.TrimSpace(source.SessionID)) + } + } + existingPi, err := m.existingImportedSessionsByNativeID(ctx, project.Id, AgentPi, piSessionIDs) + if err != nil { + return ImportSourceList{}, err + } + items := append([]ImportSourceSummary(nil), codexList.Items...) + piImportable := m.GetWebSessionRuntimeConfig().SupportsPiWebSession + for _, source := range aiSessions.PiSessions { + if source == nil { + continue + } + items = append(items, m.buildImportSourceItemFromAISession(source, AgentPi, piImportable, existingPi)) + } + sortImportSourceItems(items) + + return ImportSourceList{ + Items: items, + ScanPhase: aggregateImportScanPhase(codexList.ScanPhase, aiSessions.PiScanPhase), + BeforeCursor: strings.TrimSpace(aiSessions.PiBeforeCursor), + }, nil +} + +func aggregateImportScanPhase(phases ...string) string { + result := "complete" + for _, phase := range phases { + switch strings.ToLower(strings.TrimSpace(phase)) { + case "extended": + return "extended" + case "recent": + result = "recent" + } + } + return result +} + func (m *Manager) importCodexSessionResolved( ctx context.Context, project *model.Project, @@ -1477,6 +1594,152 @@ func (m *Manager) ImportCodexSessionBySessionID( return m.importCodexSessionResolved(ctx, project, source, mode) } +func (m *Manager) importPiSessionResolved( + ctx context.Context, + project *model.Project, + source *tables.AISessionTable, +) (ImportResult, error) { + if !m.GetWebSessionRuntimeConfig().SupportsPiWebSession { + return ImportResult{}, errors.New(errPiWebSessionUnavailable) + } + if source == nil { + return ImportResult{}, gorm.ErrRecordNotFound + } + nativeID := strings.TrimSpace(source.SessionID) + threadPath := strings.TrimSpace(source.FilePath) + if nativeID == "" || threadPath == "" { + return ImportResult{}, errors.New("Pi session identity is incomplete") + } + if model.NormalizePathCase(source.ProjectPath) != model.NormalizePathCase(project.Path) { + return ImportResult{}, errors.New("Pi session does not belong to the current project") + } + if err := m.EnsureProjectPiTrust(ctx, project.Id, project.Path); err != nil { + return ImportResult{}, err + } + identity := tables.WebSessionTable{ + Cwd: project.Path, + NativeSessionID: &nativeID, + ThreadPath: &threadPath, + } + if err := validatePiRuntimeState(identity, piRPCState{SessionID: nativeID, SessionFile: threadPath}); err != nil { + return ImportResult{}, err + } + + var records []tables.WebSessionTable + if err := model.GetDB().WithContext(ctx). + Where("project_id = ? AND agent = ? AND native_session_id = ?", project.Id, string(AgentPi), nativeID). + Order("updated_at DESC"). + Find(&records).Error; err != nil { + return ImportResult{}, err + } + if len(records) > 0 { + record := records[0] + for _, candidate := range records[1:] { + record = preferImportedCodexSession(record, candidate) + } + updates := map[string]any{ + "cwd": filepath.Clean(project.Path), + "thread_path": filepath.Clean(threadPath), + "native_session_id": nativeID, + "source_kind": defaultSourceKind(AgentPi), + "source_created_at": importedCodexSourceCreatedAt(*source), + "source_updated_at": importedCodexSourceUpdatedAt(*source), + "last_message_at": source.LastMessageAt, + "source_revision": piSourceRevision(threadPath, pointerString(record.NativeLeafID)), + "updated_at": time.Now(), + } + if err := m.updateRuntimeState(ctx, record.ID, updates); err != nil { + return ImportResult{}, err + } + if record.ArchivedAt != nil { + if _, err := m.UnarchiveSession(ctx, record.ID); err != nil { + return ImportResult{}, err + } + } + refreshed, err := m.GetSession(ctx, record.ID) + if err != nil { + return ImportResult{}, err + } + snapshot, err := m.syncImportedPiSession(ctx, refreshed) + if err != nil { + return ImportResult{}, err + } + return ImportResult{Session: snapshot.Session, History: snapshot.History, + PendingInputs: snapshot.PendingInputs, ScheduledInputs: snapshot.ScheduledInputs, + SubAgents: snapshot.SubAgents, Reused: true, Synced: true}, nil + } + + orderIndex, err := m.getNextSessionOrderIndex(ctx, project.Id) + if err != nil { + return ImportResult{}, err + } + title := strings.TrimSpace(source.Title) + titleAuto := title == "" + if titleAuto { + title = defaultTitle(AgentPi, project.Name) + } + modelName := strings.TrimSpace(source.Model) + if _, _, err := splitPiModel(modelName); err != nil { + modelName = "" + } + now := time.Now() + record := tables.WebSessionTable{ + ProjectID: project.Id, OrderIndex: orderIndex, Agent: string(AgentPi), + Backend: string(SessionBackendPiRPC), Title: title, TitleAuto: titleAuto, + Model: modelName, ReasoningEffort: string(ReasoningEffortDefault), + WorkflowMode: string(WorkflowModeDefault), PermissionLevel: string(PermissionLevelElevated), + LegacyPermissionMode: "default", Cwd: filepath.Clean(project.Path), NativeSessionID: &nativeID, + Status: string(StatusIdle), ActivityAt: now, StatusUpdatedAt: &now, + SourceKind: defaultSourceKind(AgentPi), SyncState: string(SyncStateMissing), + SourceCreatedAt: importedCodexSourceCreatedAt(*source), SourceUpdatedAt: importedCodexSourceUpdatedAt(*source), + ThreadPath: &threadPath, ThreadPreview: nilIfEmpty(source.Title), LastMessageAt: source.LastMessageAt, + SourceRevision: nilIfEmpty(piSourceRevision(threadPath, "")), + AutoRetryScope: string(AutoRetryScopeNetworkOnly), AutoRetryPreset: string(AutoRetryPresetGentleStop), + } + record.Init() + if err := model.GetDB().WithContext(ctx).Create(&record).Error; err != nil { + return ImportResult{}, err + } + snapshot, err := m.syncImportedPiSession(ctx, record) + if err != nil { + _ = m.DeleteSession(ctx, record.ID) + return ImportResult{}, err + } + return ImportResult{Session: snapshot.Session, History: snapshot.History, + PendingInputs: snapshot.PendingInputs, ScheduledInputs: snapshot.ScheduledInputs, + SubAgents: snapshot.SubAgents, Created: true, Synced: true}, nil +} + +func (m *Manager) ImportPiSessionBySessionID(ctx context.Context, projectID, sessionID string) (ImportResult, error) { + project, err := m.projectSvc.GetProject(ctx, projectID) + if err != nil { + return ImportResult{}, err + } + if m.aiSessionSvc == nil { + return ImportResult{}, errors.New("ai session service is not configured") + } + source, err := m.aiSessionSvc.ResolvePiSessionBySessionID(ctx, sessionID) + if err != nil { + return ImportResult{}, err + } + return m.importPiSessionResolved(ctx, project, source) +} + +func (m *Manager) ImportPiSession(ctx context.Context, projectID, aiSessionID string) (ImportResult, error) { + project, err := m.projectSvc.GetProject(ctx, projectID) + if err != nil { + return ImportResult{}, err + } + if m.aiSessionSvc == nil { + return ImportResult{}, errors.New("ai session service is not configured") + } + source, err := m.aiSessionSvc.ResolvePiSessionByID(ctx, aiSessionID) + if err != nil { + return ImportResult{}, err + } + return m.importPiSessionResolved(ctx, project, source) +} + func (m *Manager) GetSession(ctx context.Context, sessionID string) (tables.WebSessionTable, error) { db := model.GetDB() if db == nil { @@ -1607,7 +1870,7 @@ func (m *Manager) loadSnapshotLocal( Revision: summary.Revision, Session: summary, History: history, - PendingInputs: m.pendingInputsSnapshot(record.ID), + PendingInputs: m.pendingInputsDisplaySnapshot(record.ID), ScheduledInputs: scheduledInputs, PendingApproval: m.pendingApprovalSnapshot(record), PendingUserInput: pendingUserInputFromHistory(history.Items), @@ -1622,20 +1885,27 @@ func (m *Manager) pendingApprovalSnapshot(record tables.WebSessionTable) *Pendin m.mu.RLock() run := m.runs[record.ID] m.mu.RUnlock() - if run == nil || run.codexAppServer() == nil { + if run == nil { return nil } request, ok := run.pendingApprovalRequest() if !ok || request.Kind == pendingServerRequestPlanApproval { return nil } + if request.PiRuntime == nil && run.codexAppServer() == nil { + return nil + } + actionable := len(request.RawID) > 0 + if request.PiRuntime != nil { + actionable = strings.TrimSpace(request.PiRequestID) != "" + } return &PendingApproval{ ItemID: request.ItemID, Kind: string(request.Kind), Prompt: firstNonEmpty(request.Prompt, approvalPromptFallback(request.Kind)), Command: request.Command, RequestedAt: request.RequestedAt, - Actionable: len(request.RawID) > 0, + Actionable: actionable, } } @@ -1685,6 +1955,11 @@ func (m *Manager) UpdateModel(ctx context.Context, sessionID, modelName string) if err != nil { return SessionSummary{}, err } + if normalizeAgent(Agent(record.Agent)) == AgentPi && normalized != "" { + if _, _, err := splitPiModel(normalized); err != nil { + return SessionSummary{}, err + } + } updates := map[string]any{ "model": normalized, "updated_at": time.Now(), @@ -1719,8 +1994,18 @@ func (m *Manager) UpdateReasoningEffort( sessionID string, effort ReasoningEffort, ) (SessionSummary, error) { + record, err := m.GetSession(ctx, sessionID) + if err != nil { + return SessionSummary{}, err + } + normalized := normalizeReasoningEffort(effort) + if normalizeAgent(Agent(record.Agent)) == AgentPi { + if err := validatePiReasoningEffort(normalized); err != nil { + return SessionSummary{}, err + } + } return m.updateFields(ctx, sessionID, map[string]any{ - "reasoning_effort": string(normalizeReasoningEffort(effort)), + "reasoning_effort": string(normalized), "updated_at": time.Now(), }) } @@ -2048,7 +2333,10 @@ func (m *Manager) UpdateAutoRetryDispatchPendingOnFailure( } func (m *Manager) UpdateAgent(ctx context.Context, sessionID string, agent Agent) (SessionSummary, error) { - normalized := normalizeAgent(agent) + normalized, err := validateAgent(agent) + if err != nil { + return SessionSummary{}, err + } permissionLevel := m.resolveSessionPermissionLevel(normalized, "") modelName := m.resolveSessionModel(normalized, "") return m.updateFields(ctx, sessionID, map[string]any{ @@ -2155,6 +2443,7 @@ func (m *Manager) ArchiveSession(ctx context.Context, sessionID string) (Session if err := m.stopRunIfActive(sessionID, 5*time.Second); err != nil { return SessionSummary{}, err } + m.StopSessionPiRuntime(sessionID) now := time.Now() updates := map[string]any{ @@ -2235,6 +2524,7 @@ func (m *Manager) DeleteSession(ctx context.Context, sessionID string) error { if err := m.stopRunIfActive(sessionID, 5*time.Second); err != nil { return err } + m.StopSessionPiRuntime(sessionID) eventState := m.sessionEventState(sessionID) eventState.mu.Lock() @@ -2491,6 +2781,16 @@ func (m *Manager) HandleCommand(ctx context.Context, client *client, payload []b return m.handleConnectCommand(ctx, client, frame) case "send": return m.handleSendCommand(ctx, client, frame) + case "compact": + return m.handleCompactCommand(ctx, client, frame) + case "tree_get": + return m.handlePiTreeGetCommand(ctx, client, frame) + case "tree_nav": + return m.handlePiTreeNavigateCommand(ctx, client, frame) + case "tree_fork": + return m.handlePiTreeForkCommand(ctx, client, frame) + case "tree_clone": + return m.handlePiTreeCloneCommand(ctx, client, frame) case "hist": return m.handleHistoryCommand(ctx, client, frame) case "abort": @@ -3230,6 +3530,99 @@ func (m *Manager) handleListCommand(ctx context.Context, client *client, frame w return client.send(newAckFrame(frame.RequestID, frame.Operation, frame.SessionID, map[string]any{"items": result})) } +func (m *Manager) handleCompactCommand(ctx context.Context, client *client, frame wireCommandFrame) error { + if len(bytes.TrimSpace(frame.Payload)) > 0 && string(bytes.TrimSpace(frame.Payload)) != "{}" { + return client.send(newErrorFrame(frame.RequestID, frame.SessionID, "bad_req", "compact takes no payload", false)) + } + if err := m.CompactSession(ctx, frame.SessionID); err != nil { + return client.send(newErrorFrame(frame.RequestID, frame.SessionID, "invalid_state", err.Error(), false)) + } + return m.sendMutationAck(ctx, client, frame, nil) +} + +func (m *Manager) handlePiTreeGetCommand(ctx context.Context, client *client, frame wireCommandFrame) error { + if !m.SupportsPiSessionTree() { + return client.send(newErrorFrame(frame.RequestID, frame.SessionID, "unsupported", "Pi session tree is not supported", false)) + } + if len(bytes.TrimSpace(frame.Payload)) > 0 && string(bytes.TrimSpace(frame.Payload)) != "{}" { + return client.send(newErrorFrame(frame.RequestID, frame.SessionID, "bad_req", "tree_get takes no payload", false)) + } + tree, err := m.GetPiSessionTree(ctx, frame.SessionID) + if err != nil { + return client.send(newPiTreeErrorFrame(frame, err)) + } + return client.send(newAckFrame(frame.RequestID, frame.Operation, frame.SessionID, tree, m.currentSessionRevision(ctx, frame.SessionID))) +} + +func (m *Manager) handlePiTreeNavigateCommand(ctx context.Context, client *client, frame wireCommandFrame) error { + if !m.SupportsPiSessionTree() { + return client.send(newErrorFrame(frame.RequestID, frame.SessionID, "unsupported", "Pi session tree is not supported", false)) + } + var payload struct { + TargetID string `json:"tid"` + Revision string `json:"rev"` + Summarize bool `json:"sum"` + } + if err := json.Unmarshal(frame.Payload, &payload); err != nil { + return client.send(newErrorFrame(frame.RequestID, frame.SessionID, "bad_req", "invalid tree navigation payload", false)) + } + result, err := m.NavigatePiSessionTree(ctx, frame.SessionID, PiTreeNavigateInput{ + TargetID: payload.TargetID, Revision: payload.Revision, Summarize: payload.Summarize, + }) + if err != nil { + return client.send(newPiTreeErrorFrame(frame, err)) + } + return client.send(newAckFrame(frame.RequestID, frame.Operation, frame.SessionID, result, m.currentSessionRevision(ctx, frame.SessionID))) +} + +func (m *Manager) handlePiTreeForkCommand(ctx context.Context, client *client, frame wireCommandFrame) error { + if !m.SupportsPiSessionTree() { + return client.send(newErrorFrame(frame.RequestID, frame.SessionID, "unsupported", "Pi session tree is not supported", false)) + } + var payload struct { + TargetID string `json:"tid"` + Revision string `json:"rev"` + } + if err := json.Unmarshal(frame.Payload, &payload); err != nil { + return client.send(newErrorFrame(frame.RequestID, frame.SessionID, "bad_req", "invalid tree fork payload", false)) + } + result, err := m.ForkPiSessionTree(ctx, frame.SessionID, PiTreeForkInput{TargetID: payload.TargetID, Revision: payload.Revision}) + if err != nil { + return client.send(newPiTreeErrorFrame(frame, err)) + } + return client.send(newAckFrame(frame.RequestID, frame.Operation, frame.SessionID, mapPiTreeCreateWireResult(result))) +} + +func (m *Manager) handlePiTreeCloneCommand(ctx context.Context, client *client, frame wireCommandFrame) error { + if !m.SupportsPiSessionTree() { + return client.send(newErrorFrame(frame.RequestID, frame.SessionID, "unsupported", "Pi session tree is not supported", false)) + } + var payload struct { + Revision string `json:"rev"` + } + if err := json.Unmarshal(frame.Payload, &payload); err != nil { + return client.send(newErrorFrame(frame.RequestID, frame.SessionID, "bad_req", "invalid tree clone payload", false)) + } + result, err := m.ClonePiSessionTree(ctx, frame.SessionID, PiTreeCloneInput{Revision: payload.Revision}) + if err != nil { + return client.send(newPiTreeErrorFrame(frame, err)) + } + return client.send(newAckFrame(frame.RequestID, frame.Operation, frame.SessionID, mapPiTreeCreateWireResult(result))) +} + +func newPiTreeErrorFrame(frame wireCommandFrame, err error) wireFrame { + publicErr := ClassifyPiTreeError(err) + return newErrorFrame(frame.RequestID, frame.SessionID, publicErr.Code, publicErr.Message, false) +} + +func mapPiTreeCreateWireResult(result PiTreeCreateResult) map[string]any { + return map[string]any{ + "s": mapWireSession(result.Session), + "tree": result.Tree, + "editorText": result.EditorText, + } +} + func (m *Manager) handleSendCommand(ctx context.Context, client *client, frame wireCommandFrame) error { var payload struct { Text string `json:"txt"` @@ -3257,6 +3650,65 @@ func (m *Manager) SendMessage(ctx context.Context, sessionID, text string, attac return m.sendMessageInternal(ctx, sessionID, text, attachmentIDs, false) } +func (m *Manager) CompactSession(ctx context.Context, sessionID string) error { + dispatchLock := &m.sessionDispatchLocks[sessionRevisionLockIndex(sessionID)] + dispatchLock.Lock() + defer dispatchLock.Unlock() + + record, err := m.GetSession(ctx, sessionID) + if err != nil { + return err + } + if record.ArchivedAt != nil { + return errors.New("session is archived") + } + if normalizeAgent(Agent(record.Agent)) != AgentPi || effectiveSessionBackend(record) != SessionBackendPiRPC { + return errors.New("manual compaction is only supported for Pi RPC sessions") + } + if err := m.ensureSessionMessagingAvailable(record); err != nil { + return err + } + if record.NativeSessionID == nil || strings.TrimSpace(*record.NativeSessionID) == "" || + record.ThreadPath == nil || strings.TrimSpace(*record.ThreadPath) == "" { + return errors.New("Pi session has no native history to compact") + } + m.cancelAutoRetryTimer(sessionID) + if m.hasActiveRun(sessionID) { + return errors.New("session is already running") + } + + runID := utils.NewID() + now := time.Now() + if _, err := m.appendAndBroadcast(ctx, sessionID, record, Event{ + ID: utils.NewID(), Type: "run_st", RunID: runID, Timestamp: now, + Payload: map[string]any{ + "ag": string(AgentPi), "md": record.Model, "re": record.ReasoningEffort, + "wm": effectiveWorkflowMode(record), "pl": effectivePermissionLevel(record), "src": "compact", + }, + }); err != nil { + return err + } + if err := m.updateRuntimeState(ctx, sessionID, applyAssistantStateUpdates(map[string]any{ + "status": string(StatusRunning), "has_unread": false, "last_error": nil, + "auto_retry_attempt": 0, "auto_retry_next_at": nil, "auto_retry_last_error_code": nil, + "updated_at": now, + }, AssistantStateWorking, now)); err != nil { + return err + } + m.broadcastSessionSummary(ctx, sessionID) + + runCtx, cancel := context.WithCancel(context.Background()) + run := &activeRun{ + sessionID: sessionID, agent: AgentPi, backend: SessionBackendPiRPC, + runID: runID, cancel: cancel, done: make(chan struct{}), piCompaction: true, + } + m.mu.Lock() + m.runs[sessionID] = run + m.mu.Unlock() + go m.runSession(runCtx, run, record, "", nil) + return nil +} + func (m *Manager) ensureCodexThread(ctx context.Context, session tables.WebSessionTable) (string, error) { if normalizeAgent(Agent(session.Agent)) != AgentCodex { return "", fmt.Errorf("thread bootstrap is only supported for codex sessions") @@ -3552,15 +4004,26 @@ func (m *Manager) sendMessageInternal( } func (m *Manager) ensureSessionMessagingAvailable(record tables.WebSessionTable) error { - config := m.GetCodexRuntimeConfig() - switch normalizeAgent(Agent(record.Agent)) { + agent, err := validateAgent(Agent(record.Agent)) + if err != nil { + return err + } + config := m.GetWebSessionRuntimeConfig() + switch agent { case AgentCodex: if !config.HasCodex { - return fmt.Errorf(errCodexNotInstalled) + return fmt.Errorf("%s", errCodexNotInstalled) } case AgentClaude: if !config.HasClaudeCode { - return fmt.Errorf(errClaudeCodeNotInstalled) + return fmt.Errorf("%s", errClaudeCodeNotInstalled) + } + case AgentPi: + if !config.SupportsPiWebSession { + return fmt.Errorf("%s", errPiWebSessionUnavailable) + } + if err := m.EnsureProjectPiTrust(context.Background(), record.ProjectID, record.Cwd); err != nil { + return err } } return nil @@ -3637,6 +4100,14 @@ func (m *Manager) runSession(ctx context.Context, run *activeRun, session tables m.runCodexAppServerSession(ctx, run, session, text, attachments) return } + if run.backend == SessionBackendPiRPC && normalizeAgent(Agent(session.Agent)) == AgentPi { + if run.piCompaction { + m.runPiRPCCompaction(ctx, run, session) + } else { + m.runPiRPCSession(ctx, run, session, text, attachments) + } + return + } if run.claudeResumeOnly && normalizeAgent(Agent(session.Agent)) == AgentClaude { m.runClaudeResumeSession(ctx, run, session) return @@ -3935,6 +4406,9 @@ func (m *Manager) handleRunFailureWithCode( ) { if run != nil { run.resetActiveCallTracking() + if normalizeAgent(Agent(session.Agent)) == AgentPi { + _ = m.closePendingPiDialog(session, run, "Pi extension input ended because the runtime failed") + } } message := strings.TrimSpace(err.Error()) if message == "" { @@ -5355,15 +5829,37 @@ func upsertEnv(env []string, key, value string) []string { func (m *Manager) respondToApproval(sessionID, action string) error { dispatchLock := &m.sessionDispatchLocks[sessionRevisionLockIndex(sessionID)] dispatchLock.Lock() - defer dispatchLock.Unlock() m.mu.RLock() run, ok := m.runs[sessionID] m.mu.RUnlock() record, err := m.GetSession(context.Background(), sessionID) if err != nil { + dispatchLock.Unlock() return err } + if normalizeAgent(Agent(record.Agent)) == AgentPi { + if !ok || run == nil { + dispatchLock.Unlock() + return fmt.Errorf("no pending Pi approval") + } + pending, hasPending := run.pendingApprovalRequest() + if !hasPending || pending.PiRuntime == nil { + dispatchLock.Unlock() + return fmt.Errorf("no pending Pi approval") + } + request, taken := run.takePendingPiRequestForResponse(pending.PiRequestID, true) + if !taken { + dispatchLock.Unlock() + return fmt.Errorf("Pi approval is no longer pending") + } + dispatchLock.Unlock() + confirmed := action != "reject" && action != "cancel" + return m.respondPiExtensionRequest(record, run, request, map[string]any{"confirmed": confirmed}, "approval_res", map[string]any{ + "act": action, "prompt": request.Prompt, + }) + } + defer dispatchLock.Unlock() if normalizeAgent(Agent(record.Agent)) == AgentClaude { var pending *pendingServerRequest if ok && run != nil { @@ -5497,15 +5993,45 @@ func (m *Manager) respondToApproval(sessionID, action string) error { func (m *Manager) respondToUserInput(sessionID, itemID string, answers map[string][]string) error { dispatchLock := &m.sessionDispatchLocks[sessionRevisionLockIndex(sessionID)] dispatchLock.Lock() - defer dispatchLock.Unlock() m.mu.RLock() run, ok := m.runs[sessionID] m.mu.RUnlock() record, err := m.GetSession(context.Background(), sessionID) if err != nil { + dispatchLock.Unlock() return err } + if normalizeAgent(Agent(record.Agent)) == AgentPi { + if !ok || run == nil { + dispatchLock.Unlock() + return fmt.Errorf("no pending Pi user input request") + } + pending, hasPending := run.pendingUserInputRequest() + if !hasPending || pending.PiRuntime == nil { + dispatchLock.Unlock() + return fmt.Errorf("no pending Pi user input request") + } + if strings.TrimSpace(itemID) == "" || strings.TrimSpace(itemID) != strings.TrimSpace(pending.ItemID) { + dispatchLock.Unlock() + return fmt.Errorf("user input request does not match the active Pi prompt") + } + value := firstPiUserInputAnswer(answers) + if value == "" { + dispatchLock.Unlock() + return fmt.Errorf("no answers were provided") + } + request, taken := run.takePendingPiRequestForResponse(pending.PiRequestID, true) + if !taken { + dispatchLock.Unlock() + return fmt.Errorf("Pi user input is no longer pending") + } + dispatchLock.Unlock() + return m.respondPiExtensionRequest(record, run, request, map[string]any{"value": value}, "user_input_res", map[string]any{ + "iid": request.ItemID, "ans": answers, + }) + } + defer dispatchLock.Unlock() if normalizeAgent(Agent(record.Agent)) == AgentClaude { if ok && run != nil { if pending, hasPending := run.pendingUserInputRequest(); hasPending && @@ -5837,6 +6363,8 @@ func mapSessionRecord(record tables.WebSessionTable) SessionSummary { AutoRetryDispatchPendingOnFailure: record.AutoRetryDispatchPendingOnFailure, Cwd: record.Cwd, NativeSessionID: record.NativeSessionID, + NativeLeafID: record.NativeLeafID, + SourceRevision: record.SourceRevision, CyberPolicyFlagged: record.CyberPolicyFlagged, Status: effectiveStatus(record, assistantState), AssistantState: assistantState, @@ -6133,10 +6661,13 @@ func attachmentPayloads(items []Attachment) []map[string]any { func defaultTitle(agent Agent, projectName string) string { prefix := "Chat" - if normalizeAgent(agent) == AgentCodex { + switch normalizeAgent(agent) { + case AgentCodex: prefix = "Codex" - } else if normalizeAgent(agent) == AgentClaude { + case AgentClaude: prefix = "Claude" + case AgentPi: + prefix = "Pi" } if strings.TrimSpace(projectName) == "" { return prefix @@ -6148,10 +6679,14 @@ func defaultModel(agent Agent, provided string) string { if strings.TrimSpace(provided) != "" { return strings.TrimSpace(provided) } - if normalizeAgent(agent) == AgentCodex { + switch normalizeAgent(agent) { + case AgentCodex: return utils.DefaultWebSessionCodexModel + case AgentClaude: + return "opus" + default: + return "" } - return "opus" } func defaultReasoningEffort(agent Agent, provided ReasoningEffort) ReasoningEffort { @@ -6238,32 +6773,46 @@ func (m *Manager) resolveSessionPermissionLevel( } func defaultSessionBackend(agent Agent) SessionBackend { - if normalizeAgent(agent) == AgentCodex { + switch normalizeAgent(agent) { + case AgentCodex: return SessionBackendCodexAppServer + case AgentPi: + return SessionBackendPiRPC + default: + return SessionBackendLegacyExec } - return SessionBackendLegacyExec } func normalizeSessionBackend(backend SessionBackend, agent Agent) SessionBackend { + normalizedAgent := normalizeAgent(agent) switch strings.ToLower(strings.TrimSpace(string(backend))) { case string(SessionBackendCodexAppServer): - if normalizeAgent(agent) == AgentCodex { + if normalizedAgent == AgentCodex { return SessionBackendCodexAppServer } - return SessionBackendLegacyExec + case string(SessionBackendPiRPC): + if normalizedAgent == AgentPi { + return SessionBackendPiRPC + } case string(SessionBackendLegacyExec): - return SessionBackendLegacyExec - default: - return defaultSessionBackend(agent) + if normalizedAgent != AgentPi { + return SessionBackendLegacyExec + } } + return defaultSessionBackend(normalizedAgent) } func normalizeAgent(agent Agent) Agent { - switch agent { - case AgentCodex: - return AgentCodex + return Agent(strings.ToLower(strings.TrimSpace(string(agent)))) +} + +func validateAgent(agent Agent) (Agent, error) { + normalized := normalizeAgent(agent) + switch normalized { + case AgentClaude, AgentCodex, AgentPi: + return normalized, nil default: - return AgentClaude + return "", fmt.Errorf("invalid agent %q: expected claude, codex, or pi", strings.TrimSpace(string(agent))) } } @@ -6365,9 +6914,17 @@ func normalizePermissionLevel(level PermissionLevel) PermissionLevel { } func validateWebSessionPermissionLevel(agent Agent, level PermissionLevel) error { - if normalizeAgent(agent) == AgentClaude && normalizePermissionLevel(level) == PermissionLevelDefault { + normalizedAgent, err := validateAgent(agent) + if err != nil { + return err + } + normalizedLevel := normalizePermissionLevel(level) + if normalizedAgent == AgentClaude && normalizedLevel == PermissionLevelDefault { return fmt.Errorf("claude web sessions do not support the default permission level in claude_stream_json mode") } + if normalizedAgent == AgentPi && normalizedLevel != PermissionLevelElevated { + return fmt.Errorf("pi web sessions currently support only unrestricted access") + } return nil } @@ -6420,23 +6977,29 @@ func effectivePermissionLevel(record tables.WebSessionTable) PermissionLevel { } func effectiveSessionBackend(record tables.WebSessionTable) SessionBackend { + agent := normalizeAgent(Agent(record.Agent)) normalized := strings.ToLower(strings.TrimSpace(record.Backend)) switch normalized { case string(SessionBackendLegacyExec): - return SessionBackendLegacyExec + if agent != AgentPi { + return SessionBackendLegacyExec + } case string(SessionBackendCodexAppServer): - if normalizeAgent(Agent(record.Agent)) == AgentCodex { + if agent == AgentCodex { return SessionBackendCodexAppServer } - return SessionBackendLegacyExec + case string(SessionBackendPiRPC): + if agent == AgentPi { + return SessionBackendPiRPC + } default: - if normalizeAgent(Agent(record.Agent)) == AgentCodex { + if agent == AgentCodex { // Existing Codex sessions predate backend persistence and must continue // using the legacy exec transport unless explicitly migrated. return SessionBackendLegacyExec } - return SessionBackendLegacyExec } + return defaultSessionBackend(agent) } func preparePromptText(text string, workflowMode WorkflowMode) string { @@ -6880,6 +7443,79 @@ func (r *activeRun) blocksCodexSteerForUserInput() bool { (r.pendingServerReq != nil && r.pendingServerReq.Kind == pendingServerRequestUserInput) } +func (r *activeRun) blocksPiPendingInput() bool { + if r == nil { + return true + } + r.mu.Lock() + defer r.mu.Unlock() + return r.piResponsePending != 0 || r.pendingServerReq != nil +} + +func (r *activeRun) finishPiResponseBarrier(generation uint64) bool { + if r == nil { + return false + } + r.mu.Lock() + defer r.mu.Unlock() + if generation == 0 || r.piResponsePending != generation { + return false + } + r.piResponsePending = 0 + return r.pendingServerReq == nil +} + +func (r *activeRun) finishPiResponseHistory(generation uint64, persisted bool, err error) { + if r == nil || generation == 0 { + return + } + r.mu.Lock() + defer r.mu.Unlock() + if r.piResponseGeneration != generation || r.piResponseHistoryComplete == generation { + return + } + if persisted { + r.piResponseRequest = nil + } + r.piResponseHistoryErr = err + r.piResponseHistoryComplete = generation + if r.piResponseHistoryDone != nil { + close(r.piResponseHistoryDone) + r.piResponseHistoryDone = nil + } +} + +func (r *activeRun) waitForPiResponseHistory(ctx context.Context) (*pendingServerRequest, error) { + if r == nil { + return nil, nil + } + for { + r.mu.Lock() + generation := r.piResponsePending + if generation == 0 { + r.mu.Unlock() + return nil, nil + } + if r.piResponseHistoryComplete == generation { + request := r.piResponseRequest.clone() + err := r.piResponseHistoryErr + r.piResponseRequest = nil + r.mu.Unlock() + return request, err + } + done := r.piResponseHistoryDone + r.mu.Unlock() + if done == nil { + return nil, errors.New("Pi extension response completion state is unavailable") + } + select { + case <-done: + case <-ctx.Done(): + return nil, fmt.Errorf("wait for Pi extension response history: %w", ctx.Err()) + } + } +} + func (r *activeRun) clearPendingServerRequest() { r.mu.Lock() defer r.mu.Unlock() @@ -6899,6 +7535,39 @@ func (r *activeRun) clearPendingControlRequest(requestID string) bool { return true } +func (r *activeRun) takePendingPiRequest(requestID string) (*pendingServerRequest, bool) { + return r.takePendingPiRequestForResponse(requestID, false) +} + +func (r *activeRun) takePendingPiRequestForResponse(requestID string, markResponsePending bool) (*pendingServerRequest, bool) { + if r == nil { + return nil, false + } + r.mu.Lock() + defer r.mu.Unlock() + normalizedID := strings.TrimSpace(requestID) + if r.pendingServerReq == nil || normalizedID == "" || + strings.TrimSpace(r.pendingServerReq.PiRequestID) != normalizedID { + return nil, false + } + request := r.pendingServerReq.clone() + r.pendingServerReq = nil + if markResponsePending { + r.piResponseGeneration++ + if r.piResponseGeneration == 0 { + r.piResponseGeneration++ + } + r.piResponsePending = r.piResponseGeneration + r.piResponseHistoryDone = make(chan struct{}) + r.piResponseHistoryComplete = 0 + r.piResponseHistoryErr = nil + r.piResponseRequest = request.clone() + request.PiResponseGeneration = r.piResponseGeneration + r.piResponseRequest.PiResponseGeneration = r.piResponseGeneration + } + return request, true +} + func (r *activeRun) markCompletedPlanTool() { r.mu.Lock() defer r.mu.Unlock() diff --git a/service/websession/manager_test.go b/service/websession/manager_test.go index 15e9818f..7d86756b 100644 --- a/service/websession/manager_test.go +++ b/service/websession/manager_test.go @@ -141,6 +141,73 @@ func TestManagerCreateSessionAppendsOrderIndex(t *testing.T) { } } +func TestManagerCreateSessionRejectsInvalidAgent(t *testing.T) { + cleanup := initTestDB(t) + defer cleanup() + + project := seedProject(t) + manager, err := NewManager(Config{DataDir: t.TempDir()}, zap.NewNop()) + if err != nil { + t.Fatalf("NewManager returned error: %v", err) + } + + for _, agent := range []Agent{"", "unknown"} { + if _, err := manager.CreateSession(context.Background(), CreateParams{ + ProjectID: project.ID, + Agent: agent, + }); err == nil || !strings.Contains(err.Error(), "invalid agent") { + t.Fatalf("CreateSession agent %q error = %v, want invalid agent", agent, err) + } + } + + created, err := manager.CreateSession(context.Background(), CreateParams{ + ProjectID: project.ID, + Agent: AgentClaude, + }) + if err != nil { + t.Fatalf("CreateSession returned error: %v", err) + } + if _, err := manager.UpdateAgent(context.Background(), created.ID, "unknown"); err == nil || + !strings.Contains(err.Error(), "invalid agent") { + t.Fatalf("UpdateAgent error = %v, want invalid agent", err) + } +} + +func TestManagerCreateSessionUsesPiIdentityWithoutStartingAProcess(t *testing.T) { + cleanup := initTestDB(t) + defer cleanup() + + project := seedProject(t) + manager, err := NewManager(Config{DataDir: t.TempDir()}, zap.NewNop()) + if err != nil { + t.Fatalf("NewManager returned error: %v", err) + } + + created, err := manager.CreateSession(context.Background(), CreateParams{ + ProjectID: project.ID, + Agent: AgentPi, + }) + if err != nil { + t.Fatalf("CreateSession returned error: %v", err) + } + if created.Agent != AgentPi || created.SourceKind != string(SessionBackendPiRPC) { + t.Fatalf("unexpected Pi identity: agent=%q sourceKind=%q", created.Agent, created.SourceKind) + } + record, err := manager.GetSession(context.Background(), created.ID) + if err != nil { + t.Fatalf("GetSession returned error: %v", err) + } + if backend := effectiveSessionBackend(record); backend != SessionBackendPiRPC { + t.Fatalf("Pi backend = %q, want %q", backend, SessionBackendPiRPC) + } + manager.piRuntimeMu.Lock() + runtimeCount := len(manager.piRuntimes) + manager.piRuntimeMu.Unlock() + if runtimeCount != 0 || manager.hasActiveRun(created.ID) { + t.Fatalf("creating a Pi session started a process: runtimes=%d active=%v", runtimeCount, manager.hasActiveRun(created.ID)) + } +} + func TestManagerCreateSessionPersistsAutoRetryMaxAttempts(t *testing.T) { cleanup := initTestDB(t) defer cleanup() @@ -715,6 +782,59 @@ func TestImportCodexSessionRejectsProjectPathMismatch(t *testing.T) { } } +func TestListImportSourcesIncludesImportablePiPreview(t *testing.T) { + cleanup := initTestDB(t) + defer cleanup() + + project := seedProject(t) + piRoot := filepath.Join(t.TempDir(), "pi-sessions") + if err := os.MkdirAll(piRoot, 0o755); err != nil { + t.Fatalf("create Pi root: %v", err) + } + t.Setenv("PI_CODING_AGENT_SESSION_DIR", piRoot) + startedAt := time.Now().UTC().Add(-time.Hour) + piFile := filepath.Join(piRoot, "pi-import-source.jsonl") + piContent := fmt.Sprintf( + `{"type":"session","version":3,"id":"pi-native-id","timestamp":%q,"cwd":%q}`+"\n"+ + `{"type":"message","id":"message1","parentId":null,"timestamp":%q,"message":{"role":"user","content":"inspect Pi history","timestamp":%d}}`+"\n"+ + `{"type":"message","id":"message2","parentId":"message1","timestamp":%q,"message":{"role":"assistant","content":[{"type":"text","text":"done"}],"provider":"openai","model":"gpt-5","timestamp":%d}}`+"\n", + startedAt.Format(time.RFC3339Nano), + project.Path, + startedAt.Add(time.Second).Format(time.RFC3339Nano), + startedAt.Add(time.Second).UnixMilli(), + startedAt.Add(2*time.Second).Format(time.RFC3339Nano), + startedAt.Add(2*time.Second).UnixMilli(), + ) + if err := os.WriteFile(piFile, []byte(piContent), 0o644); err != nil { + t.Fatalf("write Pi source: %v", err) + } + + t.Setenv("CODEKANBAN_FAKE_PI_RUNTIME", "1") + t.Setenv("CODEKANBAN_FAKE_PI_SESSION", filepath.Join(piRoot, "unused-runtime.jsonl")) + manager, err := NewManager(Config{ + DataDir: t.TempDir(), + CodexPath: filepath.Join(t.TempDir(), "missing-codex"), + PiPath: fmt.Sprintf("%q -test.run=^TestPiRuntimeFakeProcess$ --", os.Args[0]), + }, zap.NewNop()) + if err != nil { + t.Fatalf("NewManager returned error: %v", err) + } + result, err := manager.ListImportSources(context.Background(), project.ID) + if err != nil { + t.Fatalf("ListImportSources returned error: %v", err) + } + if len(result.Items) != 1 { + t.Fatalf("import source count = %d, want 1: %#v", len(result.Items), result.Items) + } + item := result.Items[0] + if item.Agent != AgentPi || !item.Importable || item.SessionID != "pi-native-id" { + t.Fatalf("unexpected Pi import source: %#v", item) + } + if item.Title != "inspect Pi history" || item.Model != "openai/gpt-5" { + t.Fatalf("unexpected Pi preview metadata: %#v", item) + } +} + func TestListCodexImportSourcesUsesThreadListAndMarksDuplicates(t *testing.T) { cleanup := initTestDB(t) defer cleanup() @@ -2001,6 +2121,33 @@ func TestManagerListSessionsMarksClaudeContextWindowUnavailable(t *testing.T) { } } +func TestDecoratePiSessionSummaryPreservesObservedContextWindow(t *testing.T) { + manager := &Manager{} + observedWindow := int64(32000) + summary := SessionSummary{ + Agent: AgentPi, + ContextWindowTokens: &observedWindow, + ContextWindowSource: ContextWindowSourceSessionUsage, + } + + manager.decorateSessionSummary(&summary) + + if summary.ContextWindowTokens == nil || *summary.ContextWindowTokens != observedWindow { + t.Fatalf("expected observed Pi context window %d, got %#v", observedWindow, summary.ContextWindowTokens) + } + if summary.ContextWindowSource != ContextWindowSourceSessionUsage { + t.Fatalf("expected context window source %q, got %q", ContextWindowSourceSessionUsage, summary.ContextWindowSource) + } + + configuredWindow := int64(64000) + summary.ContextWindowTokens = &configuredWindow + summary.ContextWindowSource = ContextWindowSourceConfig + manager.decorateSessionSummary(&summary) + if summary.ContextWindowTokens != nil || summary.ContextWindowSource != ContextWindowSourceUnavailable { + t.Fatalf("Pi must not inherit a non-session context window: %#v", summary) + } +} + func TestGetCodexRuntimeConfigIncludesBinaryCapabilities(t *testing.T) { cleanup := initTestDB(t) defer cleanup() @@ -2009,6 +2156,7 @@ func TestGetCodexRuntimeConfigIncludesBinaryCapabilities(t *testing.T) { DataDir: t.TempDir(), CodexPath: writeFakeCodexVersionCLI(t, "0.146.0"), ClaudePath: writeFakeClaudeStreamCLI(t), + PiPath: filepath.Join(t.TempDir(), "missing-pi"), }, zap.NewNop()) if err != nil { t.Fatalf("NewManager returned error: %v", err) @@ -2039,6 +2187,15 @@ func TestGetCodexRuntimeConfigIncludesBinaryCapabilities(t *testing.T) { if config.GoalModeMinVersion != "0.133.0" { t.Fatalf("expected goalModeMinVersion 0.133.0, got %q", config.GoalModeMinVersion) } + if !config.Agents[AgentClaude].Installed || !config.Agents[AgentClaude].SupportsWebSession { + t.Fatalf("unexpected Claude capability: %#v", config.Agents[AgentClaude]) + } + if !config.Agents[AgentCodex].Installed || !config.Agents[AgentCodex].SupportsGoal { + t.Fatalf("unexpected Codex capability: %#v", config.Agents[AgentCodex]) + } + if config.Agents[AgentPi].Installed || config.Agents[AgentPi].SupportsWebSession { + t.Fatalf("Pi should be unavailable when PiPath is missing: %#v", config.Agents[AgentPi]) + } raw, err := json.Marshal(config) if err != nil { t.Fatalf("marshal runtime config: %v", err) @@ -2052,6 +2209,10 @@ func TestGetCodexRuntimeConfigIncludesBinaryCapabilities(t *testing.T) { payload["multiAgentV2MinCodexVersion"] != "0.146.0" { t.Fatalf("unexpected web-session capability payload: %#v", payload) } + agents, ok := payload["agents"].(map[string]any) + if !ok || agents["pi"] == nil { + t.Fatalf("runtime payload does not include Pi capability: %#v", payload["agents"]) + } } func TestGetCodexRuntimeConfigUsesCompatibilityModeForPreV2Version(t *testing.T) { diff --git a/service/websession/message_edit.go b/service/websession/message_edit.go index 2b532ce8..55399452 100644 --- a/service/websession/message_edit.go +++ b/service/websession/message_edit.go @@ -60,7 +60,12 @@ func (m *Manager) EditUserMessage( dispatchLock := &m.sessionDispatchLocks[sessionRevisionLockIndex(sessionID)] dispatchLock.Lock() - defer dispatchLock.Unlock() + dispatchLocked := true + defer func() { + if dispatchLocked { + dispatchLock.Unlock() + } + }() source, err := m.GetSession(ctx, sessionID) if err != nil { @@ -134,6 +139,11 @@ func (m *Manager) EditUserMessage( for _, attachment := range attachments { attachmentIDs = append(attachmentIDs, attachment.ID) } + + // The replacement runs on the new branch, so it must not inherit the source + // session's striped dispatch lock. The two IDs can hash to the same stripe. + dispatchLock.Unlock() + dispatchLocked = false if err := m.sendMessageInternal(ctx, branch.ID, text, attachmentIDs, false); err != nil { cleanupBranch() return SessionSnapshot{}, err diff --git a/service/websession/pending.go b/service/websession/pending.go index 5eb008a6..e0e1350a 100644 --- a/service/websession/pending.go +++ b/service/websession/pending.go @@ -2,6 +2,7 @@ package websession import ( "context" + "crypto/sha256" "errors" "fmt" "strings" @@ -51,6 +52,7 @@ func clonePendingInput(item PendingInput) PendingInput { AttachmentIDs: append([]string(nil), item.AttachmentIDs...), ReadyAt: readyAt, Paused: item.Paused, + NativeQueued: item.NativeQueued, CreatedAt: item.CreatedAt, } } @@ -99,6 +101,11 @@ func isCodexSteerSession(record tables.WebSessionTable) bool { effectiveSessionBackend(record) == SessionBackendCodexAppServer } +func isPiNativePendingSession(record tables.WebSessionTable) bool { + return normalizeAgent(Agent(record.Agent)) == AgentPi && + effectiveSessionBackend(record) == SessionBackendPiRPC +} + func (m *Manager) nextPendingSteerReadyAt() time.Time { delay := m.pendingSteerDelay if delay <= 0 { @@ -136,6 +143,56 @@ func (m *Manager) pendingInputsSnapshot(sessionID string) []PendingInput { return clonePendingInputs(m.pendingInputs[sessionID]) } +func (m *Manager) pendingInputsDisplaySnapshot(sessionID string) []PendingInput { + m.mu.RLock() + defer m.mu.RUnlock() + local := clonePendingInputs(m.pendingInputs[sessionID]) + native := clonePendingInputs(m.piNativeQueuedInputs[sessionID]) + return append(local, native...) +} + +func (m *Manager) replacePiNativeQueuedInputs(sessionID string, steering, followUp []string) { + now := time.Now() + items := make([]PendingInput, 0, len(steering)+len(followUp)) + appendItems := func(mode PendingInputMode, values []string) { + for index, value := range values { + text := strings.TrimSpace(value) + if text == "" { + continue + } + digest := sha256.Sum256([]byte(text)) + items = append(items, PendingInput{ + ID: fmt.Sprintf("pi-native:%s:%d:%x", mode, index, digest[:8]), + Mode: mode, + Text: text, + NativeQueued: true, + CreatedAt: now, + }) + } + } + appendItems(PendingInputModeRedirect, steering) + appendItems(PendingInputModeQueue, followUp) + + m.mu.Lock() + if len(items) == 0 { + delete(m.piNativeQueuedInputs, sessionID) + } else { + m.piNativeQueuedInputs[sessionID] = items + } + m.mu.Unlock() + m.broadcastPendingInputs(sessionID) +} + +func (m *Manager) clearPiNativeQueuedInputs(sessionID string) { + m.mu.Lock() + hadItems := len(m.piNativeQueuedInputs[sessionID]) > 0 + delete(m.piNativeQueuedInputs, sessionID) + m.mu.Unlock() + if hadItems { + m.broadcastPendingInputs(sessionID) + } +} + func (m *Manager) queuePendingInput( sessionID string, text string, @@ -681,7 +738,7 @@ func (m *Manager) broadcastPendingInputs(sessionID string) { return } _ = m.broadcastNextRevision(context.Background(), sessionID, func() (wireFrame, bool) { - return newPendingFrame(sessionID, m.pendingInputsSnapshot(sessionID)), true + return newPendingFrame(sessionID, m.pendingInputsDisplaySnapshot(sessionID)), true }) } @@ -733,13 +790,37 @@ func (m *Manager) runPendingProcessor(sessionID string) { } if m.hasActiveRun(sessionID) { - if next.Mode != PendingInputModeRedirect || - !isCodexSteerSession(record) { - return - } m.mu.RLock() run := m.runs[sessionID] m.mu.RUnlock() + if isPiNativePendingSession(record) { + if run == nil || run.blocksPiPendingInput() { + return + } + pending, claimed := m.claimPendingInput(sessionID, next.ID, next.Mode, time.Now()) + if !claimed { + continue + } + m.cancelPendingInputTimer(sessionID) + m.broadcastPendingInputs(sessionID) + handled, sendErr := m.sendActivePiPendingInput(ctx, record, pending) + if sendErr != nil || !handled { + m.prependPendingInput(sessionID, pending) + m.broadcastPendingInputs(sessionID) + if sendErr != nil && m.logger != nil { + m.logger.Debug("failed to send pending Pi input", + zap.String("sessionId", sessionID), + zap.String("pendingId", pending.ID), + zap.Error(sendErr), + ) + } + return + } + continue + } + if next.Mode != PendingInputModeRedirect || !isCodexSteerSession(record) { + return + } if run != nil && run.blocksCodexSteerForUserInput() { return } @@ -754,18 +835,13 @@ func (m *Manager) runPendingProcessor(sessionID string) { m.setPendingInputTimer(sessionID, readyAt) return } - steerInput, ok := m.claimPendingInput(sessionID, next.ID, next.Mode, time.Now()) - if !ok { + steerInput, claimed := m.claimPendingInput(sessionID, next.ID, next.Mode, time.Now()) + if !claimed { continue } m.cancelPendingInputTimer(sessionID) m.broadcastPendingInputs(sessionID) - handled, steerErr := m.steerActiveCodexTurn( - ctx, - record, - steerInput.Text, - steerInput.AttachmentIDs, - ) + handled, steerErr := m.steerActiveCodexTurn(ctx, record, steerInput.Text, steerInput.AttachmentIDs) if steerErr != nil || !handled { m.prependPendingInput(sessionID, steerInput) m.broadcastPendingInputs(sessionID) @@ -822,7 +898,7 @@ func (m *Manager) maybeInterruptForRedirect(sessionID string) { return } record, err := m.GetSession(context.Background(), sessionID) - if err == nil && isCodexSteerSession(record) { + if err == nil && (isCodexSteerSession(record) || isPiNativePendingSession(record)) { m.triggerPendingProcessing(sessionID) return } diff --git a/service/websession/pi_bridge.go b/service/websession/pi_bridge.go new file mode 100644 index 00000000..c302a078 --- /dev/null +++ b/service/websession/pi_bridge.go @@ -0,0 +1,153 @@ +package websession + +import ( + "bytes" + "crypto/sha256" + _ "embed" + "encoding/hex" + "errors" + "fmt" + "os" + "path/filepath" + "strings" + "sync" +) + +const ( + piBridgeCommandName = "codekanban-navigate" + piBridgeMarkerType = "codekanban.active-leaf.v1" +) + +//go:embed pi_bridge/extension.ts +var piBridgeSource []byte + +var piBridgeMaterializeMu sync.Mutex + +func (m *Manager) materializePiBridge() (string, error) { + piBridgeMaterializeMu.Lock() + defer piBridgeMaterializeMu.Unlock() + if m == nil || strings.TrimSpace(m.cfg.DataDir) == "" { + return "", errors.New("Pi bridge data directory is not configured") + } + digest := sha256.Sum256(piBridgeSource) + hash := hex.EncodeToString(digest[:]) + root := filepath.Join(m.cfg.DataDir, "pi-bridge") + dir := filepath.Join(root, hash) + path := filepath.Join(dir, "extension.ts") + + if err := os.MkdirAll(dir, 0o700); err != nil { + return "", fmt.Errorf("create Pi bridge directory: %w", err) + } + if err := ensurePiBridgeContained(root, path); err != nil { + return "", err + } + if _, err := validatePiBridgeArtifact(path); err == nil { + return filepath.Clean(path), nil + } else if !os.IsNotExist(err) { + return "", fmt.Errorf("inspect Pi bridge artifact: %w", err) + } + + temp, err := os.CreateTemp(dir, ".extension-*.tmp") + if err != nil { + return "", fmt.Errorf("create Pi bridge artifact: %w", err) + } + tempPath := temp.Name() + removeTemp := true + defer func() { + _ = temp.Close() + if removeTemp { + _ = os.Remove(tempPath) + } + }() + if err := temp.Chmod(0o600); err != nil { + return "", fmt.Errorf("secure Pi bridge artifact: %w", err) + } + if _, err := temp.Write(piBridgeSource); err != nil { + return "", fmt.Errorf("write Pi bridge artifact: %w", err) + } + if err := temp.Sync(); err != nil { + return "", fmt.Errorf("sync Pi bridge artifact: %w", err) + } + if err := temp.Close(); err != nil { + return "", fmt.Errorf("close Pi bridge artifact: %w", err) + } + if err := os.Rename(tempPath, path); err != nil { + // Another runtime may have installed this immutable hash artifact first. + // Accept only the exact embedded regular file; every other race fails closed. + if _, validateErr := validatePiBridgeArtifact(path); validateErr != nil { + return "", fmt.Errorf("install Pi bridge artifact: %w", err) + } + return filepath.Clean(path), nil + } + removeTemp = false + return filepath.Clean(path), nil +} + +func validatePiBridgeArtifact(path string) (os.FileInfo, error) { + info, err := os.Lstat(path) + if err != nil { + return nil, err + } + if !info.Mode().IsRegular() { + return nil, errors.New("Pi bridge artifact is not a regular file") + } + onDisk, err := os.ReadFile(path) + if err != nil { + return nil, fmt.Errorf("read Pi bridge artifact: %w", err) + } + if !bytes.Equal(onDisk, piBridgeSource) { + return nil, errors.New("Pi bridge artifact content does not match the embedded extension") + } + return info, nil +} + +func ensurePiBridgeContained(root, path string) error { + canonicalRoot, err := canonicalPiRuntimePath(root) + if err != nil { + return fmt.Errorf("resolve Pi bridge root: %w", err) + } + canonicalPath, err := canonicalPiRuntimePath(path) + if err != nil { + return fmt.Errorf("resolve Pi bridge artifact: %w", err) + } + relative, err := filepath.Rel(canonicalRoot, canonicalPath) + if err != nil || relative == "." || filepath.IsAbs(relative) || relative == ".." || strings.HasPrefix(relative, ".."+string(filepath.Separator)) { + return errors.New("Pi bridge artifact is outside the managed bridge root") + } + return nil +} + +type piRPCSourceInfo struct { + Path string `json:"path"` + Source string `json:"source"` + Scope string `json:"scope"` + Origin string `json:"origin"` + BaseDir string `json:"baseDir"` +} + +type piRPCSlashCommand struct { + Name string `json:"name"` + Source string `json:"source"` + SourceInfo piRPCSourceInfo `json:"sourceInfo"` +} + +func validatePiBridgeCommands(commands []piRPCSlashCommand, expectedPath string) error { + matches := make([]piRPCSlashCommand, 0, 1) + for _, command := range commands { + if strings.TrimSpace(command.Name) == piBridgeCommandName { + matches = append(matches, command) + } + } + if len(matches) != 1 { + return fmt.Errorf("Pi bridge command provenance is ambiguous: found %d registrations", len(matches)) + } + command := matches[0] + if strings.TrimSpace(command.Source) != "extension" || + strings.TrimSpace(command.SourceInfo.Source) != "cli" || + strings.TrimSpace(command.SourceInfo.Scope) != "temporary" || + strings.TrimSpace(command.SourceInfo.Origin) != "top-level" || + !samePiRuntimePath(command.SourceInfo.Path, expectedPath) { + return errors.New("Pi bridge command provenance validation failed") + } + return nil +} diff --git a/service/websession/pi_bridge/extension.ts b/service/websession/pi_bridge/extension.ts new file mode 100644 index 00000000..26f8c5cc --- /dev/null +++ b/service/websession/pi_bridge/extension.ts @@ -0,0 +1,44 @@ +import type { ExtensionAPI } from "@earendil-works/pi-coding-agent"; + +const COMMAND_NAME = "codekanban-navigate"; +const MARKER_TYPE = "codekanban.active-leaf.v1"; + +interface NavigatePayload { + targetId: string; + summarize: boolean; + nonce: string; +} + +function decodePayload(raw: string): NavigatePayload { + const parsed = JSON.parse(Buffer.from(raw.trim(), "base64url").toString("utf8")) as Partial; + if ( + typeof parsed.targetId !== "string" || + parsed.targetId.trim() === "" || + typeof parsed.summarize !== "boolean" || + typeof parsed.nonce !== "string" || + parsed.nonce.trim() === "" + ) { + throw new Error("invalid CodeKanban navigation payload"); + } + return { + targetId: parsed.targetId.trim(), + summarize: parsed.summarize, + nonce: parsed.nonce.trim(), + }; +} + +export default function (pi: ExtensionAPI) { + pi.registerCommand(COMMAND_NAME, { + description: "Internal CodeKanban session-tree navigation bridge", + handler: async (args, ctx) => { + const payload = decodePayload(args); + const result = await ctx.navigateTree(payload.targetId, { + summarize: payload.summarize, + }); + if (result.cancelled) { + throw new Error("Pi session-tree navigation was cancelled"); + } + pi.appendEntry(MARKER_TYPE, payload); + }, + }); +} diff --git a/service/websession/pi_bridge_test.go b/service/websession/pi_bridge_test.go new file mode 100644 index 00000000..c47d5a7d --- /dev/null +++ b/service/websession/pi_bridge_test.go @@ -0,0 +1,126 @@ +package websession + +import ( + "os" + "path/filepath" + "sync" + "testing" + + "go.uber.org/zap" +) + +func TestPiBridgeMaterializerConcurrentInstall(t *testing.T) { + cleanup := initTestDB(t) + defer cleanup() + manager, err := NewManager(Config{DataDir: t.TempDir()}, zap.NewNop()) + if err != nil { + t.Fatalf("NewManager: %v", err) + } + + const workers = 32 + paths := make([]string, workers) + errors := make([]error, workers) + var wait sync.WaitGroup + wait.Add(workers) + for index := range workers { + go func() { + defer wait.Done() + paths[index], errors[index] = manager.materializePiBridge() + }() + } + wait.Wait() + + for index := range workers { + if errors[index] != nil { + t.Fatalf("materialize %d: %v", index, errors[index]) + } + if paths[index] != paths[0] { + t.Fatalf("materialize %d path = %q, want %q", index, paths[index], paths[0]) + } + } + content, err := os.ReadFile(paths[0]) + if err != nil { + t.Fatalf("read bridge: %v", err) + } + if string(content) != string(piBridgeSource) { + t.Fatal("materialized bridge does not match embedded source") + } + matches, err := filepath.Glob(filepath.Join(filepath.Dir(paths[0]), ".extension-*.tmp")) + if err != nil { + t.Fatalf("glob temp artifacts: %v", err) + } + if len(matches) != 0 { + t.Fatalf("temporary bridge artifacts were not cleaned up: %v", matches) + } +} + +func TestPiBridgeMaterializerAndProvenance(t *testing.T) { + cleanup := initTestDB(t) + defer cleanup() + manager, err := NewManager(Config{DataDir: t.TempDir()}, zap.NewNop()) + if err != nil { + t.Fatalf("NewManager: %v", err) + } + path, err := manager.materializePiBridge() + if err != nil { + t.Fatalf("materializePiBridge: %v", err) + } + if !filepath.IsAbs(path) { + t.Fatalf("bridge path = %q, want absolute", path) + } + content, err := os.ReadFile(path) + if err != nil { + t.Fatalf("read bridge: %v", err) + } + if string(content) != string(piBridgeSource) { + t.Fatal("materialized bridge does not match embedded source") + } + + valid := piRPCSlashCommand{ + Name: piBridgeCommandName, + Source: "extension", + SourceInfo: piRPCSourceInfo{ + Path: path, Source: "cli", Scope: "temporary", Origin: "top-level", + }, + } + if err := validatePiBridgeCommands([]piRPCSlashCommand{valid}, path); err != nil { + t.Fatalf("valid provenance: %v", err) + } + for name, commands := range map[string][]piRPCSlashCommand{ + "missing": nil, + "duplicate": {valid, valid}, + "wrong command source": {{ + Name: piBridgeCommandName, Source: "prompt", + SourceInfo: piRPCSourceInfo{Path: path, Source: "cli", Scope: "temporary", Origin: "top-level"}, + }}, + "wrong provenance source": {{ + Name: piBridgeCommandName, Source: "extension", + SourceInfo: piRPCSourceInfo{Path: path, Source: "extension", Scope: "temporary", Origin: "top-level"}, + }}, + "wrong path": {{ + Name: piBridgeCommandName, Source: "extension", + SourceInfo: piRPCSourceInfo{Path: filepath.Join(filepath.Dir(path), "other.ts"), Source: "cli", Scope: "temporary", Origin: "top-level"}, + }}, + "wrong scope": {{ + Name: piBridgeCommandName, Source: "extension", + SourceInfo: piRPCSourceInfo{Path: path, Source: "cli", Scope: "project", Origin: "top-level"}, + }}, + "wrong origin": {{ + Name: piBridgeCommandName, Source: "extension", + SourceInfo: piRPCSourceInfo{Path: path, Source: "cli", Scope: "temporary", Origin: "package"}, + }}, + } { + t.Run(name, func(t *testing.T) { + if err := validatePiBridgeCommands(commands, path); err == nil { + t.Fatal("expected provenance validation to fail") + } + }) + } + + if err := os.WriteFile(path, []byte("tampered"), 0o600); err != nil { + t.Fatalf("tamper bridge: %v", err) + } + if _, err := manager.materializePiBridge(); err == nil { + t.Fatal("expected tampered bridge to fail closed") + } +} diff --git a/service/websession/pi_command_nonwindows.go b/service/websession/pi_command_nonwindows.go new file mode 100644 index 00000000..bc96c064 --- /dev/null +++ b/service/websession/pi_command_nonwindows.go @@ -0,0 +1,17 @@ +//go:build !windows + +package websession + +import ( + "context" + "os/exec" +) + +func buildWindowsBatchCommand( + ctx context.Context, + comspec string, + batchPath string, + args []string, +) *exec.Cmd { + return exec.CommandContext(ctx, batchPath, args...) +} diff --git a/service/websession/pi_command_windows.go b/service/websession/pi_command_windows.go new file mode 100644 index 00000000..f183e4a4 --- /dev/null +++ b/service/websession/pi_command_windows.go @@ -0,0 +1,37 @@ +//go:build windows + +package websession + +import ( + "context" + "os/exec" + "strings" + "syscall" +) + +func buildWindowsBatchCommand( + ctx context.Context, + comspec string, + batchPath string, + args []string, +) *exec.Cmd { + quotedArgs := make([]string, 0, len(args)) + for _, arg := range args { + quotedArgs = append(quotedArgs, quoteWindowsBatchArgument(arg)) + } + commandLine := `/d /s /c ""` + strings.ReplaceAll(batchPath, `"`, `""`) + `"` + if len(quotedArgs) > 0 { + commandLine += " " + strings.Join(quotedArgs, " ") + } + commandLine += `"` + + cmd := exec.CommandContext(ctx, comspec) + cmd.SysProcAttr = &syscall.SysProcAttr{CmdLine: commandLine} + return cmd +} + +func quoteWindowsBatchArgument(value string) string { + value = strings.ReplaceAll(value, "\r", " ") + value = strings.ReplaceAll(value, "\n", " ") + return `"` + strings.ReplaceAll(value, `"`, `""`) + `"` +} diff --git a/service/websession/pi_projection.go b/service/websession/pi_projection.go new file mode 100644 index 00000000..08a4083f --- /dev/null +++ b/service/websession/pi_projection.go @@ -0,0 +1,772 @@ +package websession + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "strings" + "time" + + "code-kanban/model/tables" + "code-kanban/utils" +) + +const piToolProgressInterval = 100 * time.Millisecond + +type piRPCMessage struct { + Role string `json:"role"` + Timestamp int64 `json:"timestamp"` + StopReason string `json:"stopReason"` + ErrorMessage string `json:"errorMessage"` + Content []struct { + Type string `json:"type"` + Text string `json:"text"` + Thinking string `json:"thinking"` + ID string `json:"id"` + Name string `json:"name"` + Arguments map[string]any `json:"arguments"` + } `json:"content"` +} + +type piAssistantMessageEvent struct { + Type string `json:"type"` + ContentIndex int `json:"contentIndex"` + Delta string `json:"delta"` + Content string `json:"content"` + ToolCall struct { + Type string `json:"type"` + ID string `json:"id"` + Name string `json:"name"` + Arguments map[string]any `json:"arguments"` + } `json:"toolCall"` +} + +func (m *Manager) handlePiRuntimeEvent(dispatch *piRuntimeRun, event piRPCEvent) error { + if dispatch == nil || dispatch.run == nil { + return nil + } + switch event.Type { + case "agent_start": + now := time.Now() + return m.updateRuntimeState(context.Background(), dispatch.session.ID, applyAssistantStateUpdates(map[string]any{ + "status": string(StatusRunning), "updated_at": now, + }, AssistantStateWorking, now)) + case "message_start": + var payload struct { + Message piRPCMessage `json:"message"` + } + if err := json.Unmarshal(event.Raw, &payload); err != nil { + return fmt.Errorf("decode Pi message_start: %w", err) + } + if strings.EqualFold(payload.Message.Role, "assistant") { + return m.startPiAssistantMessage(dispatch) + } + case "message_update": + var payload struct { + AssistantMessageEvent piAssistantMessageEvent `json:"assistantMessageEvent"` + } + if err := json.Unmarshal(event.Raw, &payload); err != nil { + return fmt.Errorf("decode Pi message_update: %w", err) + } + return m.handlePiAssistantMessageEvent(dispatch, payload.AssistantMessageEvent) + case "message_end": + var payload struct { + Message piRPCMessage `json:"message"` + } + if err := json.Unmarshal(event.Raw, &payload); err != nil { + return fmt.Errorf("decode Pi message_end: %w", err) + } + if strings.EqualFold(payload.Message.Role, "assistant") { + return m.finishPiAssistantMessage(dispatch, payload.Message) + } + case "tool_execution_start", "tool_execution_update", "tool_execution_end": + return m.handlePiToolExecution(dispatch, event) + case "compaction_start", "compaction_end": + return m.handlePiCompactionEvent(dispatch, event) + case "queue_update": + return m.handlePiQueueUpdate(dispatch, event) + case "auto_retry_start", "auto_retry_end", "summarization_retry_scheduled", "summarization_retry_attempt_start", "summarization_retry_finished": + return m.handlePiRetryEvent(dispatch, event) + case "extension_ui_request": + return m.handlePiExtensionUIRequest(dispatch, event.Raw) + case "extension_error": + var payload struct { + Error string `json:"error"` + Path string `json:"extensionPath"` + } + if err := json.Unmarshal(event.Raw, &payload); err != nil { + return fmt.Errorf("decode Pi extension_error: %w", err) + } + text := "Pi extension failed" + if strings.TrimSpace(payload.Path) != "" { + text += ": " + strings.TrimSpace(payload.Path) + } + m.appendRunNote(dispatch.session.ID, dispatch.session, dispatch.run, "warning", text, map[string]any{"code": "pi_extension_error"}) + case "agent_settled": + return m.finishPiSettledProjection(dispatch) + } + return nil +} + +func (m *Manager) startPiAssistantMessage(dispatch *piRuntimeRun) error { + dispatch.mu.Lock() + if dispatch.assistantMessageOpen { + dispatch.mu.Unlock() + return nil + } + messageID := utils.NewID() + dispatch.assistantMessageID = messageID + dispatch.assistantMessageOpen = true + dispatch.contents = make(map[int]*piRuntimeContentState) + dispatch.mu.Unlock() + + dispatch.run.mu.Lock() + dispatch.run.assistantMessageID = messageID + dispatch.run.mu.Unlock() + _, err := m.appendAndBroadcast(context.Background(), dispatch.session.ID, dispatch.session, Event{ + ID: utils.NewID(), Type: "msg_a_st", RunID: dispatch.run.runID, + ParentID: messageID, Timestamp: time.Now(), Payload: map[string]any{"mid": messageID}, + }) + return err +} + +func (m *Manager) ensurePiAssistantMessage(dispatch *piRuntimeRun) error { + dispatch.mu.Lock() + open := dispatch.assistantMessageOpen + dispatch.mu.Unlock() + if open { + return nil + } + return m.startPiAssistantMessage(dispatch) +} + +func (m *Manager) handlePiAssistantMessageEvent(dispatch *piRuntimeRun, update piAssistantMessageEvent) error { + if err := m.ensurePiAssistantMessage(dispatch); err != nil { + return err + } + dispatch.mu.Lock() + messageID := dispatch.assistantMessageID + state := dispatch.contents[update.ContentIndex] + if state == nil { + state = &piRuntimeContentState{} + dispatch.contents[update.ContentIndex] = state + } + switch update.Type { + case "text_start", "text_delta", "text_end": + state.kind = "text" + if update.Type == "text_delta" { + state.text += update.Delta + } else if update.Type == "text_end" { + state.text = update.Content + } + case "thinking_start", "thinking_delta", "thinking_end": + state.kind = "thinking" + if update.Type == "thinking_delta" { + state.text += update.Delta + } else if update.Type == "thinking_end" { + state.text = update.Content + } + case "toolcall_start", "toolcall_delta", "toolcall_end": + state.kind = "toolCall" + if update.Type == "toolcall_delta" { + state.text += update.Delta + } else if update.Type == "toolcall_end" { + state.toolID = strings.TrimSpace(update.ToolCall.ID) + if state.toolID != "" { + tool := dispatch.tools[state.toolID] + if tool == nil { + tool = &piRuntimeToolState{id: state.toolID} + dispatch.tools[state.toolID] = tool + } + tool.name = strings.TrimSpace(update.ToolCall.Name) + tool.args = update.ToolCall.Arguments + tool.parentID = messageID + } + } + } + content := state.text + dispatch.mu.Unlock() + + switch update.Type { + case "text_delta": + if update.Delta == "" { + return nil + } + dispatch.run.markAssistantDeltaSeen(messageID) + _, err := m.appendAndBroadcast(context.Background(), dispatch.session.ID, dispatch.session, Event{ + ID: utils.NewID(), Type: "txt_d", RunID: dispatch.run.runID, ParentID: messageID, + Timestamp: time.Now(), Payload: map[string]any{"mid": messageID, "txt": update.Delta}, + }) + return err + case "thinking_start", "thinking_delta": + return m.emitPiThinking(dispatch, messageID, update.ContentIndex, content, false) + case "thinking_end": + return m.emitPiThinking(dispatch, messageID, update.ContentIndex, content, true) + } + return nil +} + +func (m *Manager) emitPiThinking(dispatch *piRuntimeRun, messageID string, index int, text string, done bool) error { + toolID := fmt.Sprintf("pi-thinking:%s:%s:%d", dispatch.run.runID, messageID, index) + eventType := "tool_st" + if done { + eventType = "tool_end" + } + _, err := m.appendAndBroadcast(context.Background(), dispatch.session.ID, dispatch.session, Event{ + ID: utils.NewID(), Type: eventType, RunID: dispatch.run.runID, ParentID: messageID, + Timestamp: time.Now(), Payload: map[string]any{ + "tid": toolID, "name": "Reasoning", "kind": "reasoning", "out": text, "ok": true, + }, + }) + return err +} + +func (m *Manager) finishPiAssistantMessage(dispatch *piRuntimeRun, message piRPCMessage) error { + if err := m.ensurePiAssistantMessage(dispatch); err != nil { + return err + } + dispatch.mu.Lock() + messageID := dispatch.assistantMessageID + dispatch.assistantMessageOpen = false + dispatch.lastAttemptError = strings.TrimSpace(message.ErrorMessage) + if dispatch.lastAttemptError == "" && (strings.EqualFold(message.StopReason, "error") || strings.EqualFold(message.StopReason, "aborted")) { + dispatch.lastAttemptError = "Pi assistant attempt failed" + } + dispatch.mu.Unlock() + + for index, block := range message.Content { + if block.Type != "thinking" { + continue + } + if err := m.emitPiThinking(dispatch, messageID, index, block.Thinking, true); err != nil { + return err + } + } + text := piMessageText(message) + _, err := m.appendAndBroadcast(context.Background(), dispatch.session.ID, dispatch.session, Event{ + ID: utils.NewID(), Type: "txt_end", RunID: dispatch.run.runID, ParentID: messageID, + Timestamp: time.Now(), Payload: map[string]any{"mid": messageID, "txt": text}, + }) + if err == nil && dispatch.lastAttemptError == "" { + dispatch.run.mu.Lock() + dispatch.run.completedReply = true + dispatch.run.mu.Unlock() + } + return err +} + +func (m *Manager) handlePiToolExecution(dispatch *piRuntimeRun, event piRPCEvent) error { + var payload struct { + ToolCallID string `json:"toolCallId"` + ToolName string `json:"toolName"` + Args map[string]any `json:"args"` + PartialResult any `json:"partialResult"` + Result any `json:"result"` + IsError bool `json:"isError"` + } + if err := json.Unmarshal(event.Raw, &payload); err != nil { + return fmt.Errorf("decode Pi %s: %w", event.Type, err) + } + toolID := strings.TrimSpace(payload.ToolCallID) + if toolID == "" { + return errors.New("Pi tool event is missing toolCallId") + } + now := time.Now() + dispatch.mu.Lock() + tool := dispatch.tools[toolID] + if tool == nil { + tool = &piRuntimeToolState{id: toolID} + dispatch.tools[toolID] = tool + } + if strings.TrimSpace(payload.ToolName) != "" { + tool.name = strings.TrimSpace(payload.ToolName) + } + if payload.Args != nil { + tool.args = payload.Args + } + if tool.parentID == "" { + tool.parentID = dispatch.assistantMessageID + } + outputValue := payload.PartialResult + if event.Type == "tool_execution_end" { + outputValue = payload.Result + } + if outputValue != nil { + tool.output = piToolResultText(outputValue) + } + if event.Type == "tool_execution_update" && !tool.lastEmit.IsZero() && now.Sub(tool.lastEmit) < piToolProgressInterval { + dispatch.mu.Unlock() + return nil + } + tool.lastEmit = now + if event.Type == "tool_execution_end" { + tool.completed = true + } + snapshot := *tool + dispatch.mu.Unlock() + + eventType := "tool_st" + if event.Type == "tool_execution_end" { + eventType = "tool_end" + } + _, err := m.appendAndBroadcast(context.Background(), dispatch.session.ID, dispatch.session, Event{ + ID: utils.NewID(), Type: eventType, RunID: dispatch.run.runID, ParentID: snapshot.parentID, + Timestamp: now, Payload: map[string]any{ + "tid": snapshot.id, "name": firstNonEmpty(snapshot.name, "Tool"), "kind": "tool", + "in": snapshot.args, "out": snapshot.output, "ok": !payload.IsError, + }, + }) + return err +} + +func (m *Manager) handlePiCompactionEvent(dispatch *piRuntimeRun, event piRPCEvent) error { + var payload map[string]any + if err := json.Unmarshal(event.Raw, &payload); err != nil { + return fmt.Errorf("decode Pi %s: %w", event.Type, err) + } + dispatch.mu.Lock() + if dispatch.compactionToolID == "" { + dispatch.compactionToolID = "pi-compaction:" + utils.NewID() + dispatch.compactionStarted = time.Now() + } + toolID := dispatch.compactionToolID + parentID := dispatch.assistantMessageID + if event.Type == "compaction_end" { + dispatch.compactionToolID = "" + dispatch.compactionStarted = time.Time{} + } + dispatch.mu.Unlock() + + reason := strings.TrimSpace(stringValue(payload["reason"])) + output := "Pi is compacting the conversation context." + eventType := "tool_st" + ok := true + if event.Type == "compaction_end" { + eventType = "tool_end" + ok = !boolValue(payload["aborted"]) && strings.TrimSpace(stringValue(payload["errorMessage"])) == "" + output = firstNonEmpty(strings.TrimSpace(stringValue(decodeRawObject(payload["result"])["summary"])), "Context compacted") + if ok { + dispatch.mu.Lock() + dispatch.compactionCompleted = time.Now() + dispatch.mu.Unlock() + } else { + output = firstNonEmpty(strings.TrimSpace(stringValue(payload["errorMessage"])), "Pi context compaction did not complete") + } + } + _, err := m.appendAndBroadcast(context.Background(), dispatch.session.ID, dispatch.session, Event{ + ID: utils.NewID(), Type: eventType, RunID: dispatch.run.runID, ParentID: parentID, + Timestamp: time.Now(), Payload: map[string]any{ + "tid": toolID, "name": "ContextCompaction", "kind": "context_compaction", + "in": map[string]any{"reason": reason}, "out": output, "ok": ok, + }, + }) + return err +} + +func (m *Manager) handlePiQueueUpdate(dispatch *piRuntimeRun, event piRPCEvent) error { + var payload struct { + Steering []string `json:"steering"` + FollowUp []string `json:"followUp"` + } + if err := json.Unmarshal(event.Raw, &payload); err != nil { + return fmt.Errorf("decode Pi queue_update: %w", err) + } + m.replacePiNativeQueuedInputs(dispatch.session.ID, payload.Steering, payload.FollowUp) + return nil +} + +func (m *Manager) handlePiRetryEvent(dispatch *piRuntimeRun, event piRPCEvent) error { + var payload map[string]any + if err := json.Unmarshal(event.Raw, &payload); err != nil { + return fmt.Errorf("decode Pi %s: %w", event.Type, err) + } + text := "Pi retry state changed" + level := "info" + switch event.Type { + case "auto_retry_start": + text = fmt.Sprintf("Pi is retrying the model request (%d/%d)", int(numberValue(payload["attempt"])), int(numberValue(payload["maxAttempts"]))) + level = "warning" + case "auto_retry_end": + if boolValue(payload["success"]) { + text = "Pi model request retry succeeded" + dispatch.mu.Lock() + dispatch.lastAttemptError = "" + dispatch.mu.Unlock() + } else { + text = firstNonEmpty(strings.TrimSpace(stringValue(payload["finalError"])), "Pi model request retry failed") + level = "error" + dispatch.mu.Lock() + dispatch.lastAttemptError = text + dispatch.mu.Unlock() + } + case "summarization_retry_scheduled": + text = fmt.Sprintf("Pi scheduled a summarization retry (%d/%d)", int(numberValue(payload["attempt"])), int(numberValue(payload["maxAttempts"]))) + level = "warning" + case "summarization_retry_attempt_start": + text = "Pi is retrying conversation summarization" + case "summarization_retry_finished": + text = "Pi summarization retry finished" + } + extra := cloneMap(payload) + delete(extra, "type") + extra["code"] = "pi_" + event.Type + m.appendRunNote(dispatch.session.ID, dispatch.session, dispatch.run, level, text, extra) + return nil +} + +func (m *Manager) handlePiExtensionUIRequest(dispatch *piRuntimeRun, raw json.RawMessage) error { + var request struct { + ID string `json:"id"` + Method string `json:"method"` + Title string `json:"title"` + Message string `json:"message"` + Options []string `json:"options"` + Placeholder string `json:"placeholder"` + Prefill string `json:"prefill"` + Timeout int64 `json:"timeout"` + NotifyType string `json:"notifyType"` + StatusText string `json:"statusText"` + WidgetLines []string `json:"widgetLines"` + Text string `json:"text"` + } + if err := json.Unmarshal(raw, &request); err != nil { + return fmt.Errorf("decode Pi extension_ui_request: %w", err) + } + request.ID = strings.TrimSpace(request.ID) + request.Method = strings.TrimSpace(request.Method) + if request.ID == "" || request.Method == "" { + return errors.New("Pi extension UI request is missing id or method") + } + if request.Method != "select" && request.Method != "confirm" && request.Method != "input" && request.Method != "editor" { + text := firstNonEmpty(strings.TrimSpace(request.Message), strings.TrimSpace(request.StatusText), strings.Join(request.WidgetLines, "\n"), strings.TrimSpace(request.Text), strings.TrimSpace(request.Title)) + if text != "" { + level := strings.ToLower(strings.TrimSpace(request.NotifyType)) + if level != "warning" && level != "error" { + level = "info" + } + m.appendRunNote(dispatch.session.ID, dispatch.session, dispatch.run, level, truncateToolOutput("tool", text), map[string]any{"code": "pi_extension_ui_" + request.Method}) + } + return nil + } + + now := time.Now() + var expiresAt *time.Time + if request.Timeout > 0 { + value := now.Add(time.Duration(request.Timeout) * time.Millisecond) + expiresAt = &value + } + itemID := "pi-ui:" + request.ID + pending := &pendingServerRequest{ + ItemID: itemID, Prompt: firstNonEmpty(strings.TrimSpace(request.Title), strings.TrimSpace(request.Message), "Pi extension input"), + RequestedAt: &now, PiRuntime: dispatch.runtime, PiRequestID: request.ID, PiMethod: request.Method, + } + eventType := "user_input_req" + payload := map[string]any{"iid": itemID, "txt": pending.Prompt} + if request.Method == "confirm" { + pending.Kind = pendingServerRequestToolApproval + pending.Command = strings.TrimSpace(request.Message) + eventType = "approval_req" + payload["kind"] = string(pending.Kind) + payload["prompt"] = pending.Prompt + payload["command"] = pending.Command + } else { + pending.Kind = pendingServerRequestUserInput + question := toolRequestQuestion{ID: "value", Header: pending.Prompt, Question: pending.Prompt, IsOther: request.Method != "select"} + for _, option := range request.Options { + question.Options = append(question.Options, toolRequestOption{Label: option}) + } + pending.Questions = []toolRequestQuestion{question} + pending.Input = map[string]any{"placeholder": request.Placeholder, "prefill": request.Prefill} + payload["qs"] = pending.Questions + } + if _, exists := dispatch.run.pendingServerRequest(); exists { + return errors.New("Pi extension UI request arrived while another request is pending") + } + if !dispatch.run.setPendingServerRequest(pending) { + return errors.New("Pi extension UI request could not be registered") + } + dispatch.mu.Lock() + dispatch.dialog = &piRuntimeDialog{id: request.ID, method: request.Method, itemID: itemID, title: pending.Prompt, requested: now, expiresAt: expiresAt} + parentID := dispatch.assistantMessageID + dispatch.mu.Unlock() + + assistantState := AssistantStateWaitingInput + if pending.isApproval() { + assistantState = AssistantStateWaitingApproval + } + if err := m.updateRuntimeState(context.Background(), dispatch.session.ID, applyAssistantStateUpdates(map[string]any{"updated_at": now}, assistantState, now)); err != nil { + dispatch.run.clearPendingServerRequest() + return err + } + if _, err := m.appendAndBroadcast(context.Background(), dispatch.session.ID, dispatch.session, Event{ + ID: utils.NewID(), Type: eventType, RunID: dispatch.run.runID, ParentID: parentID, Timestamp: now, Payload: payload, + }); err != nil { + dispatch.run.clearPendingServerRequest() + return err + } + m.broadcastSessionSummary(context.Background(), dispatch.session.ID) + if expiresAt != nil { + delay := time.Until(*expiresAt) + time.AfterFunc(delay, func() { m.expirePiExtensionDialog(dispatch, request.ID) }) + } + return nil +} + +func (m *Manager) expirePiExtensionDialog(dispatch *piRuntimeRun, requestID string) { + if dispatch == nil || dispatch.run == nil { + return + } + pending, ok := dispatch.run.takePendingPiRequestForResponse(requestID, true) + if !ok || pending.PiRuntime != dispatch.runtime { + return + } + eventType := "user_input_res" + eventPayload := map[string]any{ + "iid": pending.ItemID, + "err": "Pi extension input timed out", + } + if pending.isApproval() { + eventType = "approval_res" + eventPayload = map[string]any{ + "iid": pending.ItemID, + "act": "cancel", + "prompt": pending.Prompt, + } + } + _ = m.respondPiExtensionRequest(dispatch.session, dispatch.run, pending, map[string]any{"cancelled": true}, eventType, eventPayload) +} + +func firstPiUserInputAnswer(answers map[string][]string) string { + for _, key := range []string{"value", "0"} { + for _, value := range answers[key] { + if normalized := strings.TrimSpace(value); normalized != "" { + return normalized + } + } + } + for _, values := range answers { + for _, value := range values { + if normalized := strings.TrimSpace(value); normalized != "" { + return normalized + } + } + } + return "" +} + +func piExtensionCancellationEvent(request *pendingServerRequest, reason string) (string, map[string]any) { + if request != nil && request.isApproval() { + return "approval_res", map[string]any{ + "iid": request.ItemID, + "act": "cancel", + "prompt": request.Prompt, + } + } + return "user_input_res", map[string]any{ + "iid": request.ItemID, + "err": firstNonEmpty(strings.TrimSpace(reason), "Pi extension input ended before a response"), + } +} + +func (m *Manager) appendPiExtensionCompletion( + session tables.WebSessionTable, + run *activeRun, + eventType string, + eventPayload map[string]any, +) error { + if run == nil { + return nil + } + if eventPayload == nil { + eventPayload = map[string]any{} + } + _, err := m.appendAndBroadcast(context.Background(), session.ID, session, Event{ + ID: utils.NewID(), Type: eventType, RunID: run.runID, ParentID: run.assistantMessageIDSnapshot(), + Timestamp: time.Now(), Payload: eventPayload, + }) + return err +} + +func (m *Manager) closePendingPiDialog(session tables.WebSessionTable, run *activeRun, reason string) error { + if run == nil { + return nil + } + waitCtx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + inFlight, waitErr := run.waitForPiResponseHistory(waitCtx) + cancel() + if waitErr != nil && inFlight == nil { + return waitErr + } + if inFlight != nil { + eventType, eventPayload := piExtensionCancellationEvent(inFlight, reason) + return m.appendPiExtensionCompletion(session, run, eventType, eventPayload) + } + pending, ok := run.pendingServerRequest() + if !ok || strings.TrimSpace(pending.PiRequestID) == "" { + return nil + } + request, taken := run.takePendingPiRequest(pending.PiRequestID) + if !taken { + return nil + } + eventType, eventPayload := piExtensionCancellationEvent(request, reason) + return m.appendPiExtensionCompletion(session, run, eventType, eventPayload) +} + +func (m *Manager) respondPiExtensionRequest( + session tables.WebSessionTable, + run *activeRun, + request *pendingServerRequest, + response map[string]any, + eventType string, + eventPayload map[string]any, +) error { + if run == nil || request == nil || request.PiRuntime == nil || strings.TrimSpace(request.PiRequestID) == "" { + return errors.New("Pi extension response channel is unavailable") + } + historyFinished := false + defer func() { + if !historyFinished { + run.finishPiResponseHistory(request.PiResponseGeneration, false, errors.New("Pi extension response ended before history was persisted")) + } + }() + runtime := request.PiRuntime + m.piRuntimeMu.Lock() + registered := m.piRuntimes[session.ID] + m.piRuntimeMu.Unlock() + if registered != runtime { + return errors.New("Pi extension response runtime is no longer active") + } + runtime.mu.Lock() + dispatch := runtime.active + valid := !runtime.stopped && dispatch != nil && dispatch.run == run + runtime.mu.Unlock() + if !valid { + return errors.New("Pi extension response run is no longer active") + } + now := time.Now() + if err := m.updateRuntimeState(context.Background(), session.ID, applyAssistantStateUpdates(map[string]any{"updated_at": now}, AssistantStateWorking, now)); err != nil { + runtime.stop(errors.New("Pi extension response state update failed")) + return err + } + dispatch.mu.Lock() + if dispatch.dialog != nil && dispatch.dialog.id == request.PiRequestID { + dispatch.dialog = nil + } + dispatch.mu.Unlock() + if err := m.appendPiExtensionCompletion(session, run, eventType, eventPayload); err != nil { + run.finishPiResponseHistory(request.PiResponseGeneration, false, err) + historyFinished = true + runtime.stop(errors.New("Pi extension response history update failed")) + return err + } + m.broadcastSessionSummary(context.Background(), session.ID) + + response["type"] = "extension_ui_response" + response["id"] = request.PiRequestID + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + err := runtime.client.Send(ctx, response) + cancel() + if err != nil { + run.finishPiResponseHistory(request.PiResponseGeneration, true, err) + historyFinished = true + runtime.stop(errors.New("Pi extension response failed")) + return err + } + run.finishPiResponseHistory(request.PiResponseGeneration, true, nil) + historyFinished = true + barrierCtx, barrierCancel := context.WithTimeout(context.Background(), 10*time.Second) + err = runtime.client.Barrier(barrierCtx, run.runID, request.PiResponseGeneration) + barrierCancel() + if err != nil { + runtime.stop(errors.New("Pi extension response barrier failed")) + return err + } + return nil +} + +func (m *Manager) finishPiSettledProjection(dispatch *piRuntimeRun) error { + m.clearPiNativeQueuedInputs(dispatch.session.ID) + dispatch.mu.Lock() + messageOpen := dispatch.assistantMessageOpen + messageID := dispatch.assistantMessageID + lastError := dispatch.lastAttemptError + compactionID := dispatch.compactionToolID + tools := make([]piRuntimeToolState, 0, len(dispatch.tools)) + for _, tool := range dispatch.tools { + if tool != nil && !tool.completed { + tools = append(tools, *tool) + } + } + dispatch.assistantMessageOpen = false + dispatch.compactionToolID = "" + dispatch.dialog = nil + dispatch.mu.Unlock() + + if messageOpen { + if _, err := m.appendAndBroadcast(context.Background(), dispatch.session.ID, dispatch.session, Event{ + ID: utils.NewID(), Type: "txt_end", RunID: dispatch.run.runID, ParentID: messageID, + Timestamp: time.Now(), Payload: map[string]any{"mid": messageID}, + }); err != nil { + return err + } + } + for _, tool := range tools { + if _, err := m.appendAndBroadcast(context.Background(), dispatch.session.ID, dispatch.session, Event{ + ID: utils.NewID(), Type: "tool_end", RunID: dispatch.run.runID, ParentID: tool.parentID, + Timestamp: time.Now(), Payload: map[string]any{ + "tid": tool.id, "name": firstNonEmpty(tool.name, "Tool"), "kind": "tool", + "in": tool.args, "out": tool.output, "ok": false, + }, + }); err != nil { + return err + } + } + if compactionID != "" { + _, _ = m.appendAndBroadcast(context.Background(), dispatch.session.ID, dispatch.session, Event{ + ID: utils.NewID(), Type: "tool_end", RunID: dispatch.run.runID, ParentID: messageID, + Timestamp: time.Now(), Payload: map[string]any{ + "tid": compactionID, "name": "ContextCompaction", "kind": "context_compaction", + "out": "Pi context compaction ended without a completion event", "ok": false, + }, + }) + } + if err := m.closePendingPiDialog(dispatch.session, dispatch.run, "Pi extension input ended before the run settled"); err != nil { + return err + } + if strings.TrimSpace(lastError) != "" { + return errors.New("Pi assistant run failed") + } + return nil +} + +func piMessageText(message piRPCMessage) string { + var builder strings.Builder + for _, block := range message.Content { + if block.Type == "text" { + builder.WriteString(block.Text) + } + } + return builder.String() +} + +func piToolResultText(value any) string { + record := decodeRawObject(value) + content, ok := record["content"].([]any) + if !ok { + if text := strings.TrimSpace(stringValue(record["text"])); text != "" { + return truncateToolOutput("tool", text) + } + encoded, _ := json.Marshal(value) + return truncateToolOutput("tool", string(encoded)) + } + parts := make([]string, 0, len(content)) + for _, item := range content { + block := decodeRawObject(item) + if strings.EqualFold(stringValue(block["type"]), "text") { + parts = append(parts, stringValue(block["text"])) + } + } + return truncateToolOutput("tool", strings.Join(parts, "\n")) +} diff --git a/service/websession/pi_rpc.go b/service/websession/pi_rpc.go new file mode 100644 index 00000000..e8298129 --- /dev/null +++ b/service/websession/pi_rpc.go @@ -0,0 +1,530 @@ +package websession + +import ( + "bufio" + "bytes" + "context" + "encoding/json" + "errors" + "fmt" + "io" + "os/exec" + "strings" + "sync" + "sync/atomic" + "time" +) + +const ( + piRPCMaxFrameBytes = 8 * 1024 * 1024 + piRPCStderrLimit = 64 * 1024 + piRPCRequestTimeout = 30 * time.Second +) + +var errPiRPCClosed = errors.New("Pi RPC process is closed") + +type piRPCRequestResult struct { + response piRPCResponse + err error +} + +type piRPCPendingRequest struct { + command string + result chan piRPCRequestResult +} + +type piRPCWrite struct { + data []byte + result chan error +} + +type piRPCClient struct { + cmd *exec.Cmd + stdin io.WriteCloser + stdout io.ReadCloser + stderr io.ReadCloser + + mu sync.Mutex + pending map[string]piRPCPendingRequest + exitErr error + events chan piRPCEvent + writes chan piRPCWrite + closing chan struct{} + done chan struct{} + readerDone chan struct{} + seq atomic.Uint64 + + stderrBuffer *piRPCBoundedBuffer + closeOnce sync.Once +} + +func startPiRPCClient(cmd *exec.Cmd) (*piRPCClient, error) { + if cmd == nil { + return nil, errors.New("Pi RPC command is nil") + } + stdin, err := cmd.StdinPipe() + if err != nil { + return nil, fmt.Errorf("open Pi RPC stdin: %w", err) + } + stdout, err := cmd.StdoutPipe() + if err != nil { + _ = stdin.Close() + return nil, fmt.Errorf("open Pi RPC stdout: %w", err) + } + stderr, err := cmd.StderrPipe() + if err != nil { + _ = stdin.Close() + _ = stdout.Close() + return nil, fmt.Errorf("open Pi RPC stderr: %w", err) + } + client := &piRPCClient{ + cmd: cmd, + stdin: stdin, + stdout: stdout, + stderr: stderr, + pending: make(map[string]piRPCPendingRequest), + events: make(chan piRPCEvent, 256), + writes: make(chan piRPCWrite, 64), + closing: make(chan struct{}), + done: make(chan struct{}), + readerDone: make(chan struct{}), + stderrBuffer: newPiRPCBoundedBuffer(piRPCStderrLimit), + } + if err := cmd.Start(); err != nil { + _ = stdin.Close() + _ = stdout.Close() + _ = stderr.Close() + return nil, fmt.Errorf("start Pi RPC process: %w", err) + } + go client.writeStdin() + go client.readStdout() + go func() { + _, _ = io.Copy(client.stderrBuffer, stderr) + }() + go client.wait() + return client, nil +} + +func (c *piRPCClient) Events() <-chan piRPCEvent { + if c == nil { + return nil + } + return c.events +} + +func (c *piRPCClient) Done() <-chan struct{} { + if c == nil { + closed := make(chan struct{}) + close(closed) + return closed + } + return c.done +} + +func (c *piRPCClient) Stderr() string { + if c == nil || c.stderrBuffer == nil { + return "" + } + return c.stderrBuffer.String() +} + +func (c *piRPCClient) Send(ctx context.Context, payload map[string]any) error { + if c == nil { + return errPiRPCClosed + } + if len(payload) == 0 || strings.TrimSpace(stringValue(payload["type"])) == "" { + return errors.New("Pi RPC frame type is required") + } + encoded, err := json.Marshal(payload) + if err != nil { + return fmt.Errorf("encode Pi RPC frame: %w", err) + } + encoded = append(encoded, '\n') + return c.enqueueWrite(ctx, encoded, "frame") +} + +func (c *piRPCClient) Barrier(ctx context.Context, runID string, generation uint64) error { + if generation == 0 { + return errors.New("Pi RPC response barrier generation is required") + } + return c.barrier(ctx, runID, generation, false, true) +} + +func (c *piRPCClient) BarrierAndWait(ctx context.Context, runID string) error { + return c.barrier(ctx, runID, 0, true, false) +} + +func (c *piRPCClient) barrier(ctx context.Context, runID string, generation uint64, wait, wakePending bool) error { + if c == nil { + return errPiRPCClosed + } + if strings.TrimSpace(runID) == "" { + return errors.New("Pi RPC barrier run id is required") + } + if err := c.Request(ctx, "get_state", nil, nil); err != nil { + return err + } + var consumed chan struct{} + if wait { + consumed = make(chan struct{}) + } + marker := piRPCEvent{ + Type: "codekanban_barrier", BarrierRunID: runID, BarrierGeneration: generation, + BarrierDone: consumed, WakePending: wakePending, + } + select { + case c.events <- marker: + case <-ctx.Done(): + err := fmt.Errorf("queue Pi RPC barrier: %w", ctx.Err()) + c.fail(err) + killCmdTree(c.cmd) + return err + case <-c.done: + return c.processError() + } + if consumed == nil { + return nil + } + select { + case <-consumed: + return nil + case <-ctx.Done(): + return fmt.Errorf("wait for Pi RPC barrier: %w", ctx.Err()) + case <-c.done: + return c.processError() + } +} + +func (c *piRPCClient) Request(ctx context.Context, command string, payload map[string]any, target any) error { + if c == nil { + return errPiRPCClosed + } + command = strings.TrimSpace(command) + if command == "" { + return errors.New("Pi RPC command type is required") + } + id := fmt.Sprintf("ck_%d", c.seq.Add(1)) + request := make(map[string]any, len(payload)+2) + for key, value := range payload { + if key != "id" && key != "type" { + request[key] = value + } + } + request["id"] = id + request["type"] = command + encoded, err := json.Marshal(request) + if err != nil { + return fmt.Errorf("encode Pi RPC %s request: %w", command, err) + } + encoded = append(encoded, '\n') + + resultCh := make(chan piRPCRequestResult, 1) + c.mu.Lock() + if c.exitErr != nil { + err := c.exitErr + c.mu.Unlock() + return err + } + c.pending[id] = piRPCPendingRequest{command: command, result: resultCh} + c.mu.Unlock() + + requestCtx := ctx + if _, hasDeadline := ctx.Deadline(); !hasDeadline { + var cancel context.CancelFunc + requestCtx, cancel = context.WithTimeout(ctx, piRPCRequestTimeout) + defer cancel() + } + if err := c.enqueueWrite(requestCtx, encoded, command+" request"); err != nil { + c.removePending(id) + return err + } + + select { + case result := <-resultCh: + if result.err != nil { + return result.err + } + if !result.response.Success { + message := strings.TrimSpace(result.response.Error) + if message == "" { + message = "request failed" + } + return fmt.Errorf("Pi RPC %s failed: %s", command, message) + } + if target == nil || len(result.response.Data) == 0 || bytes.Equal(result.response.Data, []byte("null")) { + return nil + } + if err := json.Unmarshal(result.response.Data, target); err != nil { + return fmt.Errorf("decode Pi RPC %s response: %w", command, err) + } + return nil + case <-requestCtx.Done(): + err := fmt.Errorf("Pi RPC %s request: %w", command, requestCtx.Err()) + c.removePending(id) + c.fail(err) + killCmdTree(c.cmd) + return err + case <-c.done: + c.removePending(id) + return c.processError() + } +} + +func (c *piRPCClient) enqueueWrite(ctx context.Context, encoded []byte, description string) error { + if ctx == nil { + ctx = context.Background() + } + writeCtx := ctx + if _, hasDeadline := ctx.Deadline(); !hasDeadline { + var cancel context.CancelFunc + writeCtx, cancel = context.WithTimeout(ctx, piRPCRequestTimeout) + defer cancel() + } + writeResult := make(chan error, 1) + select { + case c.writes <- piRPCWrite{data: encoded, result: writeResult}: + case <-writeCtx.Done(): + err := fmt.Errorf("queue Pi RPC %s: %w", description, writeCtx.Err()) + c.fail(err) + killCmdTree(c.cmd) + return err + case <-c.closing: + return c.processError() + case <-c.done: + return c.processError() + } + select { + case writeErr := <-writeResult: + if writeErr != nil { + return fmt.Errorf("write Pi RPC %s: %w", description, writeErr) + } + return nil + case <-writeCtx.Done(): + err := fmt.Errorf("write Pi RPC %s: %w", description, writeCtx.Err()) + c.fail(err) + killCmdTree(c.cmd) + return err + case <-c.done: + return c.processError() + } +} + +func (c *piRPCClient) Close() error { + if c == nil { + return nil + } + c.closeOnce.Do(func() { + close(c.closing) + _ = c.stdin.Close() + select { + case <-c.done: + return + case <-time.After(500 * time.Millisecond): + } + killCmdTree(c.cmd) + }) + select { + case <-c.done: + return nil + case <-time.After(2 * time.Second): + return errors.New("timed out stopping Pi RPC process") + } +} + +func (c *piRPCClient) writeStdin() { + for { + select { + case write := <-c.writes: + _, err := c.stdin.Write(write.data) + write.result <- err + if err != nil { + c.fail(fmt.Errorf("write Pi RPC stdin: %w", err)) + killCmdTree(c.cmd) + return + } + case <-c.closing: + return + case <-c.done: + return + } + } +} + +func (c *piRPCClient) readStdout() { + defer close(c.readerDone) + reader := bufio.NewReaderSize(c.stdout, 64*1024) + for { + line, err := readPiRPCJSONLFrame(reader, piRPCMaxFrameBytes) + if err != nil { + if errors.Is(err, io.EOF) { + return + } + c.fail(fmt.Errorf("read Pi RPC stdout: %w", err)) + killCmdTree(c.cmd) + return + } + if err := c.handleLine(line); err != nil { + c.fail(err) + killCmdTree(c.cmd) + return + } + } +} + +func readPiRPCJSONLFrame(reader *bufio.Reader, limit int) ([]byte, error) { + if limit <= 0 { + limit = piRPCMaxFrameBytes + } + frame := make([]byte, 0, 1024) + for { + part, err := reader.ReadSlice('\n') + if len(frame)+len(part) > limit { + return nil, fmt.Errorf("frame exceeds %d bytes", limit) + } + frame = append(frame, part...) + switch { + case err == nil: + frame = frame[:len(frame)-1] + if len(frame) > 0 && frame[len(frame)-1] == '\r' { + frame = frame[:len(frame)-1] + } + return frame, nil + case errors.Is(err, bufio.ErrBufferFull): + continue + case errors.Is(err, io.EOF): + if len(frame) == 0 { + return nil, io.EOF + } + return nil, errors.New("EOF after partial JSONL frame") + default: + return nil, err + } + } +} + +func (c *piRPCClient) handleLine(line []byte) error { + var envelope struct { + Type string `json:"type"` + } + if len(bytes.TrimSpace(line)) == 0 { + return errors.New("Pi RPC emitted an empty JSONL frame") + } + if err := json.Unmarshal(line, &envelope); err != nil { + return fmt.Errorf("Pi RPC emitted malformed JSON: %w", err) + } + if strings.TrimSpace(envelope.Type) == "" { + return errors.New("Pi RPC frame has no type") + } + if envelope.Type != "response" { + event := piRPCEvent{Type: envelope.Type, Raw: append(json.RawMessage(nil), line...)} + select { + case c.events <- event: + return nil + default: + return errors.New("Pi RPC event buffer overflow") + } + } + + var response piRPCResponse + if err := json.Unmarshal(line, &response); err != nil { + return fmt.Errorf("decode Pi RPC response: %w", err) + } + if strings.TrimSpace(response.ID) == "" || strings.TrimSpace(response.Command) == "" { + return errors.New("Pi RPC response is missing id or command") + } + c.mu.Lock() + pending, ok := c.pending[response.ID] + c.mu.Unlock() + if !ok { + return fmt.Errorf("Pi RPC response has unknown id %q", response.ID) + } + if pending.command != response.Command { + return fmt.Errorf( + "Pi RPC response command mismatch: got %s, want %s", + response.Command, + pending.command, + ) + } + c.removePending(response.ID) + pending.result <- piRPCRequestResult{response: response} + return nil +} + +func (c *piRPCClient) wait() { + err := c.cmd.Wait() + <-c.readerDone + if err != nil { + c.fail(fmt.Errorf("Pi RPC process exited: %w", err)) + } else { + c.fail(errPiRPCClosed) + } + close(c.events) + close(c.done) +} + +func (c *piRPCClient) fail(err error) { + if err == nil { + err = errPiRPCClosed + } + c.mu.Lock() + if c.exitErr == nil { + c.exitErr = err + } + resolved := c.exitErr + pending := c.pending + c.pending = make(map[string]piRPCPendingRequest) + c.mu.Unlock() + for _, request := range pending { + request.result <- piRPCRequestResult{err: resolved} + } +} + +func (c *piRPCClient) removePending(id string) { + c.mu.Lock() + delete(c.pending, id) + c.mu.Unlock() +} + +func (c *piRPCClient) processError() error { + c.mu.Lock() + defer c.mu.Unlock() + if c.exitErr != nil { + return c.exitErr + } + return errPiRPCClosed +} + +type piRPCBoundedBuffer struct { + mu sync.Mutex + limit int + data []byte +} + +func newPiRPCBoundedBuffer(limit int) *piRPCBoundedBuffer { + return &piRPCBoundedBuffer{limit: limit} +} + +func (b *piRPCBoundedBuffer) Write(data []byte) (int, error) { + b.mu.Lock() + defer b.mu.Unlock() + originalLength := len(data) + if b.limit <= 0 { + return originalLength, nil + } + if len(data) >= b.limit { + b.data = append(b.data[:0], data[len(data)-b.limit:]...) + return originalLength, nil + } + if overflow := len(b.data) + len(data) - b.limit; overflow > 0 { + copy(b.data, b.data[overflow:]) + b.data = b.data[:len(b.data)-overflow] + } + b.data = append(b.data, data...) + return originalLength, nil +} + +func (b *piRPCBoundedBuffer) String() string { + b.mu.Lock() + defer b.mu.Unlock() + return string(append([]byte(nil), b.data...)) +} diff --git a/service/websession/pi_rpc_test.go b/service/websession/pi_rpc_test.go new file mode 100644 index 00000000..282e4329 --- /dev/null +++ b/service/websession/pi_rpc_test.go @@ -0,0 +1,269 @@ +package websession + +import ( + "bufio" + "bytes" + "context" + "encoding/json" + "errors" + "fmt" + "io" + "os" + "os/exec" + "strings" + "sync" + "testing" + "time" +) + +func TestPiRPCFakeProcess(t *testing.T) { + if os.Getenv("CODEKANBAN_PI_RPC_FAKE") != "1" { + return + } + mode := os.Getenv("CODEKANBAN_PI_RPC_FAKE_MODE") + scanner := bufio.NewScanner(os.Stdin) + encoder := json.NewEncoder(os.Stdout) + switch mode { + case "ordered": + commands := make([]map[string]any, 0, 2) + for scanner.Scan() { + var command map[string]any + if json.Unmarshal(scanner.Bytes(), &command) != nil { + os.Exit(2) + } + commands = append(commands, command) + if len(commands) != 2 { + continue + } + _ = encoder.Encode(map[string]any{"type": "agent_start"}) + for index := len(commands) - 1; index >= 0; index-- { + item := commands[index] + _ = encoder.Encode(map[string]any{ + "type": "response", "id": item["id"], "command": item["type"], + "success": true, "data": map[string]any{"value": item["type"]}, + }) + } + } + case "malformed": + if scanner.Scan() { + _, _ = fmt.Fprintln(os.Stdout, `{not-json`) + } + case "partial": + if scanner.Scan() { + _, _ = io.WriteString(os.Stdout, `{"type":"response"`) + } + case "unknown", "mismatch": + if scanner.Scan() { + var command map[string]any + _ = json.Unmarshal(scanner.Bytes(), &command) + if mode == "unknown" { + command["id"] = "unknown-request" + } else { + command["type"] = "different_command" + } + _ = encoder.Encode(map[string]any{ + "type": "response", "id": command["id"], "command": command["type"], + "success": true, "data": map[string]any{"ok": true}, + }) + } + case "barrier": + for scanner.Scan() { + var command map[string]any + if json.Unmarshal(scanner.Bytes(), &command) != nil { + os.Exit(2) + } + if command["type"] != "get_state" { + continue + } + _ = encoder.Encode(map[string]any{"type": "extension_ui_request", "id": "next-dialog", "method": "input"}) + _ = encoder.Encode(map[string]any{ + "type": "response", "id": command["id"], "command": command["type"], + "success": true, "data": map[string]any{"isStreaming": true}, + }) + } + case "stderr": + if scanner.Scan() { + _, _ = io.WriteString(os.Stderr, strings.Repeat("x", piRPCStderrLimit*2)) + var command map[string]any + _ = json.Unmarshal(scanner.Bytes(), &command) + _ = encoder.Encode(map[string]any{ + "type": "response", "id": command["id"], "command": command["type"], + "success": true, "data": map[string]any{"ok": true}, + }) + } + case "exit": + if scanner.Scan() { + os.Exit(7) + } + case "hang": + if scanner.Scan() { + time.Sleep(30 * time.Second) + } + default: + os.Exit(3) + } + os.Exit(0) +} + +func fakePiRPCCommand(t *testing.T, mode string) *exec.Cmd { + t.Helper() + cmd := exec.Command(os.Args[0], "-test.run=^TestPiRPCFakeProcess$") + cmd.Env = append(os.Environ(), + "CODEKANBAN_PI_RPC_FAKE=1", + "CODEKANBAN_PI_RPC_FAKE_MODE="+mode, + ) + return cmd +} + +func TestReadPiRPCJSONLFrameBoundaries(t *testing.T) { + payload := "{\"value\":\"line\u2028separator\u2029ok\"}\r\n{\"next\":true}\n" + reader := bufio.NewReaderSize(strings.NewReader(payload), 8) + first, err := readPiRPCJSONLFrame(reader, 1024) + if err != nil { + t.Fatal(err) + } + if string(first) != "{\"value\":\"line\u2028separator\u2029ok\"}" { + t.Fatalf("unexpected first frame: %q", first) + } + second, err := readPiRPCJSONLFrame(reader, 1024) + if err != nil || string(second) != `{"next":true}` { + t.Fatalf("unexpected second frame %q: %v", second, err) + } + if _, err := readPiRPCJSONLFrame(reader, 1024); !errors.Is(err, io.EOF) { + t.Fatalf("expected EOF, got %v", err) + } +} + +func TestReadPiRPCJSONLFrameRejectsPartialAndOversized(t *testing.T) { + if _, err := readPiRPCJSONLFrame(bufio.NewReader(strings.NewReader(`{"type":"event"}`)), 1024); err == nil || !strings.Contains(err.Error(), "partial") { + t.Fatalf("expected partial frame error, got %v", err) + } + if _, err := readPiRPCJSONLFrame(bufio.NewReader(bytes.NewReader(append(bytes.Repeat([]byte("x"), 20), '\n'))), 10); err == nil || !strings.Contains(err.Error(), "exceeds") { + t.Fatalf("expected oversized frame error, got %v", err) + } +} + +func TestPiRPCClientCorrelatesOutOfOrderResponsesAndEvents(t *testing.T) { + client, err := startPiRPCClient(fakePiRPCCommand(t, "ordered")) + if err != nil { + t.Fatal(err) + } + defer client.Close() + + type response struct { + Value string `json:"value"` + } + results := make(map[string]string) + var mu sync.Mutex + var wg sync.WaitGroup + for _, command := range []string{"first", "second"} { + command := command + wg.Add(1) + go func() { + defer wg.Done() + var decoded response + if err := client.Request(context.Background(), command, nil, &decoded); err != nil { + t.Errorf("request %s: %v", command, err) + return + } + mu.Lock() + results[command] = decoded.Value + mu.Unlock() + }() + } + wg.Wait() + if results["first"] != "first" || results["second"] != "second" { + t.Fatalf("unexpected correlated results: %#v", results) + } + select { + case event := <-client.Events(): + if event.Type != "agent_start" { + t.Fatalf("unexpected event: %#v", event) + } + case <-time.After(time.Second): + t.Fatal("timed out waiting for interleaved event") + } +} + +func TestPiRPCClientBarrierFollowsEarlierRuntimeEvents(t *testing.T) { + client, err := startPiRPCClient(fakePiRPCCommand(t, "barrier")) + if err != nil { + t.Fatal(err) + } + defer client.Close() + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + if err := client.Send(ctx, map[string]any{"type": "extension_ui_response", "id": "previous-dialog", "value": "ok"}); err != nil { + t.Fatal(err) + } + if err := client.Barrier(ctx, "run-1", 1); err != nil { + t.Fatal(err) + } + first := <-client.Events() + second := <-client.Events() + if first.Type != "extension_ui_request" || second.Type != "codekanban_barrier" || second.BarrierRunID != "run-1" || second.BarrierGeneration != 1 { + t.Fatalf("unexpected barrier ordering: first=%#v second=%#v", first, second) + } +} + +func TestPiRPCClientRejectsMalformedAndPartialFrames(t *testing.T) { + for _, mode := range []string{"malformed", "partial", "unknown", "mismatch"} { + t.Run(mode, func(t *testing.T) { + client, err := startPiRPCClient(fakePiRPCCommand(t, mode)) + if err != nil { + t.Fatal(err) + } + defer client.Close() + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + err = client.Request(ctx, "get_state", nil, nil) + if err == nil { + t.Fatal("expected protocol error") + } + }) + } +} + +func TestPiRPCClientTimeoutTerminatesProcess(t *testing.T) { + client, err := startPiRPCClient(fakePiRPCCommand(t, "hang")) + if err != nil { + t.Fatal(err) + } + defer client.Close() + ctx, cancel := context.WithTimeout(context.Background(), 50*time.Millisecond) + defer cancel() + if err := client.Request(ctx, "get_state", nil, nil); err == nil || !errors.Is(err, context.DeadlineExceeded) { + t.Fatalf("expected request deadline error, got %v", err) + } + select { + case <-client.Done(): + case <-time.After(2 * time.Second): + t.Fatal("timed out waiting for hung Pi process to terminate") + } +} + +func TestPiRPCClientBoundsStderrAndRejectsEarlyExit(t *testing.T) { + client, err := startPiRPCClient(fakePiRPCCommand(t, "stderr")) + if err != nil { + t.Fatal(err) + } + var response struct { + OK bool `json:"ok"` + } + if err := client.Request(context.Background(), "get_state", nil, &response); err != nil { + t.Fatal(err) + } + _ = client.Close() + if !response.OK || len(client.Stderr()) != piRPCStderrLimit { + t.Fatalf("response=%#v stderr length=%d", response, len(client.Stderr())) + } + + exiting, err := startPiRPCClient(fakePiRPCCommand(t, "exit")) + if err != nil { + t.Fatal(err) + } + defer exiting.Close() + if err := exiting.Request(context.Background(), "get_state", nil, nil); err == nil || !strings.Contains(err.Error(), "exited") { + t.Fatalf("expected process exit error, got %v", err) + } +} diff --git a/service/websession/pi_rpc_types.go b/service/websession/pi_rpc_types.go new file mode 100644 index 00000000..ed7f9920 --- /dev/null +++ b/service/websession/pi_rpc_types.go @@ -0,0 +1,82 @@ +package websession + +import "encoding/json" + +type piRPCResponse struct { + Type string `json:"type"` + ID string `json:"id"` + Command string `json:"command"` + Success bool `json:"success"` + Data json.RawMessage `json:"data"` + Error string `json:"error"` +} + +type piRPCEvent struct { + Type string `json:"type"` + Raw json.RawMessage `json:"-"` + BarrierRunID string `json:"-"` + BarrierGeneration uint64 `json:"-"` + BarrierDone chan struct{} `json:"-"` + WakePending bool `json:"-"` +} + +type piRPCState struct { + Model *struct { + Provider string `json:"provider"` + ID string `json:"id"` + Name string `json:"name"` + } `json:"model"` + ThinkingLevel string `json:"thinkingLevel"` + IsStreaming bool `json:"isStreaming"` + SessionID string `json:"sessionId"` + SessionFile string `json:"sessionFile"` + SessionName string `json:"sessionName"` +} + +type piRPCModel struct { + Provider string `json:"provider"` + ID string `json:"id"` + Name string `json:"name"` + Reasoning bool `json:"reasoning"` + Input []string `json:"input"` + ContextWindow int64 `json:"contextWindow"` + MaxTokens int64 `json:"maxTokens"` +} + +type piRPCAvailableModels struct { + Models []piRPCModel `json:"models"` +} + +type piRPCAvailableThinkingLevels struct { + Levels []string `json:"levels"` +} + +type piRPCSetModelResult struct { + Provider string `json:"provider"` + ID string `json:"id"` + Name string `json:"name"` +} + +type piRPCImage struct { + Type string `json:"type"` + Data string `json:"data"` + MimeType string `json:"mimeType"` +} + +type piRPCSessionStats struct { + Tokens struct { + Input int64 `json:"input"` + Output int64 `json:"output"` + CacheRead int64 `json:"cacheRead"` + CacheWrite int64 `json:"cacheWrite"` + Total int64 `json:"total"` + } `json:"tokens"` + Cost float64 `json:"cost"` + SessionID string `json:"sessionId"` + SessionFile string `json:"sessionFile"` + ContextUsage *struct { + Tokens int64 `json:"tokens"` + ContextWindow int64 `json:"contextWindow"` + Percent float64 `json:"percent"` + } `json:"contextUsage"` +} diff --git a/service/websession/pi_runtime.go b/service/websession/pi_runtime.go new file mode 100644 index 00000000..bd2c957e --- /dev/null +++ b/service/websession/pi_runtime.go @@ -0,0 +1,982 @@ +package websession + +import ( + "bufio" + "context" + "encoding/base64" + "encoding/json" + "errors" + "fmt" + "net/http" + "os" + "path/filepath" + "runtime" + "strings" + "sync" + "time" + + "code-kanban/model/tables" + "code-kanban/utils" + "code-kanban/utils/ai_assistant2/log_watcher" + + "go.uber.org/zap" +) + +const piSessionHeaderLimit = 1024 * 1024 + +type piRuntimeContentState struct { + kind string + text string + toolID string +} + +type piRuntimeToolState struct { + id string + name string + args any + output string + parentID string + lastEmit time.Time + completed bool +} + +type piRuntimeDialog struct { + id string + method string + itemID string + title string + requested time.Time + expiresAt *time.Time +} + +type piRuntimeRun struct { + runtime *piSessionRuntime + run *activeRun + session tables.WebSessionTable + settled chan error + once sync.Once + + mu sync.Mutex + assistantMessageID string + assistantMessageOpen bool + contents map[int]*piRuntimeContentState + tools map[string]*piRuntimeToolState + dialog *piRuntimeDialog + compactionToolID string + compactionStarted time.Time + compactionCompleted time.Time + lastAttemptError string +} + +func (r *piRuntimeRun) settle(err error) { + if r == nil { + return + } + r.once.Do(func() { + r.settled <- err + close(r.settled) + }) +} + +type piSessionRuntime struct { + manager *Manager + client *piRPCClient + sessionID string + projectID string + cwd string + + mu sync.Mutex + active *piRuntimeRun + idleTimer *time.Timer + stopped bool + stopOnce sync.Once +} + +func (m *Manager) getOrStartPiRuntime( + ctx context.Context, + session tables.WebSessionTable, +) (*piSessionRuntime, error) { + if normalizeAgent(Agent(session.Agent)) != AgentPi || effectiveSessionBackend(session) != SessionBackendPiRPC { + return nil, errors.New("Pi RPC runtime requires a Pi web session") + } + if err := m.EnsureProjectPiTrust(ctx, session.ProjectID, session.Cwd); err != nil { + return nil, err + } + + m.piRuntimeMu.Lock() + existing := m.piRuntimes[session.ID] + m.piRuntimeMu.Unlock() + if existing != nil && existing.acquire() { + return existing, nil + } + + created, err := m.startPiRuntime(ctx, session) + if err != nil { + return nil, err + } + m.piRuntimeMu.Lock() + if existing = m.piRuntimes[session.ID]; existing == nil { + m.piRuntimes[session.ID] = created + m.piRuntimeTerminators[session.ID] = piRuntimeTerminator{ + projectID: session.ProjectID, + terminate: func() { created.stop(errors.New("Pi RPC runtime terminated")) }, + } + m.piRuntimeMu.Unlock() + go created.consumeEvents() + return created, nil + } + m.piRuntimeMu.Unlock() + created.stop(errors.New("duplicate Pi RPC runtime")) + if existing.acquire() { + return existing, nil + } + return nil, errors.New("Pi RPC runtime stopped during startup") +} + +func (m *Manager) startPiRuntime( + ctx context.Context, + session tables.WebSessionTable, +) (*piSessionRuntime, error) { + if !m.GetWebSessionRuntimeConfig().SupportsPiWebSession { + return nil, errors.New(errPiWebSessionUnavailable) + } + bridgePath, err := m.materializePiBridge() + if err != nil { + return nil, err + } + args := make([]string, 0, 4) + threadPath := pointerString(session.ThreadPath) + if threadPath != "" { + args = append(args, "--session", threadPath) + } else { + args = append(args, "--name", session.Title) + } + // The reusable RPC process must outlive the active run that triggered its start. + cmd, err := m.buildTrustedPiRPCCommand(context.Background(), session.ProjectID, session.Cwd, args...) + if err != nil { + return nil, err + } + client, err := startPiRPCClient(cmd) + if err != nil { + return nil, err + } + failed := true + defer func() { + if failed { + _ = client.Close() + } + }() + + requestCtx, cancel := context.WithTimeout(ctx, 10*time.Second) + defer cancel() + var state piRPCState + if err := client.Request(requestCtx, "get_state", nil, &state); err != nil { + return nil, err + } + if err := validatePiRuntimeStartupState(session, state, threadPath == ""); err != nil { + return nil, err + } + var commands struct { + Commands []piRPCSlashCommand `json:"commands"` + } + if err := client.Request(requestCtx, "get_commands", nil, &commands); err != nil { + return nil, err + } + if err := validatePiBridgeCommands(commands.Commands, bridgePath); err != nil { + return nil, err + } + + var entries struct { + LeafID *string `json:"leafId"` + } + if err := client.Request(requestCtx, "get_entries", nil, &entries); err != nil { + return nil, err + } + modelName := canonicalPiModel(state.Model) + updates := map[string]any{ + "source_kind": string(SessionBackendPiRPC), + "sync_error": nil, + "updated_at": time.Now(), + } + if _, statErr := os.Stat(state.SessionFile); statErr == nil { + updates["native_session_id"] = strings.TrimSpace(state.SessionID) + updates["thread_path"] = filepath.Clean(state.SessionFile) + updates["native_leaf_id"] = nilIfEmpty(pointerString(entries.LeafID)) + updates["source_revision"] = nilIfEmpty(piSourceRevision(state.SessionFile, pointerString(entries.LeafID))) + updates["sync_state"] = string(SyncStateFresh) + } + if modelName != "" { + updates["model"] = modelName + } + if normalized := piThinkingLevelToReasoning(state.ThinkingLevel); normalized != ReasoningEffortDefault { + updates["reasoning_effort"] = string(normalized) + } + if err := m.updateRuntimeState(ctx, session.ID, updates); err != nil { + return nil, err + } + + failed = false + return &piSessionRuntime{ + manager: m, + client: client, + sessionID: session.ID, + projectID: session.ProjectID, + cwd: session.Cwd, + }, nil +} + +func validatePiRuntimeState(session tables.WebSessionTable, state piRPCState) error { + return validatePiRuntimeStartupState(session, state, false) +} + +func validatePiRuntimeStartupState(session tables.WebSessionTable, state piRPCState, allowMissingNewFile bool) error { + sessionID := strings.TrimSpace(state.SessionID) + sessionFile := strings.TrimSpace(state.SessionFile) + if sessionID == "" || sessionFile == "" || !filepath.IsAbs(sessionFile) { + return errors.New("Pi RPC get_state returned an invalid session identity") + } + if expected := strings.TrimSpace(pointerString(session.NativeSessionID)); expected != "" && expected != sessionID { + return fmt.Errorf("Pi session id mismatch: got %q", sessionID) + } + if expected := strings.TrimSpace(pointerString(session.ThreadPath)); expected != "" && !samePiRuntimePath(expected, sessionFile) { + return errors.New("Pi session file mismatch") + } + if err := validatePiSessionRoot(sessionFile); err != nil { + return err + } + header, err := readPiRuntimeSessionHeader(sessionFile) + if err != nil { + if allowMissingNewFile && os.IsNotExist(err) && strings.TrimSpace(pointerString(session.NativeSessionID)) == "" && strings.TrimSpace(pointerString(session.ThreadPath)) == "" { + return nil + } + return fmt.Errorf("verify Pi session header: %w", err) + } + if header.ID != sessionID { + return errors.New("Pi session header id does not match get_state") + } + if !samePiRuntimePath(header.Cwd, session.Cwd) { + return errors.New("Pi session cwd does not match the web session") + } + return nil +} + +type piRuntimeSessionHeader struct { + Type string `json:"type"` + ID string `json:"id"` + Cwd string `json:"cwd"` +} + +func readPiRuntimeSessionHeader(path string) (piRuntimeSessionHeader, error) { + file, err := os.Open(path) + if err != nil { + return piRuntimeSessionHeader{}, err + } + defer file.Close() + scanner := bufio.NewScanner(file) + scanner.Buffer(make([]byte, 0, 64*1024), piSessionHeaderLimit) + if !scanner.Scan() { + if err := scanner.Err(); err != nil { + return piRuntimeSessionHeader{}, err + } + return piRuntimeSessionHeader{}, errors.New("empty Pi session file") + } + var header piRuntimeSessionHeader + if err := json.Unmarshal(scanner.Bytes(), &header); err != nil { + return piRuntimeSessionHeader{}, err + } + header.Type = strings.TrimSpace(header.Type) + header.ID = strings.TrimSpace(header.ID) + header.Cwd = strings.TrimSpace(header.Cwd) + if header.Type != "session" || header.ID == "" || header.Cwd == "" { + return piRuntimeSessionHeader{}, errors.New("invalid Pi session header") + } + return header, nil +} + +func canonicalPiRuntimePath(value string) (string, error) { + value = strings.TrimSpace(value) + absolute, err := filepath.Abs(value) + if err != nil { + return "", err + } + value = filepath.Clean(absolute) + unresolved := make([]string, 0, 2) + current := value + for { + if _, statErr := os.Lstat(current); statErr == nil { + if resolved, resolveErr := filepath.EvalSymlinks(current); resolveErr == nil { + current = filepath.Clean(resolved) + } + break + } else if !os.IsNotExist(statErr) { + return "", statErr + } + parent := filepath.Dir(current) + if parent == current { + break + } + unresolved = append(unresolved, filepath.Base(current)) + current = parent + } + for index := len(unresolved) - 1; index >= 0; index-- { + current = filepath.Join(current, unresolved[index]) + } + value = filepath.Clean(current) + if runtime.GOOS == "windows" { + value = strings.ToLower(value) + } + return value, nil +} + +func samePiRuntimePath(left, right string) bool { + canonicalLeft, leftErr := canonicalPiRuntimePath(left) + canonicalRight, rightErr := canonicalPiRuntimePath(right) + return leftErr == nil && rightErr == nil && canonicalLeft == canonicalRight +} + +func validatePiSessionRoot(sessionFile string) error { + root, err := log_watcher.ResolvePiSessionDir() + if err != nil { + return fmt.Errorf("resolve Pi session root: %w", err) + } + canonicalRoot, err := canonicalPiRuntimePath(root) + if err != nil { + return fmt.Errorf("resolve Pi session root: %w", err) + } + canonicalFile, err := canonicalPiRuntimePath(sessionFile) + if err != nil { + return fmt.Errorf("resolve Pi session file: %w", err) + } + relative, err := filepath.Rel(canonicalRoot, canonicalFile) + if err != nil || relative == "." || filepath.IsAbs(relative) || relative == ".." || strings.HasPrefix(relative, ".."+string(filepath.Separator)) { + return errors.New("Pi session file is outside the configured session root") + } + return nil +} + +func piSourceRevision(path, leafID string) string { + info, err := os.Stat(path) + if err != nil { + return "" + } + return fmt.Sprintf("%d:%d:%s", info.ModTime().UnixNano(), info.Size(), strings.TrimSpace(leafID)) +} + +func canonicalPiModel(model *struct { + Provider string `json:"provider"` + ID string `json:"id"` + Name string `json:"name"` +}) string { + if model == nil { + return "" + } + provider := strings.Trim(strings.TrimSpace(model.Provider), "/") + id := strings.Trim(strings.TrimSpace(model.ID), "/") + if provider == "" || id == "" { + return "" + } + return provider + "/" + id +} + +func splitPiModel(value string) (string, string, error) { + parts := strings.SplitN(strings.TrimSpace(value), "/", 2) + if len(parts) != 2 || strings.TrimSpace(parts[0]) == "" || strings.TrimSpace(parts[1]) == "" { + return "", "", errors.New("Pi model must use provider/modelId") + } + return strings.TrimSpace(parts[0]), strings.TrimSpace(parts[1]), nil +} + +func validatePiReasoningEffort(effort ReasoningEffort) error { + normalized := normalizeReasoningEffort(effort) + if normalized != ReasoningEffortDefault && piReasoningToThinkingLevel(normalized) == "" { + return fmt.Errorf("Pi does not support reasoning effort %q", normalized) + } + return nil +} + +func piReasoningToThinkingLevel(effort ReasoningEffort) string { + switch normalizeReasoningEffort(effort) { + case ReasoningEffortNone: + return "off" + case ReasoningEffortMinimal: + return "minimal" + case ReasoningEffortLow: + return "low" + case ReasoningEffortMedium: + return "medium" + case ReasoningEffortHigh: + return "high" + case ReasoningEffortXHigh: + return "xhigh" + case ReasoningEffortMax: + return "max" + default: + return "" + } +} + +func piThinkingLevelToReasoning(level string) ReasoningEffort { + switch strings.ToLower(strings.TrimSpace(level)) { + case "off": + return ReasoningEffortNone + case "minimal": + return ReasoningEffortMinimal + case "low": + return ReasoningEffortLow + case "medium": + return ReasoningEffortMedium + case "high": + return ReasoningEffortHigh + case "xhigh": + return ReasoningEffortXHigh + case "max": + return ReasoningEffortMax + default: + return ReasoningEffortDefault + } +} + +func (m *Manager) piPromptImages(attachments []Attachment) ([]piRPCImage, error) { + images := make([]piRPCImage, 0, len(attachments)) + attachmentRoot, err := canonicalPiRuntimePath(m.store.attachmentsDir) + if err != nil { + return nil, fmt.Errorf("resolve attachment root: %w", err) + } + for _, attachment := range attachments { + declaredMime := strings.ToLower(strings.TrimSpace(attachment.Mime)) + if !strings.HasPrefix(declaredMime, "image/") { + return nil, fmt.Errorf("Pi only supports image attachments; %s is %s", attachment.Name, attachment.Mime) + } + attachmentPath, err := canonicalPiRuntimePath(attachment.Path) + if err != nil { + return nil, fmt.Errorf("resolve attachment %s: %w", attachment.Name, err) + } + relative, err := filepath.Rel(attachmentRoot, attachmentPath) + if err != nil || relative == "." || filepath.IsAbs(relative) || relative == ".." || strings.HasPrefix(relative, ".."+string(filepath.Separator)) { + return nil, fmt.Errorf("attachment %s is outside the attachment root", attachment.Name) + } + info, err := os.Stat(attachmentPath) + if err != nil { + return nil, err + } + if !info.Mode().IsRegular() || info.Size() <= 0 || info.Size() > m.cfg.AttachmentSizeLimit { + return nil, fmt.Errorf("attachment %s has an invalid size or type", attachment.Name) + } + data, err := os.ReadFile(attachmentPath) + if err != nil { + return nil, err + } + if int64(len(data)) != info.Size() || !strings.HasPrefix(strings.ToLower(http.DetectContentType(data)), "image/") { + return nil, fmt.Errorf("attachment %s is not a valid image", attachment.Name) + } + images = append(images, piRPCImage{ + Type: "image", + Data: base64.StdEncoding.EncodeToString(data), + MimeType: declaredMime, + }) + } + return images, nil +} + +func (r *piSessionRuntime) acquire() bool { + if r == nil { + return false + } + r.mu.Lock() + defer r.mu.Unlock() + if r.stopped { + return false + } + if r.idleTimer != nil { + r.idleTimer.Stop() + r.idleTimer = nil + } + return true +} + +func (r *piSessionRuntime) activate(run *activeRun, session tables.WebSessionTable) (*piRuntimeRun, error) { + r.mu.Lock() + defer r.mu.Unlock() + if r.stopped { + return nil, errPiRPCClosed + } + if r.active != nil { + return nil, errors.New("Pi RPC runtime is already processing a prompt") + } + if r.idleTimer != nil { + r.idleTimer.Stop() + r.idleTimer = nil + } + dispatch := &piRuntimeRun{ + runtime: r, + run: run, + session: session, + settled: make(chan error, 1), + contents: make(map[int]*piRuntimeContentState), + tools: make(map[string]*piRuntimeToolState), + } + r.active = dispatch + return dispatch, nil +} + +func (r *piSessionRuntime) deactivate(dispatch *piRuntimeRun) { + r.mu.Lock() + if r.active == dispatch { + r.active = nil + } + r.scheduleIdleLocked() + r.mu.Unlock() +} + +func (r *piSessionRuntime) scheduleIdle() { + r.mu.Lock() + r.scheduleIdleLocked() + r.mu.Unlock() +} + +func (r *piSessionRuntime) scheduleIdleLocked() { + if r.stopped || r.active != nil { + return + } + if r.idleTimer != nil { + r.idleTimer.Stop() + } + ttl := r.manager.cfg.PiRuntimeIdleTTL + r.idleTimer = time.AfterFunc(ttl, func() { + r.stop(errors.New("Pi RPC idle timeout")) + }) +} + +func (r *piSessionRuntime) consumeEvents() { + for event := range r.client.Events() { + r.mu.Lock() + dispatch := r.active + r.mu.Unlock() + if dispatch == nil { + continue + } + if event.Type == "codekanban_barrier" { + if event.BarrierDone != nil { + close(event.BarrierDone) + } + if !event.WakePending || event.BarrierRunID != dispatch.run.runID || + !dispatch.run.finishPiResponseBarrier(event.BarrierGeneration) { + continue + } + r.manager.triggerPendingProcessing(dispatch.session.ID) + continue + } + if err := r.manager.handlePiRuntimeEvent(dispatch, event); err != nil { + dispatch.settle(err) + continue + } + if event.Type == "agent_settled" { + dispatch.settle(nil) + } + } + err := r.client.processError() + if err == nil { + err = errPiRPCClosed + } + r.mu.Lock() + dispatch := r.active + r.mu.Unlock() + if dispatch != nil { + dispatch.settle(err) + } + r.manager.clearPiNativeQueuedInputs(r.sessionID) + r.removeFromManager() +} + +func (r *piSessionRuntime) stop(reason error) { + if r == nil { + return + } + r.stopOnce.Do(func() { + r.mu.Lock() + r.stopped = true + if r.idleTimer != nil { + r.idleTimer.Stop() + r.idleTimer = nil + } + dispatch := r.active + r.mu.Unlock() + if dispatch != nil { + abortCtx, cancel := context.WithTimeout(context.Background(), time.Second) + _ = r.client.Request(abortCtx, "abort", nil, nil) + cancel() + dispatch.settle(reason) + } + _ = r.client.Close() + r.manager.clearPiNativeQueuedInputs(r.sessionID) + r.removeFromManager() + }) +} + +func (r *piSessionRuntime) removeFromManager() { + m := r.manager + if m == nil { + return + } + m.piRuntimeMu.Lock() + if m.piRuntimes[r.sessionID] == r { + delete(m.piRuntimes, r.sessionID) + delete(m.piRuntimeTerminators, r.sessionID) + } + m.piRuntimeMu.Unlock() +} + +func (m *Manager) runPiRPCSession( + ctx context.Context, + run *activeRun, + session tables.WebSessionTable, + text string, + attachments []Attachment, +) { + runtime, err := m.getOrStartPiRuntime(ctx, session) + if err != nil { + m.handleRunFailure(session.ID, session, run, err) + return + } + dispatch, err := runtime.activate(run, session) + if err != nil { + runtime.scheduleIdle() + m.handleRunFailure(session.ID, session, run, err) + return + } + defer runtime.deactivate(dispatch) + + if model := strings.TrimSpace(session.Model); model != "" { + provider, modelID, err := splitPiModel(model) + if err != nil { + m.handleRunFailure(session.ID, session, run, err) + return + } + var selected piRPCSetModelResult + if err := requestPiRuntimeControl(runtime.client, "set_model", map[string]any{ + "provider": provider, + "modelId": modelID, + }, &selected); err != nil { + m.handleRunFailure(session.ID, session, run, err) + return + } + if !strings.EqualFold(strings.TrimSpace(selected.Provider), provider) || + !strings.EqualFold(strings.TrimSpace(selected.ID), modelID) { + m.handleRunFailure(session.ID, session, run, errors.New("Pi did not select the requested model")) + return + } + } + effort := normalizeReasoningEffort(ReasoningEffort(session.ReasoningEffort)) + level := piReasoningToThinkingLevel(effort) + if effort != ReasoningEffortDefault && level == "" { + m.handleRunFailure(session.ID, session, run, fmt.Errorf("Pi does not support reasoning effort %q", effort)) + return + } + if ctx.Err() != nil { + m.finishPiAbortedRun(session, run) + return + } + if level != "" { + if err := requestPiRuntimeControl(runtime.client, "set_thinking_level", map[string]any{"level": level}, nil); err != nil { + m.handleRunFailure(session.ID, session, run, err) + return + } + } + if ctx.Err() != nil { + m.finishPiAbortedRun(session, run) + return + } + images, err := m.piPromptImages(attachments) + if err != nil { + m.handleRunFailure(session.ID, session, run, err) + return + } + payload := map[string]any{"message": preparePromptText(text, effectiveWorkflowMode(session))} + if len(images) > 0 { + payload["images"] = images + } + if ctx.Err() != nil { + m.finishPiAbortedRun(session, run) + return + } + if err := requestPiRuntimeControl(runtime.client, "prompt", payload, nil); err != nil { + m.handleRunFailure(session.ID, session, run, err) + return + } + + select { + case err := <-dispatch.settled: + if ctx.Err() != nil { + m.finishPiAbortedRun(session, run) + return + } + if err != nil { + m.handleRunFailure(session.ID, session, run, err) + return + } + case <-ctx.Done(): + abortCtx, cancel := context.WithTimeout(context.Background(), 3*time.Second) + _ = runtime.client.Request(abortCtx, "abort", nil, nil) + cancel() + select { + case <-dispatch.settled: + case <-time.After(3 * time.Second): + runtime.stop(errors.New("Pi RPC abort timed out")) + } + m.finishPiAbortedRun(session, run) + return + } + + if err := m.finishSuccessfulPiRun(runtime, dispatch, session, run); err != nil { + m.handleRunFailure(session.ID, session, run, err) + } +} + +func (m *Manager) runPiRPCCompaction(ctx context.Context, run *activeRun, session tables.WebSessionTable) { + runtime, err := m.getOrStartPiRuntime(ctx, session) + if err != nil { + m.handleRunFailure(session.ID, session, run, err) + return + } + dispatch, err := runtime.activate(run, session) + if err != nil { + runtime.scheduleIdle() + m.handleRunFailure(session.ID, session, run, err) + return + } + defer runtime.deactivate(dispatch) + + requestDone := make(chan error, 1) + go func() { + requestCtx, cancel := context.WithTimeout(context.Background(), piRPCRequestTimeout) + defer cancel() + requestDone <- runtime.client.Request(requestCtx, "compact", nil, nil) + }() + select { + case err := <-requestDone: + if ctx.Err() != nil { + m.finishPiAbortedRun(session, run) + return + } + if err != nil { + m.handleRunFailure(session.ID, session, run, err) + return + } + case <-ctx.Done(): + abortCtx, cancel := context.WithTimeout(context.Background(), 3*time.Second) + _ = runtime.client.Request(abortCtx, "abort", nil, nil) + cancel() + select { + case <-requestDone: + case <-time.After(3 * time.Second): + runtime.stop(errors.New("Pi RPC compaction abort timed out")) + } + m.finishPiAbortedRun(session, run) + return + } + + barrierCtx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + err = runtime.client.BarrierAndWait(barrierCtx, run.runID) + cancel() + if err != nil { + runtime.stop(errors.New("Pi RPC compaction event barrier failed")) + m.handleRunFailure(session.ID, session, run, err) + return + } + if err := m.finishSuccessfulPiRun(runtime, dispatch, session, run); err != nil { + m.handleRunFailure(session.ID, session, run, err) + } +} + +func (m *Manager) finishSuccessfulPiRun(runtime *piSessionRuntime, dispatch *piRuntimeRun, session tables.WebSessionTable, run *activeRun) error { + var compactionAt *time.Time + if dispatch != nil { + dispatch.mu.Lock() + if !dispatch.compactionCompleted.IsZero() { + value := dispatch.compactionCompleted + compactionAt = &value + } + dispatch.mu.Unlock() + } + if err := m.syncPiRuntimeSnapshot(context.Background(), runtime, session, compactionAt); err != nil { + return err + } + if err := m.broadcastSnapshot(context.Background(), session.ID); err != nil { + return err + } + now := time.Now() + if _, err := m.appendAndBroadcast(context.Background(), session.ID, session, Event{ + ID: utils.NewID(), Type: "run_done", RunID: run.runID, Timestamp: now, + Payload: map[string]any{"ok": true, "st": string(StatusDone)}, + }); err != nil { + return err + } + if err := m.updateRuntimeState(context.Background(), session.ID, applyAssistantStateUpdates(map[string]any{ + "status": string(StatusDone), "updated_at": now, "last_error": nil, + "auto_retry_attempt": 0, "auto_retry_next_at": nil, "auto_retry_last_error_code": nil, + }, AssistantStateNone, now)); err != nil { + return err + } + m.cancelAutoRetryTimer(session.ID) + m.broadcastSessionSummary(context.Background(), session.ID) + run.syncSourceAfterRun = true + return nil +} + +func requestPiRuntimeControl(client *piRPCClient, command string, payload map[string]any, target any) error { + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + return client.Request(ctx, command, payload, target) +} + +func (m *Manager) sendActivePiPendingInput( + ctx context.Context, + session tables.WebSessionTable, + pending PendingInput, +) (bool, error) { + m.mu.RLock() + run := m.runs[session.ID] + m.mu.RUnlock() + if run == nil || normalizeAgent(run.agent) != AgentPi || run.backend != SessionBackendPiRPC { + return false, nil + } + if run.blocksPiPendingInput() { + return false, nil + } + m.piRuntimeMu.Lock() + runtime := m.piRuntimes[session.ID] + m.piRuntimeMu.Unlock() + if runtime == nil { + return false, nil + } + runtime.mu.Lock() + dispatch := runtime.active + valid := !runtime.stopped && dispatch != nil && dispatch.run == run + runtime.mu.Unlock() + if !valid { + return false, nil + } + + attachments := make([]Attachment, 0, len(pending.AttachmentIDs)) + for _, attachmentID := range pending.AttachmentIDs { + attachment, err := m.loadAttachment(strings.TrimSpace(attachmentID)) + if err != nil { + return true, fmt.Errorf("attachment %s not found", attachmentID) + } + attachments = append(attachments, attachment) + } + text := strings.TrimSpace(pending.Text) + if text == "" && len(attachments) == 0 { + return true, errors.New("message is empty") + } + images, err := m.piPromptImages(attachments) + if err != nil { + return true, err + } + command := "follow_up" + if pending.Mode == PendingInputModeRedirect { + command = "steer" + } + payload := map[string]any{"message": text} + if len(images) > 0 { + payload["images"] = images + } + requestCtx, cancel := context.WithTimeout(ctx, 10*time.Second) + err = runtime.client.Request(requestCtx, command, payload, nil) + cancel() + if err != nil { + return true, err + } + + messageID := utils.NewID() + if _, err := m.appendAndBroadcast(context.Background(), session.ID, session, Event{ + ID: utils.NewID(), Type: "msg_u", RunID: run.runID, ParentID: messageID, Timestamp: time.Now(), + Payload: map[string]any{ + "mid": messageID, "txt": text, "atts": attachmentPayloads(attachments), "piQueueMode": string(pending.Mode), + }, + }); err != nil && m.logger != nil { + m.logger.Error("failed to persist Pi queued message", zap.String("sessionId", session.ID), zap.Error(err)) + } + return true, nil +} + +func (m *Manager) syncPiRuntimeSnapshot(ctx context.Context, runtime *piSessionRuntime, session tables.WebSessionTable, compactionAt ...*time.Time) error { + if _, hasDeadline := ctx.Deadline(); !hasDeadline { + var cancel context.CancelFunc + ctx, cancel = context.WithTimeout(ctx, 10*time.Second) + defer cancel() + } + var state piRPCState + if err := runtime.client.Request(ctx, "get_state", nil, &state); err != nil { + return err + } + if err := validatePiRuntimeState(session, state); err != nil { + return err + } + var entries piHistoryEntriesResponse + if err := runtime.client.Request(ctx, "get_entries", nil, &entries); err != nil { + return err + } + var stats piRPCSessionStats + if err := runtime.client.Request(ctx, "get_session_stats", nil, &stats); err != nil { + return err + } + updates := map[string]any{ + "native_session_id": state.SessionID, + "thread_path": filepath.Clean(state.SessionFile), + "native_leaf_id": nilIfEmpty(pointerString(entries.LeafID)), + "source_revision": nilIfEmpty(piSourceRevision(state.SessionFile, pointerString(entries.LeafID))), + "total_input_tokens": stats.Tokens.Input, + "total_cached_input_tokens": stats.Tokens.CacheRead, + "total_output_tokens": stats.Tokens.Output, + "total_cost": stats.Cost, + "last_synced_at": time.Now(), + "sync_state": string(SyncStateFresh), + "sync_error": nil, + "updated_at": time.Now(), + } + if len(compactionAt) > 0 && compactionAt[0] != nil && !compactionAt[0].IsZero() { + updates["last_context_compaction_at"] = *compactionAt[0] + updates["context_baseline_input_tokens"] = stats.Tokens.Input + updates["context_baseline_cached_input_tokens"] = stats.Tokens.CacheRead + updates["context_baseline_output_tokens"] = stats.Tokens.Output + updates["latest_token_count_input_tokens"] = 0 + updates["latest_token_count_cached_input_tokens"] = 0 + updates["latest_token_count_output_tokens"] = 0 + updates["latest_token_count_total_tokens"] = 0 + updates["latest_token_count_updated_at"] = nil + updates["latest_turn_input_tokens"] = 0 + updates["latest_turn_cached_input_tokens"] = 0 + updates["latest_turn_output_tokens"] = 0 + updates["latest_turn_usage_updated_at"] = nil + } + if stats.ContextUsage != nil { + updates["session_context_window_tokens"] = stats.ContextUsage.ContextWindow + updates["session_context_window_observed_at"] = time.Now() + if len(compactionAt) == 0 || compactionAt[0] == nil || compactionAt[0].IsZero() { + updates["latest_token_count_total_tokens"] = stats.ContextUsage.Tokens + updates["latest_token_count_updated_at"] = time.Now() + } + } + if model := canonicalPiModel(state.Model); model != "" { + updates["model"] = model + } + if effort := piThinkingLevelToReasoning(state.ThinkingLevel); effort != ReasoningEffortDefault { + updates["reasoning_effort"] = string(effort) + } + return m.reconcileLivePiHistory(ctx, session, state.SessionID, entries, updates) +} + +func (m *Manager) finishPiAbortedRun(session tables.WebSessionTable, run *activeRun) { + _ = m.closePendingPiDialog(session, run, "Pi extension input was canceled because the run was aborted") + now := time.Now() + _, _ = m.appendAndBroadcast(context.Background(), session.ID, session, Event{ + ID: utils.NewID(), Type: "run_abort", RunID: run.runID, Timestamp: now, + }) + _ = m.updateRuntimeState(context.Background(), session.ID, applyAssistantStateUpdates(map[string]any{ + "status": string(StatusIdle), "updated_at": now, + "auto_retry_attempt": 0, "auto_retry_next_at": nil, "auto_retry_last_error_code": nil, + }, AssistantStateNone, now)) + m.cancelAutoRetryTimer(session.ID) + m.broadcastSessionSummary(context.Background(), session.ID) +} diff --git a/service/websession/pi_runtime_probe.go b/service/websession/pi_runtime_probe.go new file mode 100644 index 00000000..21f0a5cb --- /dev/null +++ b/service/websession/pi_runtime_probe.go @@ -0,0 +1,349 @@ +package websession + +import ( + "bufio" + "context" + "encoding/json" + "errors" + "fmt" + "io" + "os" + "os/exec" + "path/filepath" + "runtime" + "sort" + "strings" + "time" + + "github.com/Masterminds/semver/v3" +) + +const ( + piMinVersion = "0.84.1" + piProbeSuccessCacheTTL = 5 * time.Minute + piProbeFailureCacheTTL = 5 * time.Second + piProbeTimeout = 5 * time.Second + piProbeMaxFrameBytes = 1024 * 1024 + piDiagnosticNotInstalled = "not_installed" + piDiagnosticVersion = "version_unknown" + piDiagnosticTooOld = "version_too_old" + piDiagnosticStart = "rpc_start_failed" + piDiagnosticProtocol = "rpc_protocol_incompatible" + piDiagnosticTimeout = "rpc_timeout" +) + +var piMinimumVersion = semver.MustParse(piMinVersion) + +type piRuntimeProbeResult struct { + installed bool + version *string + compatible bool + diagnostic string + models []PiModelInfo +} + +type piRuntimeProbeCache struct { + result piRuntimeProbeResult + expiresAt time.Time + loaded bool +} + +type piProbeResponse struct { + Type string `json:"type"` + ID string `json:"id"` + Command string `json:"command"` + Success bool `json:"success"` +} + +func (m *Manager) applyPiRuntimeCapabilities(config WebSessionRuntimeConfig) WebSessionRuntimeConfig { + config.PiMinVersion = piMinVersion + if m == nil { + config.Agents = runtimeAgentCapabilities(config) + return config + } + + probe := m.getPiRuntimeProbe() + config.HasPi = probe.installed + config.PiVersion = probe.version + config.PiRPCCompatible = probe.compatible + config.PiDiagnostics = probe.diagnostic + config.PiModels = append([]PiModelInfo(nil), probe.models...) + config.SupportsPiWebSession = probe.compatible + config.Agents = runtimeAgentCapabilities(config) + return config +} + +func (m *Manager) getPiRuntimeProbe() piRuntimeProbeResult { + m.piProbeMu.Lock() + defer m.piProbeMu.Unlock() + + now := time.Now() + if m.piProbe.loaded && now.Before(m.piProbe.expiresAt) { + return m.piProbe.result + } + + result := probePiRuntime(m.cfg.PiPath, m.cfg.DataDir) + ttl := piProbeFailureCacheTTL + if result.compatible { + ttl = piProbeSuccessCacheTTL + } + m.piProbe = piRuntimeProbeCache{ + result: result, + expiresAt: now.Add(ttl), + loaded: true, + } + return result +} + +func probePiRuntime(command, workingDir string) piRuntimeProbeResult { + result := piRuntimeProbeResult{} + if !hasExecutable(command) { + result.diagnostic = piDiagnosticNotInstalled + return result + } + result.installed = true + + version := detectPiVersion(command) + if version == nil { + result.diagnostic = piDiagnosticVersion + return result + } + result.version = version + parsedVersion, err := semver.NewVersion(*version) + if err != nil { + result.diagnostic = piDiagnosticVersion + return result + } + if parsedVersion.LessThan(piMinimumVersion) { + result.diagnostic = piDiagnosticTooOld + return result + } + + ctx, cancel := context.WithTimeout(context.Background(), piProbeTimeout) + defer cancel() + if err := runPiRPCProbe(ctx, command, workingDir); err != nil { + switch { + case errors.Is(err, context.DeadlineExceeded), errors.Is(ctx.Err(), context.DeadlineExceeded): + result.diagnostic = piDiagnosticTimeout + case errors.Is(err, errPiProbeStart): + result.diagnostic = piDiagnosticStart + default: + result.diagnostic = piDiagnosticProtocol + } + return result + } + + result.compatible = true + modelCtx, modelCancel := context.WithTimeout(context.Background(), piProbeTimeout) + result.models, _ = loadPiModelCatalog(modelCtx, command, workingDir) + modelCancel() + return result +} + +func loadPiModelCatalog(ctx context.Context, command, workingDir string) ([]PiModelInfo, error) { + cmd, err := buildPiCommand( + ctx, + command, + "--mode", "rpc", + "--no-session", + "--no-approve", + "--no-extensions", + "--no-skills", + "--no-prompt-templates", + "--no-themes", + ) + if err != nil { + return nil, err + } + if dir := strings.TrimSpace(workingDir); dir != "" { + if info, statErr := os.Stat(dir); statErr == nil && info.IsDir() { + cmd.Dir = dir + } + } + client, err := startPiRPCClient(cmd) + if err != nil { + return nil, err + } + defer client.Close() + var response struct { + Models []PiModelInfo `json:"models"` + } + if err := client.Request(ctx, "get_available_models", nil, &response); err != nil { + return nil, err + } + models := response.Models[:0] + seen := make(map[string]struct{}, len(response.Models)) + for _, model := range response.Models { + model.Provider = strings.TrimSpace(model.Provider) + model.ID = strings.TrimSpace(model.ID) + model.Name = strings.TrimSpace(model.Name) + if model.Provider == "" || model.ID == "" { + continue + } + key := strings.ToLower(model.Provider + "/" + model.ID) + if _, ok := seen[key]; ok { + continue + } + seen[key] = struct{}{} + model.Input = append([]string(nil), model.Input...) + models = append(models, model) + } + sort.SliceStable(models, func(i, j int) bool { + left := strings.ToLower(models[i].Provider + "/" + models[i].Name + "/" + models[i].ID) + right := strings.ToLower(models[j].Provider + "/" + models[j].Name + "/" + models[j].ID) + return left < right + }) + return append([]PiModelInfo(nil), models...), nil +} + +func detectPiVersion(command string) *string { + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + cmd, err := buildPiCommand(ctx, command, "--version") + if err != nil { + return nil + } + output, err := cmd.CombinedOutput() + if err != nil { + return nil + } + match := codexVersionPattern.FindString(string(output)) + if strings.TrimSpace(match) == "" { + return nil + } + version := strings.TrimSpace(match) + return &version +} + +var errPiProbeStart = errors.New("pi probe process failed to start") + +func runPiRPCProbe(ctx context.Context, command, workingDir string) error { + cmd, err := buildPiCommand( + ctx, + command, + "--mode", "rpc", + "--no-session", + "--no-approve", + "--no-extensions", + "--no-skills", + "--no-prompt-templates", + "--no-themes", + ) + if err != nil { + return fmt.Errorf("%w: %v", errPiProbeStart, err) + } + if dir := strings.TrimSpace(workingDir); dir != "" { + if info, statErr := os.Stat(dir); statErr == nil && info.IsDir() { + cmd.Dir = dir + } + } + + stdin, err := cmd.StdinPipe() + if err != nil { + return fmt.Errorf("%w: stdin", errPiProbeStart) + } + stdout, err := cmd.StdoutPipe() + if err != nil { + return fmt.Errorf("%w: stdout", errPiProbeStart) + } + cmd.Stderr = io.Discard + if err := cmd.Start(); err != nil { + return fmt.Errorf("%w: start", errPiProbeStart) + } + + waitCh := make(chan error, 1) + go func() { + waitCh <- cmd.Wait() + }() + defer stopPiProbeProcess(cmd, stdin, waitCh) + + expected := map[string]string{ + "state": "get_state", + "entries": "get_entries", + "tree": "get_tree", + "stats": "get_session_stats", + } + encoder := json.NewEncoder(stdin) + for id, commandType := range expected { + if err := encoder.Encode(map[string]string{"id": id, "type": commandType}); err != nil { + return err + } + } + if err := stdin.Close(); err != nil { + return err + } + + received := make(map[string]struct{}, len(expected)) + scanner := bufio.NewScanner(stdout) + scanner.Buffer(make([]byte, 0, 64*1024), piProbeMaxFrameBytes) + for scanner.Scan() { + line := scanner.Bytes() + var response piProbeResponse + if err := json.Unmarshal(line, &response); err != nil { + return err + } + if response.Type != "response" { + continue + } + expectedCommand, ok := expected[response.ID] + if !ok { + continue + } + if !response.Success || response.Command != expectedCommand { + return fmt.Errorf("unexpected response for %s", response.ID) + } + received[response.ID] = struct{}{} + if len(received) == len(expected) { + return nil + } + } + if err := scanner.Err(); err != nil { + if ctx.Err() != nil { + return ctx.Err() + } + return err + } + if ctx.Err() != nil { + return ctx.Err() + } + return io.ErrUnexpectedEOF +} + +func stopPiProbeProcess(cmd *exec.Cmd, stdin io.Closer, waitCh <-chan error) { + if stdin != nil { + _ = stdin.Close() + } + select { + case <-waitCh: + return + case <-time.After(300 * time.Millisecond): + } + killCmdTree(cmd) + select { + case <-waitCh: + case <-time.After(time.Second): + } +} + +func buildPiCommand(ctx context.Context, command string, args ...string) (*exec.Cmd, error) { + parts := splitCommandParts(command) + if len(parts) == 0 { + return nil, errors.New("pi command is empty") + } + commandArgs := append(append([]string{}, parts[1:]...), args...) + executable := parts[0] + if resolved, err := exec.LookPath(executable); err == nil { + executable = resolved + } + if runtime.GOOS == "windows" { + ext := strings.ToLower(filepath.Ext(executable)) + if ext == ".cmd" || ext == ".bat" { + comspec := strings.TrimSpace(os.Getenv("ComSpec")) + if comspec == "" { + comspec = "cmd.exe" + } + return buildWindowsBatchCommand(ctx, comspec, executable, commandArgs), nil + } + } + return exec.CommandContext(ctx, executable, commandArgs...), nil +} diff --git a/service/websession/pi_runtime_probe_test.go b/service/websession/pi_runtime_probe_test.go new file mode 100644 index 00000000..8248481f --- /dev/null +++ b/service/websession/pi_runtime_probe_test.go @@ -0,0 +1,150 @@ +package websession + +import ( + "bufio" + "context" + "encoding/json" + "fmt" + "os" + "path/filepath" + "strings" + "testing" +) + +func TestPiProbeHelperProcess(t *testing.T) { + if os.Getenv("CODEKANBAN_PI_PROBE_HELPER") != "1" { + return + } + if marker := os.Getenv("CODEKANBAN_PI_PROBE_MARKER"); marker != "" { + file, _ := os.OpenFile(marker, os.O_CREATE|os.O_APPEND|os.O_WRONLY, 0o600) + if file != nil { + _, _ = file.WriteString("start\n") + _ = file.Close() + } + } + for _, arg := range os.Args { + if arg == "--version" { + fmt.Fprintln(os.Stdout, os.Getenv("CODEKANBAN_PI_PROBE_VERSION")) + os.Exit(0) + } + } + + skip := os.Getenv("CODEKANBAN_PI_PROBE_SKIP") + scanner := bufio.NewScanner(os.Stdin) + for scanner.Scan() { + var request struct { + ID string `json:"id"` + Type string `json:"type"` + } + if json.Unmarshal(scanner.Bytes(), &request) != nil || request.ID == skip { + continue + } + data := any(map[string]any{}) + if request.Type == "get_available_models" { + data = map[string]any{"models": []map[string]any{{ + "provider": "anthropic", "id": "claude-sonnet-4", "name": "Claude Sonnet 4", + "reasoning": true, "input": []string{"text", "image"}, "contextWindow": 200000, + }}} + } + _ = json.NewEncoder(os.Stdout).Encode(map[string]any{ + "id": request.ID, + "type": "response", + "command": request.Type, + "success": true, + "data": data, + }) + } + os.Exit(0) +} + +func piProbeTestCommand() string { + return `"` + os.Args[0] + `" -test.run=TestPiProbeHelperProcess --` +} + +func TestGetWebSessionRuntimeConfigProbesPiRPCAndCachesSuccess(t *testing.T) { + marker := filepath.Join(t.TempDir(), "probe-starts.txt") + t.Setenv("CODEKANBAN_PI_PROBE_HELPER", "1") + t.Setenv("CODEKANBAN_PI_PROBE_VERSION", "0.84.1") + t.Setenv("CODEKANBAN_PI_PROBE_MARKER", marker) + + manager := &Manager{cfg: Config{ + PiPath: piProbeTestCommand(), + DataDir: t.TempDir(), + }} + config := manager.GetWebSessionRuntimeConfig() + if !config.HasPi || !config.PiRPCCompatible { + t.Fatalf("unexpected Pi probe result: %#v", config) + } + if config.PiVersion == nil || *config.PiVersion != "0.84.1" { + t.Fatalf("Pi version = %#v, want 0.84.1", config.PiVersion) + } + if config.PiDiagnostics != "" || config.PiMinVersion != piMinVersion { + t.Fatalf("unexpected Pi diagnostics: code=%q minimum=%q", config.PiDiagnostics, config.PiMinVersion) + } + if !config.SupportsPiWebSession || !config.Agents[AgentPi].SupportsWebSession { + t.Fatal("compatible Pi RPC should enable Pi Web Sessions") + } + piCapability := config.Agents[AgentPi] + if !piCapability.Installed || !piCapability.SupportsImages || !piCapability.SupportsCompaction || + !piCapability.SupportsSteer || !piCapability.SupportsFollowUp || !piCapability.SupportsTree { + t.Fatal("compatible Pi RPC should expose messaging and native tree controls") + } + if len(config.PiModels) != 1 || config.PiModels[0].Provider != "anthropic" || !config.PiModels[0].Reasoning { + t.Fatalf("Pi model catalog = %#v", config.PiModels) + } + + _ = manager.GetWebSessionRuntimeConfig() + raw, err := os.ReadFile(marker) + if err != nil { + t.Fatalf("read probe marker: %v", err) + } + if starts := strings.Count(string(raw), "start\n"); starts != 3 { + t.Fatalf("helper starts = %d, want 3 (version + RPC + models) after cached second read", starts) + } +} + +func TestGetWebSessionRuntimeConfigReportsSafePiDiagnostics(t *testing.T) { + tests := []struct { + name string + version string + skip string + path string + diagnostic string + installed bool + }{ + {name: "not installed", path: filepath.Join(t.TempDir(), "missing-pi"), diagnostic: piDiagnosticNotInstalled}, + {name: "unknown version", version: "development", diagnostic: piDiagnosticVersion, installed: true}, + {name: "old version", version: "0.84.0", diagnostic: piDiagnosticTooOld, installed: true}, + {name: "missing RPC command", version: "0.84.1", skip: "tree", diagnostic: piDiagnosticProtocol, installed: true}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Setenv("CODEKANBAN_PI_PROBE_HELPER", "1") + t.Setenv("CODEKANBAN_PI_PROBE_VERSION", tt.version) + t.Setenv("CODEKANBAN_PI_PROBE_SKIP", tt.skip) + command := tt.path + if command == "" { + command = piProbeTestCommand() + } + manager := &Manager{cfg: Config{PiPath: command, DataDir: t.TempDir()}} + config := manager.GetWebSessionRuntimeConfig() + if config.HasPi != tt.installed || config.PiRPCCompatible { + t.Fatalf("unexpected Pi state: hasPi=%v compatible=%v", config.HasPi, config.PiRPCCompatible) + } + if config.PiDiagnostics != tt.diagnostic { + t.Fatalf("Pi diagnostics = %q, want %q", config.PiDiagnostics, tt.diagnostic) + } + }) + } +} + +func TestBuildPiCommandSupportsConfiguredArguments(t *testing.T) { + cmd, err := buildPiCommand(context.Background(), `"C:\Program Files\node.exe" "C:\pi\dist\cli.js"`, "--version") + if err != nil { + t.Fatalf("buildPiCommand returned error: %v", err) + } + if cmd.Path == "" || len(cmd.Args) != 3 || cmd.Args[1] != `C:\pi\dist\cli.js` || cmd.Args[2] != "--version" { + t.Fatalf("unexpected command: path=%q args=%#v", cmd.Path, cmd.Args) + } +} diff --git a/service/websession/pi_runtime_probe_windows_test.go b/service/websession/pi_runtime_probe_windows_test.go new file mode 100644 index 00000000..a4f81392 --- /dev/null +++ b/service/websession/pi_runtime_probe_windows_test.go @@ -0,0 +1,47 @@ +//go:build windows + +package websession + +import ( + "context" + "fmt" + "os" + "path/filepath" + "strings" + "testing" +) + +func TestBuildPiCommandUsesCmdForWindowsShim(t *testing.T) { + cmd, err := buildPiCommand(context.Background(), `C:\tools\pi.cmd`, "--version") + if err != nil { + t.Fatalf("buildPiCommand returned error: %v", err) + } + if !strings.EqualFold(filepath.Base(cmd.Path), "cmd.exe") { + t.Fatalf("command path = %q, want cmd.exe", cmd.Path) + } + commandLine := cmd.SysProcAttr.CmdLine + if !strings.Contains(commandLine, `C:\tools\pi.cmd`) || !strings.Contains(commandLine, "--version") { + t.Fatalf("unexpected shim command line: %q", commandLine) + } +} + +func TestProbePiRuntimeExecutesWindowsCmdShimWithSpaces(t *testing.T) { + t.Setenv("CODEKANBAN_PI_PROBE_HELPER", "1") + t.Setenv("CODEKANBAN_PI_PROBE_VERSION", "0.84.1") + dir := filepath.Join(t.TempDir(), "Pi Command") + if err := os.MkdirAll(dir, 0o755); err != nil { + t.Fatalf("create shim dir: %v", err) + } + shimPath := filepath.Join(dir, "pi.cmd") + content := fmt.Sprintf( + "@echo off\r\n\"%s\" -test.run=TestPiProbeHelperProcess -- %%*\r\n", + os.Args[0], + ) + if err := os.WriteFile(shimPath, []byte(content), 0o644); err != nil { + t.Fatalf("write shim: %v", err) + } + result := probePiRuntime(`"`+shimPath+`"`, t.TempDir()) + if !result.installed || !result.compatible || result.version == nil || *result.version != "0.84.1" { + t.Fatalf("unexpected cmd shim probe result: %#v", result) + } +} diff --git a/service/websession/pi_runtime_test.go b/service/websession/pi_runtime_test.go new file mode 100644 index 00000000..d56fe3fa --- /dev/null +++ b/service/websession/pi_runtime_test.go @@ -0,0 +1,1983 @@ +package websession + +import ( + "bufio" + "context" + "crypto/rand" + "encoding/base64" + "encoding/hex" + "encoding/json" + "errors" + "fmt" + "os" + "path/filepath" + "strconv" + "strings" + "testing" + "time" + + "code-kanban/model" + "code-kanban/model/tables" + + "go.uber.org/zap" +) + +func TestPiRuntimeFakeProcess(t *testing.T) { + if os.Getenv("CODEKANBAN_FAKE_PI_RUNTIME") != "1" { + return + } + args := argsAfterDoubleDash(os.Args) + if containsString(args, "--version") { + fmt.Println("0.84.1") + return + } + noSession := containsString(args, "--no-session") + bridgePath := fakePiFlagValue(args, "--extension") + sessionPath := fakePiFlagValue(args, "--session") + if sessionPath == "" { + sessionPath = os.Getenv("CODEKANBAN_FAKE_PI_SESSION") + } + sessionID := "fake-pi-session" + if restoredID := readFakePiSessionHeaderID(sessionPath); restoredID != "" { + sessionID = restoredID + } + cwd, _ := os.Getwd() + if !noSession { + appendFakePiLog(map[string]any{"startup": os.Getpid(), "session": true, "args": args}) + } + + modelProvider := "openai" + modelID := "gpt-fake" + thinkingLevel := "medium" + entries, leafID := readFakePiSessionEntries(sessionPath) + pr5Active := false + pr5InputAnswered := false + pr5Steered := false + pr5FollowedUp := false + pr5Finished := false + entrySequence := len(entries) + mutationSequence := 0 + appendNativeEntry := func(entry map[string]any) { + parentID := leafID + entrySequence++ + leafID = "leaf-" + strconv.Itoa(entrySequence) + entry["id"] = leafID + entry["parentId"] = nilIfEmpty(parentID) + entry["timestamp"] = time.Now().UTC().Format(time.RFC3339Nano) + entries = append(entries, entry) + file, _ := os.OpenFile(sessionPath, os.O_APPEND|os.O_WRONLY, 0o600) + if file != nil { + encoded, _ := json.Marshal(entry) + _, _ = fmt.Fprintln(file, string(encoded)) + _ = file.Close() + } + } + appendEntry := func(role, text string) { + appendNativeEntry(map[string]any{ + "type": "message", + "message": map[string]any{"role": role, "content": []any{map[string]any{"type": "text", "text": text}}}, + }) + } + finishPR5 := func() {} + finishPR5 = func() { + if pr5Finished || !pr5Active || !pr5InputAnswered || !pr5Steered || !pr5FollowedUp { + return + } + pr5Finished = true + writeFakePiEvent(map[string]any{"type": "auto_retry_start", "attempt": 1, "maxAttempts": 3, "delayMs": 1, "errorMessage": "retryable"}) + writeFakePiEvent(map[string]any{"type": "auto_retry_end", "success": true, "attempt": 1}) + writeFakePiEvent(map[string]any{"type": "compaction_start", "reason": "threshold"}) + writeFakePiEvent(map[string]any{"type": "compaction_end", "result": map[string]any{"summary": "compacted context"}, "aborted": false}) + writeFakePiEvent(map[string]any{"type": "tool_execution_update", "toolCallId": "parallel-b", "toolName": "Read", "partialResult": map[string]any{"content": []any{map[string]any{"type": "text", "text": "B partial"}}}}) + writeFakePiEvent(map[string]any{"type": "tool_execution_end", "toolCallId": "parallel-b", "toolName": "Read", "result": map[string]any{"content": []any{map[string]any{"type": "text", "text": "B done"}}}, "isError": false}) + writeFakePiEvent(map[string]any{"type": "tool_execution_end", "toolCallId": "parallel-a", "toolName": "Write", "result": map[string]any{"content": []any{map[string]any{"type": "text", "text": "A done"}}}, "isError": false}) + writeFakePiEvent(map[string]any{"type": "message_update", "assistantMessageEvent": map[string]any{"type": "text_delta", "contentIndex": 1, "delta": "stale delta"}}) + writeFakePiEvent(map[string]any{"type": "message_end", "message": map[string]any{"role": "assistant", "content": []any{map[string]any{"type": "thinking", "thinking": "final reasoning"}, map[string]any{"type": "text", "text": "authoritative reply"}}}}) + writeFakePiEvent(map[string]any{"type": "queue_update", "steering": []string{}, "followUp": []string{}}) + appendEntry("assistant", "authoritative reply") + writeFakePiEvent(map[string]any{"type": "agent_end", "messages": []any{}}) + writeFakePiEvent(map[string]any{"type": "agent_settled"}) + } + scanner := bufio.NewScanner(os.Stdin) + for scanner.Scan() { + var command map[string]any + if err := json.Unmarshal(scanner.Bytes(), &command); err != nil { + return + } + appendFakePiLog(command) + id := command["id"] + kind, _ := command["type"].(string) + if noSession { + data := map[string]any{} + if kind == "get_available_models" { + data["models"] = []any{map[string]any{ + "provider": "openai", "id": "gpt-test", "name": "GPT Test", + "reasoning": true, "input": []string{"text", "image"}, + "contextWindow": 32000, "maxTokens": 4096, + }} + } + writeFakePiResponse(id, kind, data) + continue + } + switch kind { + case "get_state": + finishPR5() + stateSessionPath := sessionPath + if mutationSequence > 0 && os.Getenv("CODEKANBAN_FAKE_PI_INVALID_MUTATION_STATE") == "1" { + stateSessionPath = "" + } + writeFakePiResponse(id, kind, map[string]any{ + "sessionId": sessionID, "sessionFile": stateSessionPath, + "model": map[string]any{"provider": modelProvider, "id": modelID, "name": "Fake"}, + "thinkingLevel": thinkingLevel, "isStreaming": false, + }) + case "get_commands": + writeFakePiResponse(id, kind, map[string]any{"commands": []any{map[string]any{ + "name": piBridgeCommandName, "source": "extension", + "sourceInfo": map[string]any{ + "path": bridgePath, "source": "cli", "scope": "temporary", "origin": "top-level", + }, + }}}) + case "get_entries": + var leaf any + if leafID != "" { + leaf = leafID + } + writeFakePiResponse(id, kind, map[string]any{"entries": entries, "leafId": leaf}) + case "get_tree": + var leaf any + if leafID != "" { + leaf = leafID + } + writeFakePiResponse(id, kind, map[string]any{"tree": buildFakePiTree(entries), "leafId": leaf}) + case "get_session_stats": + writeFakePiResponse(id, kind, map[string]any{ + "sessionId": sessionID, "sessionFile": sessionPath, + "tokens": map[string]any{"input": 10, "output": 5, "cacheRead": 2, "cacheWrite": 0, "total": 17}, + "cost": 0.01, + "contextUsage": map[string]any{"tokens": 17, "contextWindow": 32000, "percent": 0.1}, + }) + case "fork", "clone": + targetLeafID := leafID + selectedText := "" + if kind == "fork" { + entryID, _ := command["entryId"].(string) + target := findFakePiEntry(entries, entryID) + if target == nil || fakePiEntryRole(target) != "user" { + writeFakePiError(id, kind, "invalid fork target") + continue + } + targetLeafID = fakePiParentID(target) + selectedText = fakePiEntryText(target) + } + if kind == "clone" && targetLeafID == "" { + writeFakePiError(id, kind, "cannot clone empty session") + continue + } + mutationSequence++ + newPath, newID, newEntries, newLeaf, err := createFakePiBranchedSession( + sessionPath, sessionID, cwd, entries, targetLeafID, mutationSequence, + ) + if err != nil { + writeFakePiError(id, kind, err.Error()) + continue + } + sessionPath, sessionID, entries, leafID = newPath, newID, newEntries, newLeaf + entrySequence = len(entries) + if kind == "fork" { + writeFakePiResponse(id, kind, map[string]any{"text": selectedText, "cancelled": false}) + } else { + writeFakePiResponse(id, kind, map[string]any{"cancelled": false}) + } + case "set_model": + modelProvider, _ = command["provider"].(string) + modelID, _ = command["modelId"].(string) + writeFakePiResponse(id, kind, map[string]any{"provider": modelProvider, "id": modelID}) + case "set_thinking_level": + thinkingLevel, _ = command["level"].(string) + writeFakePiResponse(id, kind, nil) + case "prompt": + if err := ensureFakePiSessionFile(sessionPath, sessionID, cwd); err != nil { + writeFakePiError(id, kind, "failed to persist fake session") + continue + } + if delay, err := time.ParseDuration(os.Getenv("CODEKANBAN_FAKE_PI_PROMPT_ACK_DELAY")); err == nil && delay > 0 { + time.Sleep(delay) + } + writeFakePiResponse(id, kind, nil) + message, _ := command["message"].(string) + if payload, ok := parseFakePiBridgeCommand(message); ok { + target := findFakePiEntry(entries, payload.TargetID) + if target == nil { + continue + } + targetType, _ := target["type"].(string) + targetParent, _ := target["parentId"].(string) + leafID = payload.TargetID + if targetType == "custom_message" || fakePiEntryRole(target) == "user" { + leafID = targetParent + } + if payload.Summarize { + appendNativeEntry(map[string]any{"type": "branch_summary", "summary": "fake branch summary"}) + } + appendNativeEntry(map[string]any{ + "type": "custom", "customType": piBridgeMarkerType, + "data": map[string]any{"targetId": payload.TargetID, "summarize": payload.Summarize, "nonce": payload.Nonce}, + }) + continue + } + if message == "pr6-tree-seed" { + appendEntry("user", message) + branchParent := leafID + appendEntry("assistant", "abandoned branch") + leafID = branchParent + appendEntry("assistant", "active branch") + writeFakePiEvent(map[string]any{"type": "agent_start"}) + writeFakePiEvent(map[string]any{"type": "message_start", "message": map[string]any{"role": "assistant", "content": []any{}}}) + writeFakePiEvent(map[string]any{"type": "message_update", "assistantMessageEvent": map[string]any{"type": "text_delta", "contentIndex": 0, "delta": "active branch"}}) + writeFakePiEvent(map[string]any{"type": "message_end", "message": map[string]any{"role": "assistant", "content": []any{map[string]any{"type": "text", "text": "active branch"}}}}) + writeFakePiEvent(map[string]any{"type": "agent_end", "messages": []any{}}) + writeFakePiEvent(map[string]any{"type": "agent_settled"}) + continue + } + if message == "hold" { + continue + } + if message == "pr5-events" { + pr5Active = true + appendEntry("user", message) + writeFakePiEvent(map[string]any{"type": "agent_start"}) + writeFakePiEvent(map[string]any{"type": "message_start", "message": map[string]any{"role": "assistant", "content": []any{}}}) + writeFakePiEvent(map[string]any{"type": "message_update", "assistantMessageEvent": map[string]any{"type": "thinking_start", "contentIndex": 0}}) + writeFakePiEvent(map[string]any{"type": "message_update", "assistantMessageEvent": map[string]any{"type": "thinking_delta", "contentIndex": 0, "delta": "live reasoning"}}) + writeFakePiEvent(map[string]any{"type": "message_update", "assistantMessageEvent": map[string]any{"type": "thinking_end", "contentIndex": 0, "content": "stale reasoning"}}) + writeFakePiEvent(map[string]any{"type": "tool_execution_start", "toolCallId": "parallel-a", "toolName": "Write", "args": map[string]any{"path": "a.txt"}}) + writeFakePiEvent(map[string]any{"type": "tool_execution_start", "toolCallId": "parallel-b", "toolName": "Read", "args": map[string]any{"path": "b.txt"}}) + writeFakePiEvent(map[string]any{"type": "extension_ui_request", "id": "confirm-1", "method": "confirm", "title": "Confirm action", "message": "Proceed?", "timeout": 5000}) + continue + } + if message == "timeout-confirm" { + appendEntry("user", message) + writeFakePiEvent(map[string]any{"type": "agent_start"}) + writeFakePiEvent(map[string]any{"type": "message_start", "message": map[string]any{"role": "assistant", "content": []any{}}}) + writeFakePiEvent(map[string]any{"type": "extension_ui_request", "id": "timeout-confirm-1", "method": "confirm", "title": "Timed confirmation", "message": "Respond before timeout", "timeout": 20}) + continue + } + if message == "settle-with-dialog" { + appendEntry("user", message) + writeFakePiEvent(map[string]any{"type": "agent_start"}) + writeFakePiEvent(map[string]any{"type": "message_start", "message": map[string]any{"role": "assistant", "content": []any{}}}) + writeFakePiEvent(map[string]any{"type": "extension_ui_request", "id": "settle-dialog-1", "method": "confirm", "title": "Unanswered confirmation", "message": "Pi settled early"}) + writeFakePiEvent(map[string]any{"type": "agent_settled"}) + continue + } + if message == "hold-dialog" { + appendEntry("user", message) + writeFakePiEvent(map[string]any{"type": "agent_start"}) + writeFakePiEvent(map[string]any{"type": "message_start", "message": map[string]any{"role": "assistant", "content": []any{}}}) + writeFakePiEvent(map[string]any{"type": "extension_ui_request", "id": "abort-dialog-1", "method": "confirm", "title": "Abort confirmation", "message": "Waiting for abort"}) + continue + } + appendEntry("assistant", "fake reply") + writeFakePiEvent(map[string]any{"type": "agent_start"}) + writeFakePiEvent(map[string]any{"type": "message_start", "message": map[string]any{"role": "assistant", "content": []any{}}}) + writeFakePiEvent(map[string]any{"type": "message_update", "assistantMessageEvent": map[string]any{"type": "text_delta", "delta": "fake reply"}}) + writeFakePiEvent(map[string]any{"type": "message_end", "message": map[string]any{"role": "assistant", "content": []any{map[string]any{"type": "text", "text": "fake reply"}}}}) + writeFakePiEvent(map[string]any{"type": "agent_end", "messages": []any{}}) + writeFakePiEvent(map[string]any{"type": "agent_settled"}) + case "compact": + writeFakePiEvent(map[string]any{"type": "compaction_start", "reason": "manual"}) + appendNativeEntry(map[string]any{ + "type": "compaction", "summary": "manual compact summary", "tokensBefore": 17, + "firstKeptEntryId": nil, "retainedTail": []any{}, + }) + result := map[string]any{"summary": "manual compact summary", "tokensBefore": 17} + writeFakePiEvent(map[string]any{"type": "compaction_end", "result": result, "aborted": false}) + writeFakePiResponse(id, kind, result) + case "steer", "follow_up": + message, _ := command["message"].(string) + appendEntry("user", message) + if kind == "steer" { + pr5Steered = true + } else { + pr5FollowedUp = true + } + steering := []string{} + followUp := []string{} + if pr5Steered { + steering = append(steering, "redirect while active") + } + if pr5FollowedUp { + followUp = append(followUp, "queue while active") + } + writeFakePiEvent(map[string]any{"type": "queue_update", "steering": steering, "followUp": followUp}) + writeFakePiResponse(id, kind, nil) + case "extension_ui_response": + requestID, _ := command["id"].(string) + switch requestID { + case "confirm-1": + writeFakePiEvent(map[string]any{"type": "extension_ui_request", "id": "input-1", "method": "input", "title": "Provide value", "placeholder": "value", "timeout": 5000}) + case "input-1": + pr5InputAnswered = true + finishPR5() + case "timeout-confirm-1": + appendEntry("assistant", "continued after timeout") + writeFakePiEvent(map[string]any{"type": "message_update", "assistantMessageEvent": map[string]any{"type": "text_delta", "contentIndex": 0, "delta": "continued after timeout"}}) + writeFakePiEvent(map[string]any{"type": "message_end", "message": map[string]any{"role": "assistant", "content": []any{map[string]any{"type": "text", "text": "continued after timeout"}}}}) + writeFakePiEvent(map[string]any{"type": "agent_end", "messages": []any{}}) + writeFakePiEvent(map[string]any{"type": "agent_settled"}) + } + case "abort": + writeFakePiResponse(id, kind, nil) + writeFakePiEvent(map[string]any{"type": "agent_settled"}) + default: + writeFakePiError(id, kind, "unknown command") + } + } +} + +func readFakePiSessionHeaderID(sessionPath string) string { + file, err := os.Open(sessionPath) + if err != nil { + return "" + } + defer file.Close() + scanner := bufio.NewScanner(file) + if !scanner.Scan() { + return "" + } + var header struct { + ID string `json:"id"` + } + if json.Unmarshal(scanner.Bytes(), &header) != nil { + return "" + } + return strings.TrimSpace(header.ID) +} + +func readFakePiSessionEntries(sessionPath string) ([]map[string]any, string) { + file, err := os.Open(sessionPath) + if err != nil { + return nil, "" + } + defer file.Close() + entries := make([]map[string]any, 0) + leafID := "" + scanner := bufio.NewScanner(file) + lineIndex := 0 + for scanner.Scan() { + lineIndex++ + if lineIndex == 1 { + continue + } + entry := map[string]any{} + if json.Unmarshal(scanner.Bytes(), &entry) != nil { + continue + } + id, _ := entry["id"].(string) + if strings.TrimSpace(id) == "" { + continue + } + entries = append(entries, entry) + leafID = id + } + return entries, leafID +} + +func buildFakePiTree(entries []map[string]any) []any { + nodes := make(map[string]map[string]any, len(entries)) + order := make([]string, 0, len(entries)) + for _, entry := range entries { + id, _ := entry["id"].(string) + if strings.TrimSpace(id) == "" { + continue + } + nodes[id] = map[string]any{"entry": entry, "children": []any{}} + order = append(order, id) + } + roots := make([]any, 0, 1) + for _, id := range order { + entry := nodes[id]["entry"].(map[string]any) + parentID := stringValue(entry["parentId"]) + if pointer, ok := entry["parentId"].(*string); ok { + parentID = pointerString(pointer) + } + parent := nodes[parentID] + if parentID == "" || parent == nil || parentID == id { + roots = append(roots, nodes[id]) + continue + } + children := parent["children"].([]any) + parent["children"] = append(children, nodes[id]) + } + return roots +} + +func parseFakePiBridgeCommand(message string) (piBridgeMarkerData, bool) { + prefix := "/" + piBridgeCommandName + " " + if !strings.HasPrefix(message, prefix) { + return piBridgeMarkerData{}, false + } + encoded := strings.TrimSpace(strings.TrimPrefix(message, prefix)) + decoded, err := base64.RawURLEncoding.DecodeString(encoded) + if err != nil { + return piBridgeMarkerData{}, false + } + var payload piBridgeMarkerData + if json.Unmarshal(decoded, &payload) != nil || strings.TrimSpace(payload.TargetID) == "" || strings.TrimSpace(payload.Nonce) == "" { + return piBridgeMarkerData{}, false + } + return payload, true +} + +func findFakePiEntry(entries []map[string]any, id string) map[string]any { + for _, entry := range entries { + if entry["id"] == id { + return entry + } + } + return nil +} + +func fakePiEntryRole(entry map[string]any) string { + message, _ := entry["message"].(map[string]any) + role, _ := message["role"].(string) + return role +} + +func fakePiParentID(entry map[string]any) string { + if entry == nil { + return "" + } + if parent, ok := entry["parentId"].(*string); ok { + return pointerString(parent) + } + return stringValue(entry["parentId"]) +} + +func fakePiEntryText(entry map[string]any) string { + message, _ := entry["message"].(map[string]any) + content := message["content"] + if text, ok := content.(string); ok { + return text + } + parts, _ := content.([]any) + var builder strings.Builder + for _, partValue := range parts { + part, _ := partValue.(map[string]any) + if part["type"] == "text" { + builder.WriteString(stringValue(part["text"])) + } + } + return builder.String() +} + +func createFakePiBranchedSession( + sourcePath string, + sourceID string, + cwd string, + entries []map[string]any, + targetLeafID string, + sequence int, +) (string, string, []map[string]any, string, error) { + byID := make(map[string]map[string]any, len(entries)) + for _, entry := range entries { + id := stringValue(entry["id"]) + if id != "" { + byID[id] = entry + } + } + chain := make([]map[string]any, 0) + seen := make(map[string]struct{}) + for id := strings.TrimSpace(targetLeafID); id != ""; { + if _, duplicate := seen[id]; duplicate { + return "", "", nil, "", errors.New("fake Pi branch contains a cycle") + } + seen[id] = struct{}{} + entry := byID[id] + if entry == nil { + return "", "", nil, "", errors.New("fake Pi branch target is missing") + } + chain = append(chain, entry) + id = fakePiParentID(entry) + } + for left, right := 0, len(chain)-1; left < right; left, right = left+1, right-1 { + chain[left], chain[right] = chain[right], chain[left] + } + + mutationID := make([]byte, 16) + if _, err := rand.Read(mutationID); err != nil { + return "", "", nil, "", err + } + newID := fmt.Sprintf("%s-branch-%d-%s", sourceID, sequence, hex.EncodeToString(mutationID)) + newPath := filepath.Join(filepath.Dir(sourcePath), newID+".jsonl") + header := map[string]any{ + "type": "session", "version": 3, "id": newID, + "timestamp": time.Now().UTC().Format(time.RFC3339Nano), "cwd": cwd, + "parentSession": sourcePath, + } + projected := make([]map[string]any, 0, len(chain)) + lines := make([][]byte, 0, len(chain)+1) + encodedHeader, _ := json.Marshal(header) + lines = append(lines, encodedHeader) + parentID := "" + for _, original := range chain { + encoded, _ := json.Marshal(original) + clone := map[string]any{} + if err := json.Unmarshal(encoded, &clone); err != nil { + return "", "", nil, "", err + } + clone["parentId"] = nilIfEmpty(parentID) + parentID = stringValue(clone["id"]) + projected = append(projected, clone) + encoded, _ = json.Marshal(clone) + lines = append(lines, encoded) + } + var content strings.Builder + for _, line := range lines { + content.Write(line) + content.WriteByte('\n') + } + if err := os.WriteFile(newPath, []byte(content.String()), 0o600); err != nil { + return "", "", nil, "", err + } + return newPath, newID, projected, parentID, nil +} + +func ensureFakePiSessionFile(sessionPath, sessionID, cwd string) error { + if _, err := os.Stat(sessionPath); err == nil { + return nil + } else if !os.IsNotExist(err) { + return err + } + header := fmt.Sprintf("{\"type\":\"session\",\"version\":3,\"id\":%q,\"timestamp\":%q,\"cwd\":%q}\n", sessionID, time.Now().UTC().Format(time.RFC3339Nano), cwd) + return os.WriteFile(sessionPath, []byte(header), 0o600) +} + +func TestPiRuntimeValidatesModelAndReasoningControls(t *testing.T) { + if _, _, err := splitPiModel("anthropic/claude-sonnet-4"); err != nil { + t.Fatalf("valid Pi model: %v", err) + } + if _, _, err := splitPiModel("claude-sonnet-4"); err == nil { + t.Fatal("model without provider should fail") + } + if err := validatePiReasoningEffort(ReasoningEffortMax); err != nil { + t.Fatalf("max reasoning should be supported: %v", err) + } + if err := validatePiReasoningEffort(ReasoningEffortUltra); err == nil { + t.Fatal("ultra reasoning should be rejected before starting a prompt") + } +} + +func TestValidatePiRuntimeStateRejectsSessionOutsideConfiguredRoot(t *testing.T) { + root := t.TempDir() + outside := filepath.Join(t.TempDir(), "outside.jsonl") + projectPath := t.TempDir() + if err := os.WriteFile(outside, []byte(fmt.Sprintf("{\"type\":\"session\",\"id\":\"native-outside\",\"cwd\":%q}\n", filepath.ToSlash(projectPath))), 0o644); err != nil { + t.Fatal(err) + } + t.Setenv("PI_CODING_AGENT_SESSION_DIR", root) + err := validatePiRuntimeState(tables.WebSessionTable{Cwd: projectPath}, piRPCState{ + SessionID: "native-outside", + SessionFile: outside, + }) + if err == nil || !strings.Contains(err.Error(), "outside the configured session root") { + t.Fatalf("expected session-root rejection, got %v", err) + } +} + +func TestManagerImportPiSessionValidatesNativeIdentity(t *testing.T) { + cleanup := initTestDB(t) + defer cleanup() + project := seedProject(t) + sessionRoot := t.TempDir() + t.Setenv("PI_CODING_AGENT_SESSION_DIR", sessionRoot) + sessionID := "imported-pi-session" + sessionPath := filepath.Join(sessionRoot, "imported.jsonl") + header := fmt.Sprintf("{\"type\":\"session\",\"version\":3,\"id\":%q,\"timestamp\":%q,\"cwd\":%q}\n", sessionID, time.Now().UTC().Format(time.RFC3339Nano), project.Path) + content := header + + "{\"type\":\"message\",\"id\":\"u1\",\"parentId\":null,\"timestamp\":\"2026-05-01T01:00:00Z\",\"message\":{\"role\":\"user\",\"content\":\"active prompt\"}}\n" + + "{\"type\":\"message\",\"id\":\"a1\",\"parentId\":\"u1\",\"timestamp\":\"2026-05-01T01:00:01Z\",\"message\":{\"role\":\"assistant\",\"content\":[{\"type\":\"text\",\"text\":\"active reply\"}]}}\n" + + "{\"type\":\"message\",\"id\":\"abandoned\",\"parentId\":\"a1\",\"timestamp\":\"2026-05-01T01:00:02Z\",\"message\":{\"role\":\"assistant\",\"content\":[{\"type\":\"text\",\"text\":\"abandoned branch\"}]}}\n" + + "{\"type\":\"message\",\"id\":\"leaf\",\"parentId\":\"a1\",\"timestamp\":\"2026-05-01T01:00:03Z\",\"message\":{\"role\":\"user\",\"content\":\"active leaf\"}}\n" + if err := os.WriteFile(sessionPath, []byte(content), 0o600); err != nil { + t.Fatalf("write Pi session fixture: %v", err) + } + info, err := os.Stat(sessionPath) + if err != nil { + t.Fatalf("stat Pi session fixture: %v", err) + } + source := tables.AISessionTable{ + SessionID: sessionID, Type: tables.AISessionTypePi, ProjectPath: project.Path, + FilePath: sessionPath, Title: "Imported Pi", Model: "openai/gpt-test", + SessionStartedAt: info.ModTime(), FileModTime: info.ModTime(), FileSize: info.Size(), + } + source.Init() + if err := model.GetDB().Create(&source).Error; err != nil { + t.Fatalf("seed Pi history record: %v", err) + } + unavailable, err := NewManager(Config{ + DataDir: t.TempDir(), PiPath: filepath.Join(t.TempDir(), "missing-pi"), + }, zap.NewNop()) + if err != nil { + t.Fatalf("create unavailable manager: %v", err) + } + if _, err := unavailable.ImportPiSessionBySessionID(context.Background(), project.ID, sessionID); err == nil || !strings.Contains(err.Error(), errPiWebSessionUnavailable) { + t.Fatalf("expected unavailable import rejection, got %v", err) + } + var importedCount int64 + if err := model.GetDB().Model(&tables.WebSessionTable{}). + Where("project_id = ? AND agent = ?", project.ID, string(AgentPi)).Count(&importedCount).Error; err != nil || importedCount != 0 { + t.Fatalf("unavailable import created %d sessions, err=%v", importedCount, err) + } + + t.Setenv("CODEKANBAN_FAKE_PI_RUNTIME", "1") + t.Setenv("CODEKANBAN_FAKE_PI_SESSION", sessionPath) + manager, err := NewManager(Config{ + DataDir: t.TempDir(), PiPath: fmt.Sprintf("%q -test.run=^TestPiRuntimeFakeProcess$ --", os.Args[0]), + PiRuntimeIdleTTL: time.Minute, + }, zap.NewNop()) + if err != nil { + t.Fatalf("NewManager returned error: %v", err) + } + defer manager.StopProjectPiRuntimes(project.ID) + if _, err := manager.TrustProjectForPi(context.Background(), project.ID); err != nil { + t.Fatalf("trust project for Pi: %v", err) + } + result, err := manager.ImportPiSessionBySessionID(context.Background(), project.ID, sessionID) + if err != nil { + t.Fatalf("ImportPiSessionBySessionID returned error: %v", err) + } + if !result.Created || result.Session.Agent != AgentPi { + t.Fatalf("unexpected imported Pi session: %+v", result.Session) + } + record, err := manager.GetSession(context.Background(), result.Session.ID) + if err != nil { + t.Fatalf("GetSession returned error: %v", err) + } + if record.Backend != string(SessionBackendPiRPC) || pointerString(result.Session.NativeSessionID) != sessionID || pointerString(result.Session.ThreadPath) != sessionPath { + t.Fatalf("unexpected imported native identity: summary=%+v record=%+v", result.Session, record) + } + if !result.Synced || pointerString(record.NativeLeafID) != "leaf" || !historyContainsText(result.History, "active prompt") || !historyContainsText(result.History, "active leaf") || historyContainsText(result.History, "abandoned branch") { + t.Fatalf("unexpected imported Pi projection: synced=%v leaf=%q history=%#v", result.Synced, pointerString(record.NativeLeafID), result.History.Items) + } + reused, err := manager.ImportPiSessionBySessionID(context.Background(), project.ID, sessionID) + if err != nil || !reused.Reused || reused.Session.ID != result.Session.ID { + t.Fatalf("expected duplicate import to reuse %q, got %+v err=%v", result.Session.ID, reused, err) + } + + outsidePath := filepath.Join(t.TempDir(), "outside.jsonl") + outsideID := "outside-pi-session" + outsideHeader := fmt.Sprintf("{\"type\":\"session\",\"version\":3,\"id\":%q,\"timestamp\":%q,\"cwd\":%q}\n", outsideID, time.Now().UTC().Format(time.RFC3339Nano), project.Path) + if err := os.WriteFile(outsidePath, []byte(outsideHeader), 0o600); err != nil { + t.Fatalf("write outside Pi fixture: %v", err) + } + outsideInfo, _ := os.Stat(outsidePath) + outsideSource := tables.AISessionTable{ + SessionID: outsideID, Type: tables.AISessionTypePi, ProjectPath: project.Path, + FilePath: outsidePath, SessionStartedAt: outsideInfo.ModTime(), FileModTime: outsideInfo.ModTime(), FileSize: outsideInfo.Size(), + } + outsideSource.Init() + if err := model.GetDB().Create(&outsideSource).Error; err != nil { + t.Fatalf("seed outside Pi history record: %v", err) + } + if _, err := manager.ImportPiSessionBySessionID(context.Background(), project.ID, outsideID); err == nil || !strings.Contains(err.Error(), "outside the configured session root") { + t.Fatalf("expected outside-root import rejection, got %v", err) + } +} + +func TestManagerPiRPCSendReusesRuntimeAndPersistsIdentity(t *testing.T) { + cleanup := initTestDB(t) + defer cleanup() + project := seedProject(t) + fakeDir := t.TempDir() + sessionPath := filepath.Join(fakeDir, "fake-session.jsonl") + logPath := filepath.Join(fakeDir, "commands.jsonl") + t.Setenv("CODEKANBAN_FAKE_PI_RUNTIME", "1") + t.Setenv("CODEKANBAN_FAKE_PI_SESSION", sessionPath) + t.Setenv("PI_CODING_AGENT_SESSION_DIR", fakeDir) + t.Setenv("CODEKANBAN_FAKE_PI_LOG", logPath) + manager, err := NewManager(Config{ + DataDir: t.TempDir(), + PiPath: fmt.Sprintf("%q -test.run=^TestPiRuntimeFakeProcess$ --", os.Args[0]), + PiRuntimeIdleTTL: time.Minute, + }, zap.NewNop()) + if err != nil { + t.Fatalf("NewManager returned error: %v", err) + } + defer manager.StopProjectPiRuntimes(project.ID) + if _, err := manager.TrustProjectForPi(context.Background(), project.ID); err != nil { + t.Fatalf("TrustProjectForPi returned error: %v", err) + } + config := manager.GetWebSessionRuntimeConfig() + if !config.SupportsPiWebSession { + t.Fatalf("fake Pi runtime unavailable: hasPi=%v version=%v compatible=%v diagnostics=%q", config.HasPi, config.PiVersion, config.PiRPCCompatible, config.PiDiagnostics) + } + created, err := manager.CreateSession(context.Background(), CreateParams{ + ProjectID: project.ID, Agent: AgentPi, Model: "openai/gpt-test", + ReasoningEffort: ReasoningEffortHigh, + }) + if err != nil { + t.Fatalf("CreateSession returned error: %v", err) + } + image, err := manager.saveAttachmentBytes("pixel.png", "image/png", []byte("\x89PNG\r\n\x1a\n\x00\x00\x00\x00")) + if err != nil { + t.Fatalf("save image attachment: %v", err) + } + if err := manager.SendMessage(context.Background(), created.ID, "first", []string{image.ID}); err != nil { + t.Fatalf("SendMessage with image returned error: %v", err) + } + waitForSessionToSettle(t, manager, created.ID) + if err := manager.SendMessage(context.Background(), created.ID, "second", nil); err != nil { + t.Fatalf("second SendMessage returned error: %v", err) + } + waitForSessionToSettle(t, manager, created.ID) + manager.StopSessionPiRuntime(created.ID) + manager.cfg.PiRuntimeIdleTTL = 25 * time.Millisecond + if err := manager.SendMessage(context.Background(), created.ID, "restored", nil); err != nil { + t.Fatalf("restored SendMessage returned error: %v", err) + } + waitForSessionToSettle(t, manager, created.ID) + deadline := time.Now().Add(2 * time.Second) + for time.Now().Before(deadline) { + manager.piRuntimeMu.Lock() + _, active := manager.piRuntimes[created.ID] + manager.piRuntimeMu.Unlock() + if !active { + break + } + time.Sleep(10 * time.Millisecond) + } + manager.piRuntimeMu.Lock() + _, idleRuntimeStillActive := manager.piRuntimes[created.ID] + manager.piRuntimeMu.Unlock() + if idleRuntimeStillActive { + t.Fatal("idle Pi runtime was not evicted after its TTL") + } + + record, err := manager.GetSession(context.Background(), created.ID) + if err != nil { + t.Fatalf("GetSession returned error: %v", err) + } + if pointerString(record.NativeSessionID) != "fake-pi-session" || !samePiRuntimePath(pointerString(record.ThreadPath), sessionPath) { + t.Fatalf("unexpected Pi identity: native=%q path=%q status=%q error=%q", pointerString(record.NativeSessionID), pointerString(record.ThreadPath), record.Status, pointerString(record.LastError)) + } + if pointerString(record.NativeLeafID) == "" || pointerString(record.SourceRevision) == "" { + t.Fatalf("missing Pi leaf/revision: leaf=%q revision=%q", pointerString(record.NativeLeafID), pointerString(record.SourceRevision)) + } + if record.Model != "openai/gpt-test" || record.ReasoningEffort != string(ReasoningEffortHigh) { + t.Fatalf("model/thinking = %q/%q", record.Model, record.ReasoningEffort) + } + window, err := manager.History(context.Background(), created.ID, DefaultHistoryWindow, nil) + if err != nil { + t.Fatalf("History returned error: %v", err) + } + if !historyContainsText(window, "fake reply") { + t.Fatalf("history does not contain fake reply: %#v", window.Items) + } + commands := readFakePiLog(t, logPath) + starts := 0 + prompts := 0 + imagePrompt := false + restoredStart := false + for _, command := range commands { + if _, ok := command["startup"]; ok { + starts++ + if args, ok := command["args"].([]any); ok { + for _, arg := range args { + if arg == "--session" { + restoredStart = true + } + } + } + } + if command["type"] == "prompt" { + prompts++ + if images, ok := command["images"].([]any); ok && len(images) == 1 { + image, _ := images[0].(map[string]any) + imagePrompt = image["data"] == "iVBORw0KGgoAAAAA" && image["mimeType"] == "image/png" + } + } + } + if starts != 2 || prompts != 3 || !imagePrompt || !restoredStart { + t.Fatalf("runtime starts=%d prompts=%d imagePrompt=%v restoredStart=%v, commands=%#v", starts, prompts, imagePrompt, restoredStart, commands) + } +} + +func TestManagerPiRPCTreeNavigatePersistsAcrossRuntimeRestore(t *testing.T) { + cleanup := initTestDB(t) + defer cleanup() + project := seedProject(t) + fakeDir := t.TempDir() + sessionPath := filepath.Join(fakeDir, "tree-session.jsonl") + logPath := filepath.Join(fakeDir, "tree-commands.jsonl") + t.Setenv("CODEKANBAN_FAKE_PI_RUNTIME", "1") + t.Setenv("CODEKANBAN_FAKE_PI_SESSION", sessionPath) + t.Setenv("PI_CODING_AGENT_SESSION_DIR", fakeDir) + t.Setenv("CODEKANBAN_FAKE_PI_LOG", logPath) + manager, err := NewManager(Config{ + DataDir: t.TempDir(), PiPath: fmt.Sprintf("%q -test.run=^TestPiRuntimeFakeProcess$ --", os.Args[0]), + PiRuntimeIdleTTL: time.Minute, + }, zap.NewNop()) + if err != nil { + t.Fatal(err) + } + defer manager.StopProjectPiRuntimes(project.ID) + if _, err := manager.TrustProjectForPi(context.Background(), project.ID); err != nil { + t.Fatal(err) + } + created, err := manager.CreateSession(context.Background(), CreateParams{ProjectID: project.ID, Agent: AgentPi}) + if err != nil { + t.Fatal(err) + } + if err := manager.SendMessage(context.Background(), created.ID, "pr6-tree-seed", nil); err != nil { + t.Fatal(err) + } + waitForSessionToSettle(t, manager, created.ID) + + tree, err := manager.GetPiSessionTree(context.Background(), created.ID) + if err != nil { + t.Fatalf("GetPiSessionTree: %v", err) + } + if tree.SessionID != "fake-pi-session" || tree.Revision == "" || len(tree.Nodes) != 3 { + t.Fatalf("unexpected Pi tree snapshot: %#v", tree) + } + var userID, abandonedID, activeID string + for _, node := range tree.Nodes { + switch node.Preview { + case "pr6-tree-seed": + userID = node.ID + case "abandoned branch": + abandonedID = node.ID + case "active branch": + activeID = node.ID + } + } + if userID == "" || abandonedID == "" || activeID == "" || pointerString(tree.LeafID) != activeID { + t.Fatalf("tree nodes/leaf mismatch: user=%q abandoned=%q active=%q tree=%#v", userID, abandonedID, activeID, tree) + } + originalRevision := tree.Revision + manager.replacePiNativeQueuedInputs(created.ID, []string{"native queued"}, nil) + if _, err := manager.GetPiSessionTree(context.Background(), created.ID); err != nil { + t.Fatalf("read tree with native queued input: %v", err) + } + if _, err := manager.NavigatePiSessionTree(context.Background(), created.ID, PiTreeNavigateInput{ + TargetID: abandonedID, Revision: originalRevision, + }); err == nil || !strings.Contains(err.Error(), "messages are pending") { + t.Fatalf("expected native queue to block navigation, got %v", err) + } + if _, err := manager.ForkPiSessionTree(context.Background(), created.ID, PiTreeForkInput{ + TargetID: userID, Revision: originalRevision, + }); err == nil || !strings.Contains(err.Error(), "messages are pending") { + t.Fatalf("expected native queue to block fork, got %v", err) + } + if _, err := manager.ClonePiSessionTree(context.Background(), created.ID, PiTreeCloneInput{ + Revision: originalRevision, + }); err == nil || !strings.Contains(err.Error(), "messages are pending") { + t.Fatalf("expected native queue to block clone, got %v", err) + } + manager.clearPiNativeQueuedInputs(created.ID) + + result, err := manager.NavigatePiSessionTree(context.Background(), created.ID, PiTreeNavigateInput{ + TargetID: abandonedID, Revision: originalRevision, + }) + if err != nil { + t.Fatalf("NavigatePiSessionTree: %v", err) + } + if pointerString(result.Tree.LeafID) != abandonedID || result.Tree.Revision == originalRevision { + t.Fatalf("navigation did not advance tree identity: %#v", result) + } + window, err := manager.History(context.Background(), created.ID, DefaultHistoryWindow, nil) + if err != nil { + t.Fatal(err) + } + if !historyContainsText(window, "abandoned branch") || historyContainsText(window, "active branch") { + t.Fatalf("navigation did not replace the active branch timeline: %#v", window.Items) + } + for _, item := range window.Items { + if strings.Contains(item.Text, piBridgeMarkerType) { + t.Fatalf("bridge marker leaked into history: %#v", item) + } + } + if _, err := manager.NavigatePiSessionTree(context.Background(), created.ID, PiTreeNavigateInput{ + TargetID: activeID, Revision: originalRevision, + }); !errors.Is(err, ErrPiTreeRevisionConflict) { + t.Fatalf("expected stale revision conflict, got %v", err) + } + + manager.StopSessionPiRuntime(created.ID) + restored, err := manager.GetPiSessionTree(context.Background(), created.ID) + if err != nil { + t.Fatalf("GetPiSessionTree after restore: %v", err) + } + if pointerString(restored.LeafID) != abandonedID || restored.Revision != result.Tree.Revision { + t.Fatalf("restored tree lost durable logical leaf: before=%#v after=%#v", result.Tree, restored) + } + rootResult, err := manager.NavigatePiSessionTree(context.Background(), created.ID, PiTreeNavigateInput{ + TargetID: userID, Revision: restored.Revision, + }) + if err != nil { + t.Fatalf("navigate to root user: %v", err) + } + if rootResult.EditorText != "pr6-tree-seed" || rootResult.Tree.LeafID != nil { + t.Fatalf("root-user navigation did not return editor text/reset leaf: %#v", rootResult) + } + window, err = manager.History(context.Background(), created.ID, DefaultHistoryWindow, nil) + if err != nil { + t.Fatal(err) + } + if len(window.Items) != 0 { + t.Fatalf("root-user navigation retained an active timeline: %#v", window.Items) + } + + commands := readFakePiLog(t, logPath) + starts, bridgePrompts := 0, 0 + for _, command := range commands { + if _, ok := command["startup"]; ok { + starts++ + } + if command["type"] == "prompt" && strings.HasPrefix(stringValue(command["message"]), "/"+piBridgeCommandName+" ") { + bridgePrompts++ + } + } + if starts != 2 || bridgePrompts != 2 { + t.Fatalf("unexpected navigate runtime lifecycle: starts=%d bridgePrompts=%d commands=%#v", starts, bridgePrompts, commands) + } +} + +func TestManagerPiRPCTreeForkAndCloneCreateIndependentSessions(t *testing.T) { + for _, testCase := range []struct { + name string + operation string + }{ + {name: "fork", operation: "fork"}, + {name: "clone", operation: "clone"}, + } { + t.Run(testCase.name, func(t *testing.T) { + cleanup := initTestDB(t) + defer cleanup() + project := seedProject(t) + fakeDir := t.TempDir() + sourcePath := filepath.Join(fakeDir, testCase.name+"-source.jsonl") + logPath := filepath.Join(fakeDir, testCase.name+"-commands.jsonl") + t.Setenv("CODEKANBAN_FAKE_PI_RUNTIME", "1") + t.Setenv("CODEKANBAN_FAKE_PI_SESSION", sourcePath) + t.Setenv("PI_CODING_AGENT_SESSION_DIR", fakeDir) + t.Setenv("CODEKANBAN_FAKE_PI_LOG", logPath) + manager, err := NewManager(Config{ + DataDir: t.TempDir(), PiPath: fmt.Sprintf("%q -test.run=^TestPiRuntimeFakeProcess$ --", os.Args[0]), + PiRuntimeIdleTTL: time.Minute, + }, zap.NewNop()) + if err != nil { + t.Fatal(err) + } + defer manager.StopProjectPiRuntimes(project.ID) + if _, err := manager.TrustProjectForPi(context.Background(), project.ID); err != nil { + t.Fatal(err) + } + created, err := manager.CreateSession(context.Background(), CreateParams{ + ProjectID: project.ID, Agent: AgentPi, Model: "openai/gpt-test", + ReasoningEffort: ReasoningEffortHigh, + }) + if err != nil { + t.Fatal(err) + } + if err := manager.SendMessage(context.Background(), created.ID, "pr6-tree-seed", nil); err != nil { + t.Fatal(err) + } + waitForSessionToSettle(t, manager, created.ID) + + tree, err := manager.GetPiSessionTree(context.Background(), created.ID) + if err != nil { + t.Fatal(err) + } + userID := "" + for _, node := range tree.Nodes { + if node.Preview == "pr6-tree-seed" { + userID = node.ID + } + } + if userID == "" { + t.Fatalf("missing forkable user node: %#v", tree) + } + sourceBefore := mustGetSession(t, manager, created.ID) + historyBefore, err := manager.History(context.Background(), created.ID, DefaultHistoryWindow, nil) + if err != nil { + t.Fatal(err) + } + + var result PiTreeCreateResult + if testCase.operation == "fork" { + result, err = manager.ForkPiSessionTree(context.Background(), created.ID, PiTreeForkInput{ + TargetID: userID, Revision: tree.Revision, + }) + } else { + result, err = manager.ClonePiSessionTree(context.Background(), created.ID, PiTreeCloneInput{ + Revision: tree.Revision, + }) + } + if err != nil { + t.Fatalf("%s Pi tree: %v", testCase.operation, err) + } + if result.Session.ID == "" || result.Session.ID == created.ID || result.Session.Agent != AgentPi { + t.Fatalf("unexpected target session: %#v", result.Session) + } + target := mustGetSession(t, manager, result.Session.ID) + if pointerString(target.NativeSessionID) == "" || pointerString(target.NativeSessionID) == pointerString(sourceBefore.NativeSessionID) || + pointerString(target.ThreadPath) == "" || samePiRuntimePath(pointerString(target.ThreadPath), pointerString(sourceBefore.ThreadPath)) { + t.Fatalf("tree mutation did not create an independent native identity: source=%#v target=%#v", sourceBefore, target) + } + if target.Backend != string(SessionBackendPiRPC) || target.Model != sourceBefore.Model || target.ReasoningEffort != sourceBefore.ReasoningEffort || + target.TotalInputTokens != 10 || target.TotalCachedInputTokens != 2 || target.TotalOutputTokens != 5 || target.SessionContextWindowTokens != 32000 { + t.Fatalf("target config/usage was not projected atomically: %#v", target) + } + sourceAfter := mustGetSession(t, manager, created.ID) + if pointerString(sourceAfter.NativeSessionID) != pointerString(sourceBefore.NativeSessionID) || + !samePiRuntimePath(pointerString(sourceAfter.ThreadPath), pointerString(sourceBefore.ThreadPath)) || + pointerString(sourceAfter.NativeLeafID) != pointerString(sourceBefore.NativeLeafID) || + pointerString(sourceAfter.SourceRevision) != pointerString(sourceBefore.SourceRevision) { + t.Fatalf("tree mutation changed source identity: before=%#v after=%#v", sourceBefore, sourceAfter) + } + historyAfter, err := manager.History(context.Background(), created.ID, DefaultHistoryWindow, nil) + if err != nil { + t.Fatal(err) + } + if len(historyAfter.Items) != len(historyBefore.Items) || !historyContainsText(historyAfter, "active branch") || historyContainsText(historyAfter, "abandoned branch") { + t.Fatalf("tree mutation changed source history: before=%#v after=%#v", historyBefore.Items, historyAfter.Items) + } + + targetHistory, err := manager.History(context.Background(), result.Session.ID, DefaultHistoryWindow, nil) + if err != nil { + t.Fatal(err) + } + if testCase.operation == "fork" { + if result.EditorText != "pr6-tree-seed" || len(targetHistory.Items) != 0 || result.Tree.LeafID != nil { + t.Fatalf("unexpected root fork result: result=%#v history=%#v", result, targetHistory.Items) + } + } else if result.EditorText != "" || !historyContainsText(targetHistory, "pr6-tree-seed") || + !historyContainsText(targetHistory, "active branch") || historyContainsText(targetHistory, "abandoned branch") { + t.Fatalf("unexpected clone projection: result=%#v history=%#v", result, targetHistory.Items) + } + + manager.piRuntimeMu.Lock() + _, sourceRuntimePresent := manager.piRuntimes[created.ID] + _, targetRuntimePresent := manager.piRuntimes[result.Session.ID] + manager.piRuntimeMu.Unlock() + if sourceRuntimePresent || targetRuntimePresent { + t.Fatalf("mutated Pi runtime was retained: source=%v target=%v", sourceRuntimePresent, targetRuntimePresent) + } + if _, err := manager.GetPiSessionTree(context.Background(), created.ID); err != nil { + t.Fatalf("restore source tree: %v", err) + } + if _, err := manager.GetPiSessionTree(context.Background(), result.Session.ID); err != nil { + t.Fatalf("restore target tree: %v", err) + } + + restoredPaths := map[string]bool{} + for _, command := range readFakePiLog(t, logPath) { + if _, startup := command["startup"]; !startup { + continue + } + args, _ := command["args"].([]any) + for index := 0; index+1 < len(args); index++ { + if args[index] == "--session" { + restoredPaths[filepath.Clean(stringValue(args[index+1]))] = true + } + } + } + if !restoredPaths[filepath.Clean(sourcePath)] || !restoredPaths[filepath.Clean(pointerString(target.ThreadPath))] { + t.Fatalf("source/target were not restored from independent files: %#v", restoredPaths) + } + }) + } +} + +func TestManagerPiRPCTreeWireCommands(t *testing.T) { + cleanup := initTestDB(t) + defer cleanup() + project := seedProject(t) + fakeDir := t.TempDir() + sourcePath := filepath.Join(fakeDir, "wire-tree-source.jsonl") + t.Setenv("CODEKANBAN_FAKE_PI_RUNTIME", "1") + t.Setenv("CODEKANBAN_FAKE_PI_SESSION", sourcePath) + t.Setenv("PI_CODING_AGENT_SESSION_DIR", fakeDir) + manager, err := NewManager(Config{ + DataDir: t.TempDir(), PiPath: fmt.Sprintf("%q -test.run=^TestPiRuntimeFakeProcess$ --", os.Args[0]), + PiRuntimeIdleTTL: time.Minute, + }, zap.NewNop()) + if err != nil { + t.Fatal(err) + } + defer manager.StopProjectPiRuntimes(project.ID) + if _, err := manager.TrustProjectForPi(context.Background(), project.ID); err != nil { + t.Fatal(err) + } + created, err := manager.CreateSession(context.Background(), CreateParams{ + ProjectID: project.ID, Agent: AgentPi, Model: "openai/gpt-test", + }) + if err != nil { + t.Fatal(err) + } + if err := manager.SendMessage(context.Background(), created.ID, "pr6-tree-seed", nil); err != nil { + t.Fatal(err) + } + waitForSessionToSettle(t, manager, created.ID) + + conn := &captureWSConn{} + client := manager.RegisterCommandClient(conn) + defer manager.UnregisterClient(client) + send := func(requestID, operation, payload string) wireFrame { + t.Helper() + conn.frames = nil + command := fmt.Sprintf(`{"v":1,"k":"cmd","rid":%q,"sid":%q,"op":%q`, requestID, created.ID, operation) + if payload != "" { + command += `,"p":` + payload + } + command += `}` + if err := manager.HandleCommand(context.Background(), client, []byte(command)); err != nil { + t.Fatalf("HandleCommand %s: %v", operation, err) + } + if len(conn.frames) != 1 { + t.Fatalf("%s returned %d frames: %#v", operation, len(conn.frames), conn.frames) + } + return conn.frames[0] + } + decodePayload := func(frame wireFrame, target any) { + t.Helper() + data, err := json.Marshal(frame.Payload) + if err != nil { + t.Fatalf("marshal %s payload: %v", frame.Operation, err) + } + if err := json.Unmarshal(data, target); err != nil { + t.Fatalf("decode %s payload: %v; payload=%s", frame.Operation, err, data) + } + } + + getFrame := send("tree-get", "tree_get", "") + if getFrame.Kind != "ack" || getFrame.Operation != "tree_get" || getFrame.SessionID != created.ID { + t.Fatalf("unexpected tree_get frame: %#v", getFrame) + } + var tree PiTreeSnapshot + decodePayload(getFrame, &tree) + var userID, abandonedID string + for _, node := range tree.Nodes { + switch node.Preview { + case "pr6-tree-seed": + userID = node.ID + case "abandoned branch": + abandonedID = node.ID + } + } + if tree.Revision == "" || userID == "" || abandonedID == "" { + t.Fatalf("tree_get omitted required tree state: %#v", tree) + } + + staleFrame := send("tree-nav-stale", "tree_nav", fmt.Sprintf(`{"tid":%q,"rev":"stale","sum":false}`, abandonedID)) + if staleFrame.Kind != "err" || staleFrame.Code != "conflict" { + t.Fatalf("stale tree_nav did not return conflict: %#v", staleFrame) + } + navFrame := send("tree-nav", "tree_nav", fmt.Sprintf(`{"tid":%q,"rev":%q,"sum":false}`, abandonedID, tree.Revision)) + if navFrame.Kind != "ack" || navFrame.Operation != "tree_nav" { + t.Fatalf("unexpected tree_nav frame: %#v", navFrame) + } + var navigation PiTreeNavigateResult + decodePayload(navFrame, &navigation) + if pointerString(navigation.Tree.LeafID) != abandonedID || navigation.Tree.Revision == tree.Revision { + t.Fatalf("tree_nav did not return the switched tree: %#v", navigation) + } + sourceBefore := mustGetSession(t, manager, created.ID) + + type createWireResult struct { + Session wireSess `json:"s"` + Tree PiTreeSnapshot `json:"tree"` + EditorText string `json:"editorText"` + } + forkFrame := send("tree-fork", "tree_fork", fmt.Sprintf(`{"tid":%q,"rev":%q}`, userID, navigation.Tree.Revision)) + if forkFrame.Kind != "ack" || forkFrame.Operation != "tree_fork" { + t.Fatalf("unexpected tree_fork frame: %#v", forkFrame) + } + var forkResult createWireResult + decodePayload(forkFrame, &forkResult) + if forkResult.Session.ID == "" || forkResult.Session.ID == created.ID || forkResult.EditorText != "pr6-tree-seed" { + t.Fatalf("unexpected tree_fork payload: %#v", forkResult) + } + + cloneFrame := send("tree-clone", "tree_clone", fmt.Sprintf(`{"rev":%q}`, navigation.Tree.Revision)) + if cloneFrame.Kind != "ack" || cloneFrame.Operation != "tree_clone" { + t.Fatalf("unexpected tree_clone frame: kind=%q op=%q code=%q message=%q", cloneFrame.Kind, cloneFrame.Operation, cloneFrame.Code, cloneFrame.Message) + } + var cloneResult createWireResult + decodePayload(cloneFrame, &cloneResult) + if cloneResult.Session.ID == "" || cloneResult.Session.ID == created.ID || cloneResult.Session.ID == forkResult.Session.ID || cloneResult.EditorText != "" { + t.Fatalf("unexpected tree_clone payload: %#v", cloneResult) + } + + sourceAfter := mustGetSession(t, manager, created.ID) + if pointerString(sourceAfter.NativeSessionID) != pointerString(sourceBefore.NativeSessionID) || + !samePiRuntimePath(pointerString(sourceAfter.ThreadPath), pointerString(sourceBefore.ThreadPath)) || + pointerString(sourceAfter.NativeLeafID) != pointerString(sourceBefore.NativeLeafID) || + pointerString(sourceAfter.SourceRevision) != pointerString(sourceBefore.SourceRevision) { + t.Fatalf("wire tree mutations changed source identity: before=%#v after=%#v", sourceBefore, sourceAfter) + } +} + +func TestManagerPiRPCTreeMutationFailureKeepsSourceAndCreatesNoTarget(t *testing.T) { + cleanup := initTestDB(t) + defer cleanup() + project := seedProject(t) + fakeDir := t.TempDir() + sourcePath := filepath.Join(fakeDir, "failed-mutation-source.jsonl") + logPath := filepath.Join(fakeDir, "failed-mutation-commands.jsonl") + t.Setenv("CODEKANBAN_FAKE_PI_RUNTIME", "1") + t.Setenv("CODEKANBAN_FAKE_PI_SESSION", sourcePath) + t.Setenv("CODEKANBAN_FAKE_PI_INVALID_MUTATION_STATE", "1") + t.Setenv("PI_CODING_AGENT_SESSION_DIR", fakeDir) + t.Setenv("CODEKANBAN_FAKE_PI_LOG", logPath) + manager, err := NewManager(Config{ + DataDir: t.TempDir(), PiPath: fmt.Sprintf("%q -test.run=^TestPiRuntimeFakeProcess$ --", os.Args[0]), + PiRuntimeIdleTTL: time.Minute, + }, zap.NewNop()) + if err != nil { + t.Fatal(err) + } + defer manager.StopProjectPiRuntimes(project.ID) + if _, err := manager.TrustProjectForPi(context.Background(), project.ID); err != nil { + t.Fatal(err) + } + created, err := manager.CreateSession(context.Background(), CreateParams{ + ProjectID: project.ID, Agent: AgentPi, Model: "openai/gpt-test", + }) + if err != nil { + t.Fatal(err) + } + if err := manager.SendMessage(context.Background(), created.ID, "pr6-tree-seed", nil); err != nil { + t.Fatal(err) + } + waitForSessionToSettle(t, manager, created.ID) + tree, err := manager.GetPiSessionTree(context.Background(), created.ID) + if err != nil { + t.Fatal(err) + } + sourceBefore := mustGetSession(t, manager, created.ID) + historyBefore, err := manager.History(context.Background(), created.ID, DefaultHistoryWindow, nil) + if err != nil { + t.Fatal(err) + } + + _, err = manager.ClonePiSessionTree(context.Background(), created.ID, PiTreeCloneInput{Revision: tree.Revision}) + if err == nil || !strings.Contains(err.Error(), "incomplete session identity") { + t.Fatalf("expected post-mutation identity failure, got %v", err) + } + sourceAfter := mustGetSession(t, manager, created.ID) + if pointerString(sourceAfter.NativeSessionID) != pointerString(sourceBefore.NativeSessionID) || + !samePiRuntimePath(pointerString(sourceAfter.ThreadPath), pointerString(sourceBefore.ThreadPath)) || + pointerString(sourceAfter.SourceRevision) != pointerString(sourceBefore.SourceRevision) { + t.Fatalf("failed mutation changed source identity: before=%#v after=%#v", sourceBefore, sourceAfter) + } + var sessionCount, turnCount, itemCount int64 + db := model.GetDB() + if err := db.Model(&tables.WebSessionTable{}).Where("project_id = ? AND agent = ?", project.ID, string(AgentPi)).Count(&sessionCount).Error; err != nil { + t.Fatal(err) + } + if err := db.Model(&tables.WebSessionTurnTable{}).Where("web_session_id <> ?", created.ID).Count(&turnCount).Error; err != nil { + t.Fatal(err) + } + if err := db.Model(&tables.WebSessionItemTable{}).Where("web_session_id <> ?", created.ID).Count(&itemCount).Error; err != nil { + t.Fatal(err) + } + if sessionCount != 1 || turnCount != 0 || itemCount != 0 { + t.Fatalf("failed mutation left target rows: sessions=%d turns=%d items=%d", sessionCount, turnCount, itemCount) + } + historyAfter, err := manager.History(context.Background(), created.ID, DefaultHistoryWindow, nil) + if err != nil { + t.Fatal(err) + } + if len(historyAfter.Items) != len(historyBefore.Items) || !historyContainsText(historyAfter, "active branch") { + t.Fatalf("failed mutation changed source history: before=%#v after=%#v", historyBefore.Items, historyAfter.Items) + } + manager.piRuntimeMu.Lock() + _, runtimePresent := manager.piRuntimes[created.ID] + manager.piRuntimeMu.Unlock() + if runtimePresent { + t.Fatal("failed mutation retained the switched Pi runtime") + } + + t.Setenv("CODEKANBAN_FAKE_PI_INVALID_MUTATION_STATE", "0") + if _, err := manager.GetPiSessionTree(context.Background(), created.ID); err != nil { + t.Fatalf("restore source after failed mutation: %v", err) + } + starts := 0 + for _, command := range readFakePiLog(t, logPath) { + if _, startup := command["startup"]; startup { + starts++ + } + } + if starts != 2 { + t.Fatalf("expected source runtime restore after failed mutation, starts=%d", starts) + } +} + +func TestManagerPiRPCProjectsPR5EventsAndNativePendingInputs(t *testing.T) { + cleanup := initTestDB(t) + defer cleanup() + project := seedProject(t) + fakeDir := t.TempDir() + sessionPath := filepath.Join(fakeDir, "pr5-session.jsonl") + logPath := filepath.Join(fakeDir, "pr5-commands.jsonl") + t.Setenv("CODEKANBAN_FAKE_PI_RUNTIME", "1") + t.Setenv("CODEKANBAN_FAKE_PI_SESSION", sessionPath) + t.Setenv("PI_CODING_AGENT_SESSION_DIR", fakeDir) + t.Setenv("CODEKANBAN_FAKE_PI_LOG", logPath) + manager, err := NewManager(Config{ + DataDir: t.TempDir(), PiPath: fmt.Sprintf("%q -test.run=^TestPiRuntimeFakeProcess$ --", os.Args[0]), + PiRuntimeIdleTTL: time.Minute, + }, zap.NewNop()) + if err != nil { + t.Fatal(err) + } + defer manager.StopProjectPiRuntimes(project.ID) + if _, err := manager.TrustProjectForPi(context.Background(), project.ID); err != nil { + t.Fatal(err) + } + if !manager.GetWebSessionRuntimeConfig().SupportsPiWebSession { + t.Fatal("fake Pi runtime should support Web Sessions") + } + created, err := manager.CreateSession(context.Background(), CreateParams{ProjectID: project.ID, Agent: AgentPi}) + if err != nil { + t.Fatal(err) + } + if err := manager.SendMessage(context.Background(), created.ID, "pr5-events", nil); err != nil { + t.Fatal(err) + } + approval := waitForPendingServerRequest(t, manager, created.ID, pendingServerRequestToolApproval) + if approval == nil || approval.PiRequestID != "confirm-1" { + t.Fatalf("unexpected Pi approval: %#v", approval) + } + snapshot, err := manager.loadSnapshotLocal(context.Background(), mustGetSession(t, manager, created.ID), DefaultHistoryWindow, false) + if err != nil { + t.Fatal(err) + } + if snapshot.PendingApproval == nil || !snapshot.PendingApproval.Actionable { + t.Fatalf("expected actionable Pi approval, got %#v", snapshot.PendingApproval) + } + if err := manager.sendMessageWithMode(context.Background(), created.ID, "redirect while active", nil, PendingInputModeRedirect, "pi-steer"); err != nil { + t.Fatal(err) + } + if err := manager.sendMessageWithMode(context.Background(), created.ID, "queue while active", nil, PendingInputModeQueue, "pi-follow-up"); err != nil { + t.Fatal(err) + } + time.Sleep(50 * time.Millisecond) + for _, command := range readFakePiLog(t, logPath) { + if command["type"] == "steer" || command["type"] == "follow_up" { + t.Fatalf("queued input crossed pending Pi dialog: %#v", command) + } + } + if err := manager.respondToApproval(created.ID, "approve"); err != nil { + t.Fatalf("respondToApproval: %v", err) + } + input := waitForPendingServerRequest(t, manager, created.ID, pendingServerRequestUserInput) + if input == nil || input.PiRequestID != "input-1" { + t.Fatalf("unexpected Pi input request: %#v", input) + } + if err := manager.respondToUserInput(created.ID, input.ItemID, map[string][]string{"value": {"typed value"}}); err != nil { + t.Fatalf("respondToUserInput: %v", err) + } + waitForPiNativeQueue(t, manager, created.ID, 2) + manager.piRuntimeMu.Lock() + managerRuntime := manager.piRuntimes[created.ID] + manager.piRuntimeMu.Unlock() + if managerRuntime == nil { + t.Fatal("expected active Pi runtime") + } + var state piRPCState + if err := managerRuntime.client.Request(context.Background(), "get_state", nil, &state); err != nil { + t.Fatalf("release fake Pi queued continuation: %v", err) + } + waitForSessionToSettle(t, manager, created.ID) + if queued := manager.pendingInputsDisplaySnapshot(created.ID); len(queued) != 0 { + t.Fatalf("native Pi queue was not cleared after settle: %#v", queued) + } + + events, err := manager.store.readEvents(created.ID) + if err != nil { + t.Fatal(err) + } + texts := userMessageTexts(events) + if !containsString(texts, "redirect while active") || !containsString(texts, "queue while active") { + t.Fatalf("queued Pi messages were not projected: %#v", texts) + } + var authoritative, stale bool + toolEnds := map[string]string{} + runDoneIndex := -1 + lastProjectionIndex := -1 + for index, event := range events { + switch event.Type { + case "txt_end": + text := stringValue(event.Payload["txt"]) + authoritative = authoritative || text == "authoritative reply" + stale = stale || text == "stale delta" + lastProjectionIndex = index + case "tool_end": + toolEnds[stringValue(event.Payload["tid"])] = stringValue(event.Payload["out"]) + lastProjectionIndex = index + case "note": + code := stringValue(event.Payload["code"]) + if strings.HasPrefix(code, "pi_auto_retry") || strings.HasPrefix(code, "pi_compaction") { + lastProjectionIndex = index + } + case "run_done": + runDoneIndex = index + } + } + if !authoritative || stale { + t.Fatalf("Pi message_end did not authoritatively calibrate text: authoritative=%v staleEnd=%v", authoritative, stale) + } + if toolEnds["parallel-a"] != "A done" || toolEnds["parallel-b"] != "B done" { + t.Fatalf("parallel Pi tools overwrote each other: %#v", toolEnds) + } + if runDoneIndex <= lastProjectionIndex { + record := mustGetSession(t, manager, created.ID) + var failures []map[string]any + for _, event := range events { + if event.Type == "run_fail" { + failures = append(failures, event.Payload) + } + } + t.Fatalf("run_done arrived before projection settled: runDone=%d lastProjection=%d status=%q lastError=%q failures=%#v", runDoneIndex, lastProjectionIndex, record.Status, pointerString(record.LastError), failures) + } + window, err := manager.History(context.Background(), created.ID, DefaultHistoryWindow, nil) + if err != nil { + t.Fatal(err) + } + if !historyContainsText(window, "authoritative reply") || historyContainsText(window, "stale delta") { + t.Fatalf("unexpected authoritative Pi history: %#v", window.Items) + } + if !historyContainsToolOutput(window, "parallel-a", "A done") || !historyContainsToolOutput(window, "parallel-b", "B done") { + t.Fatalf("missing independent Pi tool projections: %#v", window.Items) + } + nativeMessages := 0 + preservedLiveItems := 0 + finalReasoning := false + for _, item := range window.Items { + if (item.Kind == "user" || item.Kind == "assistant") && pointerString(item.SourceThreadID) == "fake-pi-session" && pointerString(item.SourceItemID) != "" { + nativeMessages++ + } + if item.Kind == "tool" || item.Kind == "system" { + preservedLiveItems++ + } + if item.Tool != nil && item.Tool.Kind == "reasoning" { + finalReasoning = finalReasoning || item.Tool.Output == "final reasoning" + if item.Tool.Output == "stale reasoning" { + t.Fatalf("Pi message_end did not authoritatively calibrate thinking: %#v", item) + } + } + } + if nativeMessages != 4 || preservedLiveItems < 5 || !finalReasoning { + t.Fatalf("Pi incremental sync lost native/live items: nativeMessages=%d preservedLiveItems=%d finalReasoning=%v items=%#v", nativeMessages, preservedLiveItems, finalReasoning, window.Items) + } + var turns []tables.WebSessionTurnTable + if err := model.GetDB().Where("web_session_id = ?", created.ID).Find(&turns).Error; err != nil { + t.Fatal(err) + } + if len(turns) != 3 { + t.Fatalf("expected three native Pi turns, got %d: %#v", len(turns), turns) + } + for _, turn := range turns { + if pointerString(turn.SourceThreadID) != "fake-pi-session" || pointerString(turn.SourceTurnID) == "" { + t.Fatalf("Pi turn missing native identity: %#v", turn) + } + } + var liveRows []tables.WebSessionItemTable + if err := model.GetDB().Where("web_session_id = ? AND item_kind IN ?", created.ID, []string{"tool", "system"}).Find(&liveRows).Error; err != nil { + t.Fatal(err) + } + for _, row := range liveRows { + if pointerString(row.WebTurnID) == "" || pointerString(row.SourceThreadID) != "fake-pi-session" || pointerString(row.SourceTurnID) == "" { + t.Fatalf("Pi live item missing native turn binding: %#v", row) + } + } + record := mustGetSession(t, manager, created.ID) + if pointerString(record.NativeLeafID) == "" || pointerString(record.SourceRevision) == "" || record.ItemCount != len(window.Items) { + t.Fatalf("Pi incremental sync identity/count mismatch: leaf=%q revision=%q itemCount=%d history=%d", pointerString(record.NativeLeafID), pointerString(record.SourceRevision), record.ItemCount, len(window.Items)) + } + commands := readFakePiLog(t, logPath) + sequence := make([]string, 0) + for _, command := range commands { + kind, _ := command["type"].(string) + if kind == "extension_ui_response" || kind == "steer" || kind == "follow_up" { + sequence = append(sequence, kind+":"+stringValue(command["id"])) + } + } + if len(sequence) != 4 || !strings.HasPrefix(sequence[0], "extension_ui_response:confirm-1") || + !strings.HasPrefix(sequence[1], "extension_ui_response:input-1") || + !strings.HasPrefix(sequence[2], "steer:") || !strings.HasPrefix(sequence[3], "follow_up:") { + t.Fatalf("unexpected Pi interaction sequence: %#v", sequence) + } +} + +func TestManagerPiRPCExtensionConfirmTimeoutCancelsApproval(t *testing.T) { + cleanup := initTestDB(t) + defer cleanup() + project := seedProject(t) + fakeDir := t.TempDir() + logPath := filepath.Join(fakeDir, "timeout-commands.jsonl") + t.Setenv("CODEKANBAN_FAKE_PI_RUNTIME", "1") + t.Setenv("CODEKANBAN_FAKE_PI_SESSION", filepath.Join(fakeDir, "timeout-session.jsonl")) + t.Setenv("PI_CODING_AGENT_SESSION_DIR", fakeDir) + t.Setenv("CODEKANBAN_FAKE_PI_LOG", logPath) + manager, err := NewManager(Config{ + DataDir: t.TempDir(), PiPath: fmt.Sprintf("%q -test.run=^TestPiRuntimeFakeProcess$ --", os.Args[0]), + PiRuntimeIdleTTL: time.Minute, + }, zap.NewNop()) + if err != nil { + t.Fatal(err) + } + defer manager.StopProjectPiRuntimes(project.ID) + if _, err := manager.TrustProjectForPi(context.Background(), project.ID); err != nil { + t.Fatal(err) + } + created, err := manager.CreateSession(context.Background(), CreateParams{ProjectID: project.ID, Agent: AgentPi}) + if err != nil { + t.Fatal(err) + } + if err := manager.SendMessage(context.Background(), created.ID, "timeout-confirm", nil); err != nil { + t.Fatal(err) + } + waitForActiveRun(t, manager, created.ID) + waitForSessionToSettle(t, manager, created.ID) + + events, err := manager.store.readEvents(created.ID) + if err != nil { + t.Fatal(err) + } + approvalCancelled := false + for _, event := range events { + if event.Type == "user_input_res" { + t.Fatalf("confirm timeout was projected as user input: %#v", event) + } + if event.Type == "approval_res" && stringValue(event.Payload["act"]) == "cancel" { + approvalCancelled = true + } + } + if !approvalCancelled || !historyHasEvent(events, "run_done") { + t.Fatalf("confirm timeout did not close approval and settle: %#v", events) + } + commands := readFakePiLog(t, logPath) + cancelResponses := 0 + for _, command := range commands { + if command["type"] == "extension_ui_response" && command["id"] == "timeout-confirm-1" && command["cancelled"] == true { + cancelResponses++ + } + } + if cancelResponses != 1 { + t.Fatalf("expected one cancelled Pi timeout response, got %d: %#v", cancelResponses, commands) + } +} + +func TestManagerPiRPCSettledAndAbortClosePendingDialogs(t *testing.T) { + for _, testCase := range []struct { + name string + prompt string + requestID string + terminal string + abort bool + }{ + {name: "settled", prompt: "settle-with-dialog", requestID: "settle-dialog-1", terminal: "run_done"}, + {name: "abort", prompt: "hold-dialog", requestID: "abort-dialog-1", terminal: "run_abort", abort: true}, + } { + t.Run(testCase.name, func(t *testing.T) { + cleanup := initTestDB(t) + defer cleanup() + project := seedProject(t) + fakeDir := t.TempDir() + logPath := filepath.Join(fakeDir, testCase.name+"-commands.jsonl") + t.Setenv("CODEKANBAN_FAKE_PI_RUNTIME", "1") + t.Setenv("CODEKANBAN_FAKE_PI_SESSION", filepath.Join(fakeDir, testCase.name+"-session.jsonl")) + t.Setenv("PI_CODING_AGENT_SESSION_DIR", fakeDir) + t.Setenv("CODEKANBAN_FAKE_PI_LOG", logPath) + manager, err := NewManager(Config{ + DataDir: t.TempDir(), PiPath: fmt.Sprintf("%q -test.run=^TestPiRuntimeFakeProcess$ --", os.Args[0]), + PiRuntimeIdleTTL: time.Minute, + }, zap.NewNop()) + if err != nil { + t.Fatal(err) + } + defer manager.StopProjectPiRuntimes(project.ID) + if _, err := manager.TrustProjectForPi(context.Background(), project.ID); err != nil { + t.Fatal(err) + } + created, err := manager.CreateSession(context.Background(), CreateParams{ProjectID: project.ID, Agent: AgentPi}) + if err != nil { + t.Fatal(err) + } + if err := manager.SendMessage(context.Background(), created.ID, testCase.prompt, nil); err != nil { + t.Fatal(err) + } + if testCase.abort { + pending := waitForPendingServerRequest(t, manager, created.ID, pendingServerRequestToolApproval) + if pending.PiRequestID != testCase.requestID { + t.Fatalf("unexpected Pi dialog: %#v", pending) + } + if err := manager.AbortSession(created.ID); err != nil { + t.Fatal(err) + } + } + waitForSessionToSettle(t, manager, created.ID) + + events, err := manager.store.readEvents(created.ID) + if err != nil { + t.Fatal(err) + } + completionIndex, terminalIndex := -1, -1 + for index, event := range events { + if event.Type == "approval_res" && stringValue(event.Payload["act"]) == "cancel" { + completionIndex = index + } + if event.Type == testCase.terminal { + terminalIndex = index + } + } + if completionIndex < 0 || terminalIndex <= completionIndex { + t.Fatalf("Pi dialog was not closed before %s: %#v", testCase.terminal, events) + } + if snapshot, err := manager.Snapshot(context.Background(), created.ID, DefaultHistoryWindow); err != nil { + t.Fatal(err) + } else if snapshot.PendingApproval != nil || snapshot.PendingUserInput != nil { + t.Fatalf("settled Pi dialog remained actionable: %#v", snapshot) + } + for _, command := range readFakePiLog(t, logPath) { + if command["type"] == "extension_ui_response" && command["id"] == testCase.requestID { + t.Fatalf("terminal cleanup wrote a stale Pi extension response: %#v", command) + } + } + }) + } +} + +func TestManagerPiRPCManualCompactionUsesNativeRuntime(t *testing.T) { + cleanup := initTestDB(t) + defer cleanup() + project := seedProject(t) + fakeDir := t.TempDir() + sessionPath := filepath.Join(fakeDir, "compact-session.jsonl") + logPath := filepath.Join(fakeDir, "compact-commands.jsonl") + t.Setenv("CODEKANBAN_FAKE_PI_RUNTIME", "1") + t.Setenv("CODEKANBAN_FAKE_PI_SESSION", sessionPath) + t.Setenv("PI_CODING_AGENT_SESSION_DIR", fakeDir) + t.Setenv("CODEKANBAN_FAKE_PI_LOG", logPath) + manager, err := NewManager(Config{ + DataDir: t.TempDir(), PiPath: fmt.Sprintf("%q -test.run=^TestPiRuntimeFakeProcess$ --", os.Args[0]), + PiRuntimeIdleTTL: time.Minute, + }, zap.NewNop()) + if err != nil { + t.Fatal(err) + } + defer manager.StopProjectPiRuntimes(project.ID) + if _, err := manager.TrustProjectForPi(context.Background(), project.ID); err != nil { + t.Fatal(err) + } + created, err := manager.CreateSession(context.Background(), CreateParams{ProjectID: project.ID, Agent: AgentPi}) + if err != nil { + t.Fatal(err) + } + if err := manager.SendMessage(context.Background(), created.ID, "before compact", nil); err != nil { + t.Fatal(err) + } + waitForSessionToSettle(t, manager, created.ID) + before, err := manager.History(context.Background(), created.ID, DefaultHistoryWindow, nil) + if err != nil { + t.Fatal(err) + } + beforeUserCount := historyKindCount(before, "user") + beforeLeaf := pointerString(mustGetSession(t, manager, created.ID).NativeLeafID) + + if err := manager.CompactSession(context.Background(), created.ID); err != nil { + t.Fatal(err) + } + waitForSessionToSettle(t, manager, created.ID) + after, err := manager.History(context.Background(), created.ID, DefaultHistoryWindow, nil) + if err != nil { + t.Fatal(err) + } + if historyKindCount(after, "user") != beforeUserCount { + t.Fatalf("manual Pi compaction added a user message: before=%d after=%d", beforeUserCount, historyKindCount(after, "user")) + } + compactions := 0 + for _, item := range after.Items { + if item.Tool != nil && item.Tool.Kind == "context_compaction" { + compactions++ + if !item.Done || item.Tool.Output != "manual compact summary" { + t.Fatalf("unexpected manual compaction item: %#v", item) + } + } + } + if compactions != 1 { + t.Fatalf("expected one manual compaction item, got %d: %#v", compactions, after.Items) + } + record := mustGetSession(t, manager, created.ID) + if pointerString(record.NativeLeafID) == "" || pointerString(record.NativeLeafID) == beforeLeaf || pointerString(record.SourceRevision) == "" { + t.Fatalf("manual compaction did not refresh native identity: before=%q after=%q revision=%q", beforeLeaf, pointerString(record.NativeLeafID), pointerString(record.SourceRevision)) + } + if record.LastContextCompactionAt == nil || record.ContextBaselineInputTokens != 10 || record.ContextBaselineCachedInputTokens != 2 || record.ContextBaselineOutputTokens != 5 { + t.Fatalf("manual compaction did not reset the Pi context baseline: %#v", record) + } + if record.LatestTokenCountUpdatedAt != nil || record.LatestTurnUsageUpdatedAt != nil { + t.Fatalf("manual compaction retained a higher-priority context estimate: %#v", record) + } + snapshot, err := manager.Snapshot(context.Background(), created.ID, DefaultHistoryWindow) + if err != nil { + t.Fatal(err) + } + if snapshot.Session.ContextWindowTokens == nil || *snapshot.Session.ContextWindowTokens != 32000 || snapshot.Session.ContextWindowSource != ContextWindowSourceSessionUsage { + t.Fatalf("manual compaction lost the Pi session context window: %#v", snapshot.Session) + } + if snapshot.Session.ContextEstimateMode != ContextEstimateModeSinceCompaction || snapshot.Session.ContextEstimate.UsedTokens != 0 { + t.Fatalf("manual compaction did not expose the reset context baseline: %#v", snapshot.Session) + } + starts, compacts := 0, 0 + for _, command := range readFakePiLog(t, logPath) { + if _, ok := command["startup"]; ok { + starts++ + } + if command["type"] == "compact" { + compacts++ + } + } + if starts != 1 || compacts != 1 { + t.Fatalf("manual compaction did not reuse the Pi runtime: starts=%d compacts=%d", starts, compacts) + } +} + +func TestManagerPiRPCAbortKeepsSessionUsable(t *testing.T) { + cleanup := initTestDB(t) + defer cleanup() + project := seedProject(t) + fakeDir := t.TempDir() + logPath := filepath.Join(fakeDir, "abort-commands.jsonl") + t.Setenv("CODEKANBAN_FAKE_PI_RUNTIME", "1") + t.Setenv("CODEKANBAN_FAKE_PI_PROMPT_ACK_DELAY", "100ms") + t.Setenv("CODEKANBAN_FAKE_PI_SESSION", filepath.Join(fakeDir, "abort-session.jsonl")) + t.Setenv("PI_CODING_AGENT_SESSION_DIR", fakeDir) + t.Setenv("CODEKANBAN_FAKE_PI_LOG", logPath) + manager, err := NewManager(Config{ + DataDir: t.TempDir(), PiPath: fmt.Sprintf("%q -test.run=^TestPiRuntimeFakeProcess$ --", os.Args[0]), + PiRuntimeIdleTTL: time.Minute, + }, zap.NewNop()) + if err != nil { + t.Fatal(err) + } + defer manager.StopProjectPiRuntimes(project.ID) + if _, err := manager.TrustProjectForPi(context.Background(), project.ID); err != nil { + t.Fatal(err) + } + config := manager.GetWebSessionRuntimeConfig() + if !config.SupportsPiWebSession { + t.Fatalf("fake Pi runtime unavailable: hasPi=%v version=%v compatible=%v diagnostics=%q", config.HasPi, config.PiVersion, config.PiRPCCompatible, config.PiDiagnostics) + } + created, err := manager.CreateSession(context.Background(), CreateParams{ProjectID: project.ID, Agent: AgentPi}) + if err != nil { + t.Fatal(err) + } + if err := manager.SendMessage(context.Background(), created.ID, "before hold", nil); err != nil { + t.Fatal(err) + } + waitForSessionToSettle(t, manager, created.ID) + if err := manager.SendMessage(context.Background(), created.ID, "hold", nil); err != nil { + t.Fatal(err) + } + waitForActiveRun(t, manager, created.ID) + tree, err := manager.GetPiSessionTree(context.Background(), created.ID) + if err != nil { + t.Fatalf("read Pi tree during active run: %v", err) + } + if tree.SessionID != "fake-pi-session" || tree.Revision == "" { + t.Fatalf("unexpected active-run Pi tree: %#v", tree) + } + starts := 0 + for _, command := range readFakePiLog(t, logPath) { + if _, startup := command["startup"]; startup { + starts++ + } + } + if starts != 1 { + t.Fatalf("active-run tree read started another Pi runtime: starts=%d", starts) + } + if err := manager.AbortSession(created.ID); err != nil { + t.Fatalf("AbortSession returned error: %v", err) + } + waitForSessionToSettle(t, manager, created.ID) + if err := manager.SendMessage(context.Background(), created.ID, "after abort", nil); err != nil { + t.Fatalf("send after abort returned error: %v", err) + } + waitForSessionToSettle(t, manager, created.ID) +} + +func waitForActiveRun(t *testing.T, manager *Manager, sessionID string) { + t.Helper() + deadline := time.Now().Add(5 * time.Second) + for time.Now().Before(deadline) { + if manager.hasActiveRun(sessionID) { + return + } + time.Sleep(10 * time.Millisecond) + } + t.Fatalf("session %s did not start a run", sessionID) +} + +func TestPiPromptImagesRevalidatesAndEncodesAttachments(t *testing.T) { + store, err := newStore(t.TempDir()) + if err != nil { + t.Fatal(err) + } + manager := &Manager{store: store, cfg: Config{AttachmentSizeLimit: 1024}} + path := filepath.Join(store.attachmentsDir, "pixel.png") + if err := os.WriteFile(path, []byte("\x89PNG\r\n\x1a\n\x00\x00\x00\x00"), 0o600); err != nil { + t.Fatal(err) + } + images, err := manager.piPromptImages([]Attachment{{Name: "pixel.png", Mime: "image/png", Path: path}}) + if err != nil { + t.Fatalf("piPromptImages returned error: %v", err) + } + if len(images) != 1 || images[0].Data != "iVBORw0KGgoAAAAA" || images[0].MimeType != "image/png" { + t.Fatalf("unexpected images: %#v", images) + } + if _, err := manager.piPromptImages([]Attachment{{Name: "notes.txt", Mime: "text/plain", Path: path}}); err == nil { + t.Fatal("expected non-image attachment rejection") + } + outside := filepath.Join(t.TempDir(), "outside.png") + if err := os.WriteFile(outside, []byte("\x89PNG\r\n\x1a\n"), 0o600); err != nil { + t.Fatal(err) + } + if _, err := manager.piPromptImages([]Attachment{{Name: "outside.png", Mime: "image/png", Path: outside}}); err == nil { + t.Fatal("expected attachment-root rejection") + } +} + +func argsAfterDoubleDash(args []string) []string { + for index, value := range args { + if value == "--" { + return args[index+1:] + } + } + return args +} + +func containsString(values []string, target string) bool { + for _, value := range values { + if value == target { + return true + } + } + return false +} + +func fakePiFlagValue(args []string, flag string) string { + for index := 0; index+1 < len(args); index++ { + if args[index] == flag { + return args[index+1] + } + } + return "" +} + +func writeFakePiResponse(id any, command string, data any) { + writeFakePiEvent(map[string]any{"type": "response", "id": id, "command": command, "success": true, "data": data}) +} + +func writeFakePiError(id any, command, message string) { + writeFakePiEvent(map[string]any{"type": "response", "id": id, "command": command, "success": false, "error": message}) +} + +func writeFakePiEvent(value map[string]any) { + encoded, _ := json.Marshal(value) + _, _ = os.Stdout.Write(append(encoded, '\n')) +} + +func appendFakePiLog(value map[string]any) { + path := os.Getenv("CODEKANBAN_FAKE_PI_LOG") + if path == "" { + return + } + encoded, _ := json.Marshal(value) + file, err := os.OpenFile(path, os.O_CREATE|os.O_APPEND|os.O_WRONLY, 0o600) + if err == nil { + _, _ = file.Write(append(encoded, '\n')) + _ = file.Close() + } +} + +func readFakePiLog(t *testing.T, path string) []map[string]any { + t.Helper() + data, err := os.ReadFile(path) + if err != nil { + t.Fatal(err) + } + var result []map[string]any + for _, line := range strings.Split(strings.TrimSpace(string(data)), "\n") { + var value map[string]any + if err := json.Unmarshal([]byte(line), &value); err != nil { + t.Fatalf("decode fake Pi log: %v", err) + } + result = append(result, value) + } + return result +} + +func mustGetSession(t *testing.T, manager *Manager, sessionID string) tables.WebSessionTable { + t.Helper() + record, err := manager.GetSession(context.Background(), sessionID) + if err != nil { + t.Fatal(err) + } + return record +} + +func waitForPiNativeQueue(t *testing.T, manager *Manager, sessionID string, count int) { + t.Helper() + deadline := time.Now().Add(3 * time.Second) + for time.Now().Before(deadline) { + items := manager.pendingInputsDisplaySnapshot(sessionID) + if len(items) == count { + allNative := true + for _, item := range items { + if !item.NativeQueued { + allNative = false + break + } + } + if allNative && len(manager.pendingInputsSnapshot(sessionID)) == 0 { + return + } + } + time.Sleep(10 * time.Millisecond) + } + t.Fatalf("native Pi queue did not reach %d items: %#v", count, manager.pendingInputsDisplaySnapshot(sessionID)) +} + +func historyKindCount(window HistoryWindow, kind string) int { + count := 0 + for _, item := range window.Items { + if item.Kind == kind { + count++ + } + } + return count +} + +func historyContainsToolOutput(window HistoryWindow, toolID, output string) bool { + for _, item := range window.Items { + if item.Tool != nil && item.Tool.ID == toolID && strings.Contains(item.Tool.Output, output) { + return true + } + } + return false +} + +func historyContainsText(window HistoryWindow, expected string) bool { + for _, item := range window.Items { + if strings.Contains(item.Text, expected) { + return true + } + } + return false +} diff --git a/service/websession/pi_sync.go b/service/websession/pi_sync.go new file mode 100644 index 00000000..17623463 --- /dev/null +++ b/service/websession/pi_sync.go @@ -0,0 +1,532 @@ +package websession + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "strings" + "time" + + "code-kanban/model" + "code-kanban/model/tables" + "code-kanban/utils" + + "gorm.io/gorm" +) + +type piHistoryEntry struct { + Type string `json:"type"` + ID string `json:"id"` + ParentID *string `json:"parentId"` + Timestamp string `json:"timestamp"` + Message piHistoryMessage `json:"message"` + CustomType string `json:"customType"` + Data json.RawMessage `json:"data"` + Content json.RawMessage `json:"content"` + Summary string `json:"summary"` + Provider string `json:"provider"` + ModelID string `json:"modelId"` + ThinkingLevel string `json:"thinkingLevel"` + Name string `json:"name"` +} + +type piHistoryMessage struct { + Role string `json:"role"` + Content json.RawMessage `json:"content"` + Timestamp int64 `json:"timestamp"` + StopReason string `json:"stopReason"` + ErrorMessage string `json:"errorMessage"` +} + +type piHistoryEntriesResponse struct { + Entries []piHistoryEntry `json:"entries"` + LeafID *string `json:"leafId"` +} + +type piHistoryTreeResponse struct { + Tree json.RawMessage `json:"tree"` + LeafID *string `json:"leafId"` +} + +func (m *Manager) syncImportedPiSession( + ctx context.Context, + session tables.WebSessionTable, +) (SessionSnapshot, error) { + if m.hasActiveRun(session.ID) { + return SessionSnapshot{}, errors.New("cannot sync an active Pi web session") + } + runtime, err := m.getOrStartPiRuntime(ctx, session) + if err != nil { + return SessionSnapshot{}, err + } + defer runtime.scheduleIdle() + + var tree piHistoryTreeResponse + if err := runtime.client.Request(ctx, "get_tree", nil, &tree); err != nil { + return SessionSnapshot{}, fmt.Errorf("read Pi session tree: %w", err) + } + if len(tree.Tree) == 0 || !json.Valid(tree.Tree) { + return SessionSnapshot{}, errors.New("Pi returned an invalid session tree") + } + var entries piHistoryEntriesResponse + if err := runtime.client.Request(ctx, "get_entries", nil, &entries); err != nil { + return SessionSnapshot{}, fmt.Errorf("read Pi session entries: %w", err) + } + if pointerString(tree.LeafID) != pointerString(entries.LeafID) { + return SessionSnapshot{}, errors.New("Pi session tree and entry leaf do not match") + } + + refreshed, err := m.GetSession(ctx, session.ID) + if err != nil { + return SessionSnapshot{}, err + } + if err := m.projectPiHistoryEntries(ctx, refreshed, entries); err != nil { + return SessionSnapshot{}, err + } + if err := m.syncPiRuntimeSnapshot(ctx, runtime, refreshed); err != nil { + return SessionSnapshot{}, err + } + refreshed, err = m.GetSession(ctx, session.ID) + if err != nil { + return SessionSnapshot{}, err + } + return m.loadSnapshotLocal(ctx, refreshed, DefaultHistoryWindow, false) +} + +type piHistoryProjection struct { + turns []tables.WebSessionTurnTable + items []tables.WebSessionItemTable + updates map[string]any +} + +func (m *Manager) projectPiHistoryEntries( + ctx context.Context, + session tables.WebSessionTable, + response piHistoryEntriesResponse, +) error { + projection, err := buildPiHistoryProjection(session, response) + if err != nil { + return err + } + if err := m.store.deleteSessionFiles(session.ID); err != nil { + return err + } + return m.replaceSessionHistoryCache(ctx, session, projection.turns, projection.items, projection.updates) +} + +func buildPiHistoryProjection( + session tables.WebSessionTable, + response piHistoryEntriesResponse, +) (piHistoryProjection, error) { + active, err := activePiHistoryEntries(response.Entries, pointerString(response.LeafID)) + if err != nil { + return piHistoryProjection{}, err + } + turns := make([]tables.WebSessionTurnTable, 0) + items := make([]tables.WebSessionItemTable, 0, len(active)) + nativeID := pointerString(session.NativeSessionID) + var currentTurn *tables.WebSessionTurnTable + var order int64 + var lastMessageAt *time.Time + + for _, entry := range active { + if entry.Type != "message" { + continue + } + role := strings.ToLower(strings.TrimSpace(entry.Message.Role)) + if role != "user" && role != "assistant" { + continue + } + text := piHistoryContentText(entry.Message.Content) + if text == "" && strings.TrimSpace(entry.Message.ErrorMessage) == "" { + continue + } + observedAt := piHistoryEntryTime(entry) + lastMessageAt = &observedAt + + if role == "user" || currentTurn == nil { + turn := tables.WebSessionTurnTable{} + turn.Init() + turn.WebSessionID = session.ID + turn.SourceThreadID = nilIfEmptyHistory(nativeID) + turn.SourceTurnID = nilIfEmptyHistory(entry.ID) + turn.OrderIndex = int64(len(turns) + 1) + turn.Status = "completed" + turn.SourceCreated = true + turns = append(turns, turn) + currentTurn = &turns[len(turns)-1] + } + + order++ + item := HistoryItem{ + ID: utils.NewID(), + SourceThreadID: nilIfEmptyHistory(nativeID), + SourceTurnID: currentTurn.SourceTurnID, + SourceItemID: nilIfEmptyHistory(entry.ID), + OrderIndex: order, + Kind: role, + ItemType: map[string]string{"user": "user_message", "assistant": "agent_message"}[role], + Text: text, + Timestamp: &observedAt, + ObservedAt: &observedAt, + Done: true, + } + if strings.EqualFold(strings.TrimSpace(entry.Message.StopReason), "error") || strings.TrimSpace(entry.Message.ErrorMessage) != "" { + item.Level = "error" + if item.Text == "" { + item.Text = "Pi assistant run failed" + } + } + row := tables.WebSessionItemTable{} + row.Init() + row.WebSessionID = session.ID + row.WebTurnID = ¤tTurn.ID + applyHistoryItemToRow(&row, session.ID, item) + items = append(items, row) + } + + now := time.Now() + updates := map[string]any{ + "source_kind": string(SessionBackendPiRPC), + "native_leaf_id": nilIfEmpty(pointerString(response.LeafID)), + "source_revision": nilIfEmpty(piSourceRevision(pointerString(session.ThreadPath), pointerString(response.LeafID))), + "last_synced_at": now, + "sync_state": string(SyncStateFresh), + "sync_error": nil, + "turn_count": len(turns), + "item_count": len(items), + "last_message_at": lastMessageAt, + "source_updated_at": lastMessageAt, + "updated_at": now, + } + return piHistoryProjection{turns: turns, items: items, updates: updates}, nil +} + +type piLiveMessage struct { + entry piHistoryEntry + role string + text string + observedAt time.Time + turnSourceID string +} + +func piLiveMessages(response piHistoryEntriesResponse) ([]piLiveMessage, error) { + active, err := activePiHistoryEntries(response.Entries, pointerString(response.LeafID)) + if err != nil { + return nil, err + } + messages := make([]piLiveMessage, 0, len(active)) + turnSourceID := "" + for _, entry := range active { + if entry.Type != "message" { + continue + } + role := strings.ToLower(strings.TrimSpace(entry.Message.Role)) + if role != "user" && role != "assistant" { + continue + } + text := piHistoryContentText(entry.Message.Content) + if text == "" && strings.TrimSpace(entry.Message.ErrorMessage) == "" { + continue + } + if role == "user" || turnSourceID == "" { + turnSourceID = strings.TrimSpace(entry.ID) + } + messages = append(messages, piLiveMessage{ + entry: entry, role: role, text: text, + observedAt: piHistoryEntryTime(entry), turnSourceID: turnSourceID, + }) + } + return messages, nil +} + +func (m *Manager) reconcileLivePiHistory( + ctx context.Context, + session tables.WebSessionTable, + nativeSessionID string, + response piHistoryEntriesResponse, + updates map[string]any, +) error { + nativeSessionID = strings.TrimSpace(nativeSessionID) + if nativeSessionID == "" { + return errors.New("Pi runtime has no native session id") + } + messages, err := piLiveMessages(response) + if err != nil { + return err + } + db := model.GetDB() + if db == nil { + return model.ErrDBNotInitialized + } + return db.WithContext(ctx).Transaction(func(tx *gorm.DB) error { + var current tables.WebSessionTable + if err := tx.Select("id").First(¤t, "id = ?", session.ID).Error; err != nil { + return err + } + + var itemRows []tables.WebSessionItemTable + if err := tx.Where("web_session_id = ?", session.ID).Order("order_index ASC").Find(&itemRows).Error; err != nil { + return err + } + var turnRows []tables.WebSessionTurnTable + if err := tx.Where("web_session_id = ?", session.ID).Order("order_index ASC").Find(&turnRows).Error; err != nil { + return err + } + + turnBySource := make(map[string]*tables.WebSessionTurnTable, len(turnRows)) + for index := range turnRows { + row := &turnRows[index] + if pointerString(row.SourceThreadID) == nativeSessionID { + turnBySource[pointerString(row.SourceTurnID)] = row + } + } + turnIDs := make(map[string]string) + activeTurnSources := make([]string, 0) + seenTurn := make(map[string]struct{}) + for _, message := range messages { + turnSourceID := strings.TrimSpace(message.turnSourceID) + if _, exists := seenTurn[turnSourceID]; exists { + continue + } + seenTurn[turnSourceID] = struct{}{} + activeTurnSources = append(activeTurnSources, turnSourceID) + row := turnBySource[turnSourceID] + if row == nil { + created := &tables.WebSessionTurnTable{} + created.Init() + created.WebSessionID = session.ID + row = created + } + row.SourceThreadID = nilIfEmptyHistory(nativeSessionID) + row.SourceTurnID = nilIfEmptyHistory(turnSourceID) + row.OrderIndex = int64(len(activeTurnSources)) + row.Status = "completed" + row.SourceCreated = true + if row.CreatedAt.IsZero() { + if err := tx.Create(row).Error; err != nil { + return err + } + } else if err := tx.Save(row).Error; err != nil { + return err + } + turnIDs[turnSourceID] = row.ID + } + + exactRows := make(map[string]*tables.WebSessionItemTable) + for index := range itemRows { + row := &itemRows[index] + if pointerString(row.SourceThreadID) == nativeSessionID { + exactRows[pointerString(row.SourceItemID)] = row + } + } + usedRows := make(map[string]struct{}, len(messages)) + maxOrder := int64(0) + for index := range itemRows { + if itemRows[index].OrderIndex > maxOrder { + maxOrder = itemRows[index].OrderIndex + } + } + activeItemIDs := make([]string, 0, len(messages)) + turnStartOrders := make(map[string]int64, len(activeTurnSources)) + var lastMessageAt *time.Time + for _, message := range messages { + entryID := strings.TrimSpace(message.entry.ID) + activeItemIDs = append(activeItemIDs, entryID) + row := exactRows[entryID] + if row == nil { + row = findPiLiveMessageCandidate(itemRows, usedRows, message.role, message.text) + } + isNew := row == nil + if isNew { + row = &tables.WebSessionItemTable{} + row.Init() + maxOrder++ + row.OrderIndex = maxOrder + } + usedRows[row.ID] = struct{}{} + item := mapHistoryItemRowWithSession(*row, session.ID) + item.SourceThreadID = nilIfEmptyHistory(nativeSessionID) + item.SourceTurnID = nilIfEmptyHistory(message.turnSourceID) + item.SourceItemID = nilIfEmptyHistory(entryID) + item.Kind = message.role + item.ItemType = map[string]string{"user": "user_message", "assistant": "agent_message"}[message.role] + item.Text = message.text + item.Timestamp = &message.observedAt + item.ObservedAt = &message.observedAt + item.Done = true + item.Level = "" + if strings.EqualFold(strings.TrimSpace(message.entry.Message.StopReason), "error") || strings.TrimSpace(message.entry.Message.ErrorMessage) != "" { + item.Level = "error" + if item.Text == "" { + item.Text = "Pi assistant run failed" + } + } + applyHistoryItemToRow(row, session.ID, item) + row.WebTurnID = nilIfEmptyHistory(turnIDs[message.turnSourceID]) + row.Role = message.role + row.Status = "completed" + if current, exists := turnStartOrders[message.turnSourceID]; !exists || row.OrderIndex < current { + turnStartOrders[message.turnSourceID] = row.OrderIndex + } + if isNew { + if err := tx.Create(row).Error; err != nil { + return err + } + } else if err := tx.Save(row).Error; err != nil { + return err + } + value := message.observedAt + lastMessageAt = &value + } + + for index, turnSourceID := range activeTurnSources { + startOrder, exists := turnStartOrders[turnSourceID] + if !exists { + continue + } + liveItems := tx.Model(&tables.WebSessionItemTable{}). + Where("web_session_id = ? AND order_index >= ? AND (source_thread_id IS NULL OR source_thread_id = '')", session.ID, startOrder) + if index+1 < len(activeTurnSources) { + if nextOrder, ok := turnStartOrders[activeTurnSources[index+1]]; ok { + liveItems = liveItems.Where("order_index < ?", nextOrder) + } + } + if err := liveItems.Updates(map[string]any{ + "web_turn_id": turnIDs[turnSourceID], + "source_thread_id": nativeSessionID, + "source_turn_id": turnSourceID, + }).Error; err != nil { + return err + } + } + + staleItems := tx.Unscoped().Where( + "web_session_id = ? AND source_thread_id = ? AND item_kind IN ?", + session.ID, nativeSessionID, []string{"user", "assistant"}, + ) + if len(activeItemIDs) > 0 { + staleItems = staleItems.Where("source_item_id NOT IN ?", activeItemIDs) + } + if err := staleItems.Delete(&tables.WebSessionItemTable{}).Error; err != nil { + return err + } + staleTurns := tx.Unscoped().Where("web_session_id = ? AND source_thread_id = ?", session.ID, nativeSessionID) + if len(activeTurnSources) > 0 { + staleTurns = staleTurns.Where("source_turn_id NOT IN ?", activeTurnSources) + } + if err := staleTurns.Delete(&tables.WebSessionTurnTable{}).Error; err != nil { + return err + } + + var itemCount int64 + if err := tx.Model(&tables.WebSessionItemTable{}).Where("web_session_id = ?", session.ID).Count(&itemCount).Error; err != nil { + return err + } + var turnCount int64 + if err := tx.Model(&tables.WebSessionTurnTable{}).Where("web_session_id = ?", session.ID).Count(&turnCount).Error; err != nil { + return err + } + nextUpdates := cloneMap(updates) + nextUpdates["turn_count"] = turnCount + nextUpdates["item_count"] = itemCount + nextUpdates["last_message_at"] = lastMessageAt + nextUpdates["source_updated_at"] = lastMessageAt + return tx.Model(&tables.WebSessionTable{}). + Where("id = ?", session.ID). + Updates(withSnapshotRevisionIncrement(nextUpdates)).Error + }) +} + +func findPiLiveMessageCandidate( + rows []tables.WebSessionItemTable, + used map[string]struct{}, + role string, + text string, +) *tables.WebSessionItemTable { + var fallback *tables.WebSessionItemTable + for index := range rows { + row := &rows[index] + if _, claimed := used[row.ID]; claimed || pointerString(row.SourceThreadID) != "" || row.ItemKind != role { + continue + } + if fallback == nil { + fallback = row + } + if row.Text == text { + return row + } + } + return fallback +} + +func activePiHistoryEntries(entries []piHistoryEntry, leafID string) ([]piHistoryEntry, error) { + if strings.TrimSpace(leafID) == "" { + if len(entries) == 0 { + return nil, nil + } + return nil, errors.New("Pi session entries have no active leaf") + } + byID := make(map[string]piHistoryEntry, len(entries)) + for _, entry := range entries { + if id := strings.TrimSpace(entry.ID); id != "" { + byID[id] = entry + } + } + path := make([]piHistoryEntry, 0, len(entries)) + seen := make(map[string]struct{}, len(entries)) + for currentID := strings.TrimSpace(leafID); currentID != ""; { + if _, duplicate := seen[currentID]; duplicate { + return nil, errors.New("Pi session active branch contains a cycle") + } + entry, ok := byID[currentID] + if !ok { + return nil, errors.New("Pi session active branch is incomplete") + } + seen[currentID] = struct{}{} + path = append(path, entry) + if entry.ParentID == nil { + break + } + currentID = strings.TrimSpace(*entry.ParentID) + } + for left, right := 0, len(path)-1; left < right; left, right = left+1, right-1 { + path[left], path[right] = path[right], path[left] + } + return path, nil +} + +func piHistoryContentText(raw json.RawMessage) string { + if len(raw) == 0 { + return "" + } + var text string + if json.Unmarshal(raw, &text) == nil { + return strings.TrimSpace(text) + } + var blocks []struct { + Type string `json:"type"` + Text string `json:"text"` + } + if json.Unmarshal(raw, &blocks) != nil { + return "" + } + parts := make([]string, 0, len(blocks)) + for _, block := range blocks { + if strings.EqualFold(block.Type, "text") && strings.TrimSpace(block.Text) != "" { + parts = append(parts, strings.TrimSpace(block.Text)) + } + } + return strings.Join(parts, "\n") +} + +func piHistoryEntryTime(entry piHistoryEntry) time.Time { + if timestamp, err := time.Parse(time.RFC3339Nano, strings.TrimSpace(entry.Timestamp)); err == nil { + return timestamp + } + if entry.Message.Timestamp > 0 { + return time.UnixMilli(entry.Message.Timestamp) + } + return time.Now() +} diff --git a/service/websession/pi_tree.go b/service/websession/pi_tree.go new file mode 100644 index 00000000..65814f45 --- /dev/null +++ b/service/websession/pi_tree.go @@ -0,0 +1,554 @@ +package websession + +import ( + "context" + "encoding/base64" + "encoding/json" + "errors" + "fmt" + "strings" + "time" + + "code-kanban/model/tables" + "code-kanban/utils" + + "go.uber.org/zap" +) + +const ( + piTreeMaxNodes = 10000 + piTreePreviewRunes = 240 +) + +var ErrPiTreeRevisionConflict = errors.New("Pi session tree changed; refresh and try again") + +type PiTreePublicError struct { + Code string + Message string +} + +func ClassifyPiTreeError(err error) PiTreePublicError { + if errors.Is(err, ErrPiTreeRevisionConflict) { + return PiTreePublicError{Code: "conflict", Message: ErrPiTreeRevisionConflict.Error()} + } + message := strings.ToLower(err.Error()) + switch { + case strings.Contains(message, "active pi web session"), + strings.Contains(message, "while messages are pending"), + strings.Contains(message, "session is archived"): + return PiTreePublicError{Code: "invalid_state", Message: "Pi session tree cannot change in the current session state"} + case strings.Contains(message, "project agent trust is required"), + strings.Contains(message, "path is not managed by the project"): + return PiTreePublicError{Code: "forbidden", Message: "Pi project trust is required"} + case strings.Contains(message, "is required"), + strings.Contains(message, "target does not exist"), + strings.Contains(message, "target is not a user message"), + strings.Contains(message, "only supported for pi"), + strings.Contains(message, "requires an existing native session"): + return PiTreePublicError{Code: "bad_req", Message: "Invalid Pi session tree request"} + default: + return PiTreePublicError{Code: "internal", Message: "Pi session tree operation failed"} + } +} + +type PiTreeNode struct { + ID string `json:"id"` + ParentID *string `json:"parentId"` + Type string `json:"type"` + Role string `json:"role,omitempty"` + Preview string `json:"preview,omitempty"` + Timestamp string `json:"timestamp,omitempty"` + Label string `json:"label,omitempty"` + Active bool `json:"active"` + Children []string `json:"children"` +} + +type PiTreeSnapshot struct { + SessionID string `json:"sessionId"` + LeafID *string `json:"leafId"` + Revision string `json:"revision"` + Nodes []PiTreeNode `json:"nodes"` +} + +type PiTreeNavigateInput struct { + TargetID string `json:"targetId"` + Revision string `json:"revision"` + Summarize bool `json:"summarize" default:"false"` +} + +type PiTreeNavigateResult struct { + Tree PiTreeSnapshot `json:"tree"` + EditorText string `json:"editorText,omitempty"` +} + +type piHistoryTreeNode struct { + Entry piHistoryEntry `json:"entry"` + Children []piHistoryTreeNode `json:"children"` + Label string `json:"label"` +} + +type piTreeRawNode struct { + entry piHistoryEntry + label string + children []string +} + +type piBridgeMarkerData struct { + TargetID string `json:"targetId"` + Summarize bool `json:"summarize"` + Nonce string `json:"nonce"` +} + +func (m *Manager) GetPiSessionTree(ctx context.Context, sessionID string) (PiTreeSnapshot, error) { + if m == nil { + return PiTreeSnapshot{}, errors.New("web session manager is not configured") + } + dispatchLock := &m.sessionDispatchLocks[sessionRevisionLockIndex(sessionID)] + dispatchLock.Lock() + defer dispatchLock.Unlock() + + session, err := m.piTreeSession(ctx, sessionID) + if err != nil { + return PiTreeSnapshot{}, err + } + runtime, err := m.getOrStartPiRuntime(ctx, session) + if err != nil { + return PiTreeSnapshot{}, err + } + defer runtime.scheduleIdle() + return m.readPiTreeSnapshot(ctx, runtime, session) +} + +func (m *Manager) NavigatePiSessionTree( + ctx context.Context, + sessionID string, + input PiTreeNavigateInput, +) (PiTreeNavigateResult, error) { + if m == nil { + return PiTreeNavigateResult{}, errors.New("web session manager is not configured") + } + dispatchLock := &m.sessionDispatchLocks[sessionRevisionLockIndex(sessionID)] + dispatchLock.Lock() + defer dispatchLock.Unlock() + + session, err := m.piTreeSession(ctx, sessionID) + if err != nil { + return PiTreeNavigateResult{}, err + } + if m.hasActiveRun(session.ID) { + return PiTreeNavigateResult{}, errors.New("cannot navigate an active Pi web session") + } + if len(m.pendingInputsDisplaySnapshot(session.ID)) > 0 { + return PiTreeNavigateResult{}, errors.New("cannot navigate while messages are pending") + } + targetID := strings.TrimSpace(input.TargetID) + if targetID == "" { + return PiTreeNavigateResult{}, errors.New("Pi tree target id is required") + } + expectedRevision := strings.TrimSpace(input.Revision) + if expectedRevision == "" { + return PiTreeNavigateResult{}, errors.New("Pi tree revision is required") + } + + runtime, err := m.getOrStartPiRuntime(ctx, session) + if err != nil { + return PiTreeNavigateResult{}, err + } + navigationSent := false + navigationComplete := false + defer func() { + if navigationSent && !navigationComplete { + runtime.stop(errors.New("Pi tree navigation did not reach its verified completion boundary")) + return + } + runtime.scheduleIdle() + }() + current, rawCurrent, err := m.readPiTreeSnapshotRaw(ctx, runtime, session) + if err != nil { + return PiTreeNavigateResult{}, err + } + if current.Revision != expectedRevision { + return PiTreeNavigateResult{}, ErrPiTreeRevisionConflict + } + target, ok := rawCurrent[targetID] + if !ok || isPiBridgeMarker(target.entry) { + return PiTreeNavigateResult{}, errors.New("Pi tree target does not exist") + } + + nonce := utils.NewID() + payload, err := json.Marshal(piBridgeMarkerData{TargetID: targetID, Summarize: input.Summarize, Nonce: nonce}) + if err != nil { + return PiTreeNavigateResult{}, err + } + message := "/" + piBridgeCommandName + " " + base64.RawURLEncoding.EncodeToString(payload) + operationCtx, cancel := context.WithTimeout(context.Background(), piRPCRequestTimeout) + defer cancel() + navigationSent = true + if err := runtime.client.Request(operationCtx, "prompt", map[string]any{"message": message}, nil); err != nil { + return PiTreeNavigateResult{}, fmt.Errorf("start Pi tree navigation: %w", err) + } + + entries, marker, err := waitForPiBridgeMarker(operationCtx, runtime.client, nonce) + if err != nil { + return PiTreeNavigateResult{}, err + } + if err := validatePiNavigationMarker(target.entry, marker, entries.Entries, input.Summarize); err != nil { + return PiTreeNavigateResult{}, err + } + + fresh, rawFresh, err := m.readPiTreeSnapshotRaw(operationCtx, runtime, session) + if err != nil { + return PiTreeNavigateResult{}, err + } + freshMarker, ok := rawFresh[marker.ID] + freshMarkerData, markerOK := parsePiBridgeMarker(freshMarker.entry) + if !ok || !markerOK || freshMarkerData.Nonce != nonce || + pointerString(entries.LeafID) != marker.ID || pointerString(freshMarker.entry.ParentID) != pointerString(marker.ParentID) { + return PiTreeNavigateResult{}, errors.New("Pi tree navigation marker does not match the active leaf") + } + if pointerString(fresh.LeafID) != pointerString(marker.ParentID) { + return PiTreeNavigateResult{}, errors.New("Pi tree navigation produced an unexpected logical leaf") + } + + if err := m.projectPiHistoryEntries(operationCtx, session, entries); err != nil { + return PiTreeNavigateResult{}, err + } + navigationComplete = true + if err := m.broadcastSnapshot(context.Background(), session.ID); err != nil && m.logger != nil { + m.logger.Warn("failed to broadcast Pi tree navigation snapshot", + zap.String("sessionId", session.ID), + zap.Error(err), + ) + } + return PiTreeNavigateResult{Tree: fresh, EditorText: piTreeEditorText(target.entry)}, nil +} + +func (m *Manager) piTreeSession(ctx context.Context, sessionID string) (tables.WebSessionTable, error) { + session, err := m.GetSession(ctx, strings.TrimSpace(sessionID)) + if err != nil { + return tables.WebSessionTable{}, err + } + if session.ArchivedAt != nil { + return tables.WebSessionTable{}, errors.New("session is archived") + } + if normalizeAgent(Agent(session.Agent)) != AgentPi || effectiveSessionBackend(session) != SessionBackendPiRPC { + return tables.WebSessionTable{}, errors.New("session tree is only supported for Pi web sessions") + } + if strings.TrimSpace(pointerString(session.NativeSessionID)) == "" || strings.TrimSpace(pointerString(session.ThreadPath)) == "" { + return tables.WebSessionTable{}, errors.New("Pi session tree requires an existing native session") + } + if err := m.EnsureProjectPiTrust(ctx, session.ProjectID, session.Cwd); err != nil { + return tables.WebSessionTable{}, err + } + return session, nil +} + +func (m *Manager) readPiTreeSnapshot( + ctx context.Context, + runtime *piSessionRuntime, + session tables.WebSessionTable, +) (PiTreeSnapshot, error) { + snapshot, _, err := m.readPiTreeSnapshotRaw(ctx, runtime, session) + return snapshot, err +} + +func (m *Manager) readPiTreeSnapshotRaw( + ctx context.Context, + runtime *piSessionRuntime, + session tables.WebSessionTable, +) (PiTreeSnapshot, map[string]piTreeRawNode, error) { + var response struct { + Tree []piHistoryTreeNode `json:"tree"` + LeafID *string `json:"leafId"` + } + if err := runtime.client.Request(ctx, "get_tree", nil, &response); err != nil { + return PiTreeSnapshot{}, nil, fmt.Errorf("read Pi session tree: %w", err) + } + revision := piSourceRevision(pointerString(session.ThreadPath), pointerString(response.LeafID)) + if revision == "" { + return PiTreeSnapshot{}, nil, errors.New("Pi session tree revision is unavailable") + } + snapshot, raw, err := projectPiTree(pointerString(session.NativeSessionID), revision, response.Tree, pointerString(response.LeafID)) + if err != nil { + return PiTreeSnapshot{}, nil, err + } + return snapshot, raw, nil +} + +func projectPiTree( + sessionID string, + revision string, + roots []piHistoryTreeNode, + rawLeafID string, +) (PiTreeSnapshot, map[string]piTreeRawNode, error) { + raw := make(map[string]piTreeRawNode) + order := make([]string, 0) + type pendingNode struct { + node piHistoryTreeNode + nestedParentID string + root bool + } + stack := make([]pendingNode, 0, len(roots)) + for index := len(roots) - 1; index >= 0; index-- { + stack = append(stack, pendingNode{node: roots[index], root: true}) + } + for len(stack) > 0 { + last := len(stack) - 1 + current := stack[last] + stack = stack[:last] + entry := current.node.Entry + entry.ID = strings.TrimSpace(entry.ID) + entry.Type = strings.TrimSpace(entry.Type) + if entry.ID == "" || entry.Type == "" { + return PiTreeSnapshot{}, nil, errors.New("Pi session tree contains an invalid node") + } + if len(raw) >= piTreeMaxNodes { + return PiTreeSnapshot{}, nil, errors.New("Pi session tree exceeds the node limit") + } + if _, duplicate := raw[entry.ID]; duplicate { + return PiTreeSnapshot{}, nil, errors.New("Pi session tree contains a duplicate node id") + } + parentID := strings.TrimSpace(pointerString(entry.ParentID)) + if current.root { + if parentID != "" { + return PiTreeSnapshot{}, nil, errors.New("Pi session tree contains an orphan root") + } + } else if parentID == "" || parentID != current.nestedParentID { + return PiTreeSnapshot{}, nil, errors.New("Pi session tree parent linkage is invalid") + } + children := make([]string, 0, len(current.node.Children)) + for _, child := range current.node.Children { + children = append(children, strings.TrimSpace(child.Entry.ID)) + } + raw[entry.ID] = piTreeRawNode{entry: entry, label: strings.TrimSpace(current.node.Label), children: children} + order = append(order, entry.ID) + for index := len(current.node.Children) - 1; index >= 0; index-- { + stack = append(stack, pendingNode{node: current.node.Children[index], nestedParentID: entry.ID}) + } + } + + rawLeafID = strings.TrimSpace(rawLeafID) + if len(raw) == 0 { + if rawLeafID != "" { + return PiTreeSnapshot{}, nil, errors.New("Pi session tree leaf is missing") + } + return PiTreeSnapshot{SessionID: strings.TrimSpace(sessionID), Revision: revision, Nodes: []PiTreeNode{}}, raw, nil + } + if _, ok := raw[rawLeafID]; !ok { + return PiTreeSnapshot{}, nil, errors.New("Pi session tree leaf does not exist") + } + + logicalLeafID := rawLeafID + for logicalLeafID != "" && isPiBridgeMarker(raw[logicalLeafID].entry) { + logicalLeafID = strings.TrimSpace(pointerString(raw[logicalLeafID].entry.ParentID)) + } + active := make(map[string]struct{}) + seen := make(map[string]struct{}) + for id := logicalLeafID; id != ""; { + if _, duplicate := seen[id]; duplicate { + return PiTreeSnapshot{}, nil, errors.New("Pi session tree active path contains a cycle") + } + seen[id] = struct{}{} + node, ok := raw[id] + if !ok { + return PiTreeSnapshot{}, nil, errors.New("Pi session tree active path is incomplete") + } + if !isPiBridgeMarker(node.entry) { + active[id] = struct{}{} + } + id = strings.TrimSpace(pointerString(node.entry.ParentID)) + } + + nodes := make([]PiTreeNode, 0, len(order)) + byVisibleID := make(map[string]int) + for _, id := range order { + node := raw[id] + if isPiBridgeMarker(node.entry) { + continue + } + parentID := strings.TrimSpace(pointerString(node.entry.ParentID)) + for parentID != "" { + parent, ok := raw[parentID] + if !ok { + return PiTreeSnapshot{}, nil, errors.New("Pi session tree contains an incomplete parent chain") + } + if !isPiBridgeMarker(parent.entry) { + break + } + parentID = strings.TrimSpace(pointerString(parent.entry.ParentID)) + } + projected := PiTreeNode{ + ID: id, ParentID: nilIfEmpty(parentID), Type: node.entry.Type, + Role: piTreeRole(node.entry), Preview: piTreePreview(node.entry), + Timestamp: strings.TrimSpace(node.entry.Timestamp), Label: node.label, + Children: []string{}, + } + _, projected.Active = active[id] + byVisibleID[id] = len(nodes) + nodes = append(nodes, projected) + } + for index := range nodes { + parentID := strings.TrimSpace(pointerString(nodes[index].ParentID)) + if parentID == "" { + continue + } + parentIndex, ok := byVisibleID[parentID] + if !ok { + return PiTreeSnapshot{}, nil, errors.New("Pi session tree projected parent is missing") + } + nodes[parentIndex].Children = append(nodes[parentIndex].Children, nodes[index].ID) + } + + return PiTreeSnapshot{ + SessionID: strings.TrimSpace(sessionID), LeafID: nilIfEmpty(logicalLeafID), + Revision: revision, Nodes: nodes, + }, raw, nil +} + +func waitForPiBridgeMarker( + ctx context.Context, + client *piRPCClient, + nonce string, +) (piHistoryEntriesResponse, piHistoryEntry, error) { + ticker := time.NewTicker(25 * time.Millisecond) + defer ticker.Stop() + for { + var response piHistoryEntriesResponse + if err := client.Request(ctx, "get_entries", nil, &response); err != nil { + return piHistoryEntriesResponse{}, piHistoryEntry{}, fmt.Errorf("verify Pi tree navigation: %w", err) + } + for _, entry := range response.Entries { + marker, ok := parsePiBridgeMarker(entry) + if ok && marker.Nonce == nonce { + if pointerString(response.LeafID) != entry.ID { + return piHistoryEntriesResponse{}, piHistoryEntry{}, errors.New("Pi tree navigation marker is not the active leaf") + } + return response, entry, nil + } + } + select { + case <-ctx.Done(): + return piHistoryEntriesResponse{}, piHistoryEntry{}, fmt.Errorf("wait for Pi tree navigation marker: %w", ctx.Err()) + case <-ticker.C: + } + } +} + +func validatePiNavigationMarker( + target piHistoryEntry, + marker piHistoryEntry, + entries []piHistoryEntry, + summarize bool, +) error { + data, ok := parsePiBridgeMarker(marker) + if !ok || data.TargetID != strings.TrimSpace(target.ID) || data.Summarize != summarize { + return errors.New("Pi tree navigation marker payload does not match the request") + } + expectedParent := strings.TrimSpace(target.ID) + if target.Type == "custom_message" || (target.Type == "message" && strings.EqualFold(strings.TrimSpace(target.Message.Role), "user")) { + expectedParent = strings.TrimSpace(pointerString(target.ParentID)) + } + actualParent := strings.TrimSpace(pointerString(marker.ParentID)) + if actualParent == expectedParent { + return nil + } + if !summarize || actualParent == "" { + return errors.New("Pi tree navigation marker parent does not match the target semantics") + } + for _, entry := range entries { + if strings.TrimSpace(entry.ID) == actualParent && entry.Type == "branch_summary" && strings.TrimSpace(pointerString(entry.ParentID)) == expectedParent { + return nil + } + } + return errors.New("Pi tree navigation summary is not attached to the target branch") +} + +func isPiBridgeMarker(entry piHistoryEntry) bool { + return strings.TrimSpace(entry.Type) == "custom" && strings.TrimSpace(entry.CustomType) == piBridgeMarkerType +} + +func parsePiBridgeMarker(entry piHistoryEntry) (piBridgeMarkerData, bool) { + if !isPiBridgeMarker(entry) || len(entry.Data) == 0 { + return piBridgeMarkerData{}, false + } + var data piBridgeMarkerData + if json.Unmarshal(entry.Data, &data) != nil { + return piBridgeMarkerData{}, false + } + data.TargetID = strings.TrimSpace(data.TargetID) + data.Nonce = strings.TrimSpace(data.Nonce) + return data, data.TargetID != "" && data.Nonce != "" +} + +func piTreeRole(entry piHistoryEntry) string { + if entry.Type == "message" { + return strings.TrimSpace(entry.Message.Role) + } + if entry.Type == "custom_message" { + return "user" + } + return "" +} + +func piTreePreview(entry piHistoryEntry) string { + var value string + switch entry.Type { + case "message": + value = piHistoryContentText(entry.Message.Content) + case "custom_message": + value = piHistoryContentText(entry.Content) + case "compaction", "branch_summary": + value = entry.Summary + case "model_change": + value = strings.Trim(strings.TrimSpace(entry.Provider)+"/"+strings.TrimSpace(entry.ModelID), "/") + case "thinking_level_change": + value = entry.ThinkingLevel + case "custom": + value = entry.CustomType + case "session_info": + value = entry.Name + } + value = strings.TrimSpace(value) + if line, _, found := strings.Cut(value, "\n"); found { + value = strings.TrimSpace(line) + } + runes := []rune(value) + if len(runes) > piTreePreviewRunes { + value = string(runes[:piTreePreviewRunes-3]) + "..." + } + return value +} + +func piTreeEditorText(entry piHistoryEntry) string { + if entry.Type == "message" && strings.EqualFold(strings.TrimSpace(entry.Message.Role), "user") { + return piTreeEditorContentText(entry.Message.Content) + } + if entry.Type == "custom_message" { + return piTreeEditorContentText(entry.Content) + } + return "" +} + +func piTreeEditorContentText(raw json.RawMessage) string { + if len(raw) == 0 { + return "" + } + var text string + if json.Unmarshal(raw, &text) == nil { + return text + } + var blocks []struct { + Type string `json:"type"` + Text string `json:"text"` + } + if json.Unmarshal(raw, &blocks) != nil { + return "" + } + var builder strings.Builder + for _, block := range blocks { + if strings.EqualFold(block.Type, "text") { + builder.WriteString(block.Text) + } + } + return builder.String() +} diff --git a/service/websession/pi_tree_mutation.go b/service/websession/pi_tree_mutation.go new file mode 100644 index 00000000..dcff5204 --- /dev/null +++ b/service/websession/pi_tree_mutation.go @@ -0,0 +1,349 @@ +package websession + +import ( + "context" + "errors" + "fmt" + "os" + "path/filepath" + "strings" + "time" + + "code-kanban/model" + "code-kanban/model/tables" + + "gorm.io/gorm" +) + +type PiTreeForkInput struct { + TargetID string `json:"targetId"` + Revision string `json:"revision"` +} + +type PiTreeCloneInput struct { + Revision string `json:"revision"` +} + +type PiTreeCreateResult struct { + Session SessionSummary `json:"session"` + Tree PiTreeSnapshot `json:"tree"` + EditorText string `json:"editorText,omitempty"` +} + +type piTreeMutationResult struct { + Text string `json:"text"` + Cancelled bool `json:"cancelled"` +} + +func (m *Manager) ForkPiSessionTree( + ctx context.Context, + sessionID string, + input PiTreeForkInput, +) (PiTreeCreateResult, error) { + return m.createPiSessionFromTree(ctx, sessionID, "fork", strings.TrimSpace(input.TargetID), strings.TrimSpace(input.Revision)) +} + +func (m *Manager) ClonePiSessionTree( + ctx context.Context, + sessionID string, + input PiTreeCloneInput, +) (PiTreeCreateResult, error) { + return m.createPiSessionFromTree(ctx, sessionID, "clone", "", strings.TrimSpace(input.Revision)) +} + +func (m *Manager) createPiSessionFromTree( + ctx context.Context, + sessionID string, + operation string, + targetID string, + expectedRevision string, +) (PiTreeCreateResult, error) { + if m == nil { + return PiTreeCreateResult{}, errors.New("web session manager is not configured") + } + if expectedRevision == "" { + return PiTreeCreateResult{}, errors.New("Pi tree revision is required") + } + if operation == "fork" && targetID == "" { + return PiTreeCreateResult{}, errors.New("Pi tree fork target id is required") + } + if operation != "fork" && operation != "clone" { + return PiTreeCreateResult{}, errors.New("unsupported Pi tree mutation") + } + + dispatchLock := &m.sessionDispatchLocks[sessionRevisionLockIndex(sessionID)] + dispatchLock.Lock() + defer dispatchLock.Unlock() + + source, err := m.piTreeSession(ctx, sessionID) + if err != nil { + return PiTreeCreateResult{}, err + } + if m.hasActiveRun(source.ID) { + return PiTreeCreateResult{}, errors.New("cannot mutate an active Pi web session") + } + if len(m.pendingInputsDisplaySnapshot(source.ID)) > 0 { + return PiTreeCreateResult{}, errors.New("cannot mutate a Pi session while messages are pending") + } + + runtime, err := m.getOrStartPiRuntime(ctx, source) + if err != nil { + return PiTreeCreateResult{}, err + } + mutationSent := false + defer func() { + if mutationSent { + runtime.stop(errors.New("Pi tree mutation changed the native session")) + return + } + runtime.scheduleIdle() + }() + + current, rawCurrent, err := m.readPiTreeSnapshotRaw(ctx, runtime, source) + if err != nil { + return PiTreeCreateResult{}, err + } + if current.Revision != expectedRevision { + return PiTreeCreateResult{}, ErrPiTreeRevisionConflict + } + if operation == "fork" { + node, ok := rawCurrent[targetID] + if !ok || isPiBridgeMarker(node.entry) || !piTreeForkableEntry(node.entry) { + return PiTreeCreateResult{}, errors.New("Pi tree fork target is not a user message") + } + } + + operationCtx, cancel := context.WithTimeout(context.Background(), piRPCRequestTimeout) + defer cancel() + payload := map[string]any(nil) + if operation == "fork" { + payload = map[string]any{"entryId": targetID} + } + mutationSent = true + var nativeResult piTreeMutationResult + if err := runtime.client.Request(operationCtx, operation, payload, &nativeResult); err != nil { + return PiTreeCreateResult{}, fmt.Errorf("Pi tree %s failed: %w", operation, err) + } + if nativeResult.Cancelled { + return PiTreeCreateResult{}, fmt.Errorf("Pi tree %s was cancelled", operation) + } + + var state piRPCState + if err := runtime.client.Request(operationCtx, "get_state", nil, &state); err != nil { + return PiTreeCreateResult{}, fmt.Errorf("read forked Pi session state: %w", err) + } + if err := validateNewPiTreeSessionIdentity(source, state); err != nil { + return PiTreeCreateResult{}, err + } + + targetRecord, err := newPiTreeSessionRecord(source, operation, state) + if err != nil { + return PiTreeCreateResult{}, err + } + var entries piHistoryEntriesResponse + if err := runtime.client.Request(operationCtx, "get_entries", nil, &entries); err != nil { + return PiTreeCreateResult{}, fmt.Errorf("read forked Pi session entries: %w", err) + } + if operation == "clone" && strings.TrimSpace(pointerString(entries.LeafID)) == "" { + return PiTreeCreateResult{}, errors.New("cloned Pi session has no active leaf") + } + targetRecord.NativeLeafID = nilIfEmpty(pointerString(entries.LeafID)) + + var treeResponse struct { + Tree []piHistoryTreeNode `json:"tree"` + LeafID *string `json:"leafId"` + } + if err := runtime.client.Request(operationCtx, "get_tree", nil, &treeResponse); err != nil { + return PiTreeCreateResult{}, fmt.Errorf("read forked Pi session tree: %w", err) + } + if pointerString(treeResponse.LeafID) != pointerString(entries.LeafID) { + return PiTreeCreateResult{}, errors.New("forked Pi tree and entry leaf do not match") + } + revision := piSourceRevision(state.SessionFile, pointerString(treeResponse.LeafID)) + if revision == "" { + return PiTreeCreateResult{}, errors.New("forked Pi tree revision is unavailable") + } + tree, _, err := projectPiTree(state.SessionID, revision, treeResponse.Tree, pointerString(treeResponse.LeafID)) + if err != nil { + return PiTreeCreateResult{}, err + } + + var stats piRPCSessionStats + if err := runtime.client.Request(operationCtx, "get_session_stats", nil, &stats); err != nil { + return PiTreeCreateResult{}, fmt.Errorf("read forked Pi session stats: %w", err) + } + if err := validatePiMutationStats(state, stats); err != nil { + return PiTreeCreateResult{}, err + } + projection, err := buildPiHistoryProjection(targetRecord, entries) + if err != nil { + return PiTreeCreateResult{}, err + } + applyPiMutationStats(projection.updates, state, stats) + if err := m.createProjectedPiTreeSession(operationCtx, &targetRecord, projection); err != nil { + return PiTreeCreateResult{}, err + } + + created, err := m.GetSession(operationCtx, targetRecord.ID) + if err != nil { + return PiTreeCreateResult{}, err + } + m.broadcastProjectSessionSummaries(context.Background(), source.ProjectID) + return PiTreeCreateResult{ + Session: m.mapSessionSummary(created), + Tree: tree, + EditorText: func() string { + if operation == "fork" { + return nativeResult.Text + } + return "" + }(), + }, nil +} + +func piTreeForkableEntry(entry piHistoryEntry) bool { + return strings.TrimSpace(entry.Type) == "message" && strings.EqualFold(strings.TrimSpace(entry.Message.Role), "user") +} + +func validateNewPiTreeSessionIdentity(source tables.WebSessionTable, state piRPCState) error { + if strings.TrimSpace(state.SessionID) == "" || strings.TrimSpace(state.SessionFile) == "" { + return errors.New("Pi tree mutation returned an incomplete session identity") + } + if strings.TrimSpace(state.SessionID) == strings.TrimSpace(pointerString(source.NativeSessionID)) { + return errors.New("Pi tree mutation did not create a new native session") + } + if samePiRuntimePath(state.SessionFile, pointerString(source.ThreadPath)) { + return errors.New("Pi tree mutation reused the source session file") + } + candidate := source + candidate.NativeSessionID = nilIfEmpty(state.SessionID) + candidate.ThreadPath = nilIfEmpty(filepath.Clean(state.SessionFile)) + return validatePiRuntimeState(candidate, state) +} + +func newPiTreeSessionRecord( + source tables.WebSessionTable, + operation string, + state piRPCState, +) (tables.WebSessionTable, error) { + info, err := os.Stat(state.SessionFile) + if err != nil { + return tables.WebSessionTable{}, fmt.Errorf("stat forked Pi session: %w", err) + } + if !info.Mode().IsRegular() { + return tables.WebSessionTable{}, errors.New("forked Pi session file is not regular") + } + now := time.Now() + prefix := "Clone of " + if operation == "fork" { + prefix = "Fork of " + } + title := prefix + strings.TrimSpace(source.Title) + if strings.TrimSpace(source.Title) == "" { + title = prefix + "Pi session" + } + modelName := canonicalPiModel(state.Model) + if modelName == "" { + modelName = strings.TrimSpace(source.Model) + } + reasoning := piThinkingLevelToReasoning(state.ThinkingLevel) + if reasoning == ReasoningEffortDefault { + reasoning = normalizeReasoningEffort(ReasoningEffort(source.ReasoningEffort)) + } + updatedAt := info.ModTime() + record := tables.WebSessionTable{ + ProjectID: source.ProjectID, WorktreeID: source.WorktreeID, + Agent: string(AgentPi), ClaudeRuntime: source.ClaudeRuntime, + Backend: string(SessionBackendPiRPC), Title: title, TitleAuto: false, + Model: modelName, ReasoningEffort: string(reasoning), WorkflowMode: source.WorkflowMode, + PermissionLevel: source.PermissionLevel, ActiveCallTimeoutEnabled: source.ActiveCallTimeoutEnabled, + AutoRetryEnabled: source.AutoRetryEnabled, AutoRetryScope: source.AutoRetryScope, + AutoRetryPreset: source.AutoRetryPreset, AutoRetryMaxAttempts: source.AutoRetryMaxAttempts, + AutoRetryDispatchPendingOnFailure: source.AutoRetryDispatchPendingOnFailure, + LegacyPermissionMode: source.LegacyPermissionMode, Cwd: source.Cwd, + NativeSessionID: nilIfEmpty(state.SessionID), ThreadPath: nilIfEmpty(filepath.Clean(state.SessionFile)), + Status: string(StatusIdle), AssistantState: "", HasUnread: false, + ActivityAt: now, StatusUpdatedAt: &now, SourceKind: string(SessionBackendPiRPC), + SyncState: string(SyncStateFresh), SourceUpdatedAt: &updatedAt, + ThreadPreview: nilIfEmpty(title), AutoRetryAttempt: 0, + } + record.Init() + return record, nil +} + +func validatePiMutationStats(state piRPCState, stats piRPCSessionStats) error { + if id := strings.TrimSpace(stats.SessionID); id != "" && id != strings.TrimSpace(state.SessionID) { + return errors.New("forked Pi session stats id does not match state") + } + if path := strings.TrimSpace(stats.SessionFile); path != "" && !samePiRuntimePath(path, state.SessionFile) { + return errors.New("forked Pi session stats file does not match state") + } + return nil +} + +func applyPiMutationStats(updates map[string]any, state piRPCState, stats piRPCSessionStats) { + updates["native_session_id"] = strings.TrimSpace(state.SessionID) + updates["thread_path"] = filepath.Clean(state.SessionFile) + updates["total_input_tokens"] = stats.Tokens.Input + updates["total_cached_input_tokens"] = stats.Tokens.CacheRead + updates["total_output_tokens"] = stats.Tokens.Output + updates["total_cost"] = stats.Cost + if stats.ContextUsage != nil { + updates["session_context_window_tokens"] = stats.ContextUsage.ContextWindow + updates["session_context_window_observed_at"] = time.Now() + updates["latest_token_count_total_tokens"] = stats.ContextUsage.Tokens + updates["latest_token_count_updated_at"] = time.Now() + } +} + +func (m *Manager) createProjectedPiTreeSession( + ctx context.Context, + record *tables.WebSessionTable, + projection piHistoryProjection, +) error { + if record == nil { + return errors.New("forked Pi web session is missing") + } + db := model.GetDB() + if db == nil { + return model.ErrDBNotInitialized + } + return db.WithContext(ctx).Transaction(func(tx *gorm.DB) error { + var duplicate int64 + if err := tx.Unscoped().Model(&tables.WebSessionTable{}). + Where("project_id = ? AND agent = ? AND native_session_id = ?", record.ProjectID, string(AgentPi), pointerString(record.NativeSessionID)). + Count(&duplicate).Error; err != nil { + return err + } + if duplicate > 0 { + return errors.New("forked Pi session is already linked to a web session") + } + var maxOrder float64 + if err := tx.Model(&tables.WebSessionTable{}). + Where("project_id = ? AND archived_at IS NULL", record.ProjectID). + Select("COALESCE(MAX(order_index), 0)").Scan(&maxOrder).Error; err != nil { + return err + } + record.OrderIndex = maxOrder + sessionOrderStep + if err := tx.Create(record).Error; err != nil { + return err + } + if len(projection.turns) > 0 { + if err := tx.Create(&projection.turns).Error; err != nil { + return err + } + } + if len(projection.items) > 0 { + if err := tx.Create(&projection.items).Error; err != nil { + return err + } + } + if len(projection.updates) > 0 { + if err := tx.Model(&tables.WebSessionTable{}).Where("id = ?", record.ID). + Updates(withSnapshotRevisionIncrement(projection.updates)).Error; err != nil { + return err + } + } + return nil + }) +} diff --git a/service/websession/pi_tree_test.go b/service/websession/pi_tree_test.go new file mode 100644 index 00000000..cdf6b9be --- /dev/null +++ b/service/websession/pi_tree_test.go @@ -0,0 +1,204 @@ +package websession + +import ( + "context" + "encoding/json" + "errors" + "strings" + "testing" +) + +func TestProjectPiTreeFiltersBridgeMarkersAndPreservesActivePath(t *testing.T) { + rootID := "user-root" + markerID := "marker-1" + markerData, err := json.Marshal(piBridgeMarkerData{TargetID: rootID, Nonce: "nonce-1"}) + if err != nil { + t.Fatal(err) + } + roots := []piHistoryTreeNode{ + { + Entry: piHistoryEntry{ + Type: "message", ID: rootID, Timestamp: "2026-05-01T01:00:00Z", + Message: piHistoryMessage{Role: "user", Content: json.RawMessage(`"root prompt"`)}, + }, + Children: []piHistoryTreeNode{ + { + Entry: piHistoryEntry{Type: "custom", ID: markerID, ParentID: &rootID, CustomType: piBridgeMarkerType, Data: markerData}, + Children: []piHistoryTreeNode{{ + Entry: piHistoryEntry{ + Type: "message", ID: "assistant-new", ParentID: &markerID, + Message: piHistoryMessage{Role: "assistant", Content: json.RawMessage(`[{"type":"text","text":"new branch"}]`)}, + }, + }}, + }, + { + Entry: piHistoryEntry{ + Type: "message", ID: "assistant-old", ParentID: &rootID, + Message: piHistoryMessage{Role: "assistant", Content: json.RawMessage(`"old branch"`)}, + }, + }, + }, + }, + } + + snapshot, raw, err := projectPiTree("native-1", "revision-1", roots, "assistant-new") + if err != nil { + t.Fatalf("projectPiTree: %v", err) + } + if snapshot.SessionID != "native-1" || snapshot.Revision != "revision-1" || pointerString(snapshot.LeafID) != "assistant-new" { + t.Fatalf("unexpected snapshot identity: %#v", snapshot) + } + if len(raw) != 4 || len(snapshot.Nodes) != 3 { + t.Fatalf("marker was not retained internally and filtered externally: raw=%d nodes=%#v", len(raw), snapshot.Nodes) + } + byID := make(map[string]PiTreeNode, len(snapshot.Nodes)) + for _, node := range snapshot.Nodes { + byID[node.ID] = node + if node.ID == markerID { + t.Fatalf("bridge marker leaked into projected tree: %#v", node) + } + } + if pointerString(byID["assistant-new"].ParentID) != rootID { + t.Fatalf("marker child was not reparented to visible ancestor: %#v", byID["assistant-new"]) + } + if !byID[rootID].Active || !byID["assistant-new"].Active || byID["assistant-old"].Active { + t.Fatalf("unexpected active path: %#v", snapshot.Nodes) + } + if byID[rootID].Preview != "root prompt" || byID["assistant-new"].Preview != "new branch" { + t.Fatalf("unexpected previews: %#v", snapshot.Nodes) + } +} + +func TestProjectPiTreeUsesMarkerParentAsLogicalLeaf(t *testing.T) { + rootID := "root-user" + markerData, _ := json.Marshal(piBridgeMarkerData{TargetID: rootID, Nonce: "nonce-root"}) + roots := []piHistoryTreeNode{ + {Entry: piHistoryEntry{Type: "message", ID: rootID, Message: piHistoryMessage{Role: "user", Content: json.RawMessage(`"original"`)}}}, + {Entry: piHistoryEntry{Type: "custom", ID: "marker-root", CustomType: piBridgeMarkerType, Data: markerData}}, + } + snapshot, _, err := projectPiTree("native-root", "revision-root", roots, "marker-root") + if err != nil { + t.Fatalf("projectPiTree: %v", err) + } + if snapshot.LeafID != nil { + t.Fatalf("root navigation marker should expose a nil logical leaf: %#v", snapshot) + } + if len(snapshot.Nodes) != 1 || snapshot.Nodes[0].Active { + t.Fatalf("root navigation should leave no visible node active: %#v", snapshot.Nodes) + } +} + +func TestProjectPiTreeRejectsMalformedForests(t *testing.T) { + root := piHistoryTreeNode{Entry: piHistoryEntry{Type: "message", ID: "root"}} + orphanParent := "missing" + duplicateParent := "root" + for name, testCase := range map[string]struct { + roots []piHistoryTreeNode + leaf string + }{ + "orphan root": { + roots: []piHistoryTreeNode{{Entry: piHistoryEntry{Type: "message", ID: "orphan", ParentID: &orphanParent}}}, + leaf: "orphan", + }, + "self parent root": { + roots: []piHistoryTreeNode{{Entry: piHistoryEntry{Type: "message", ID: "self", ParentID: stringPointer("self")}}}, + leaf: "self", + }, + "duplicate id": { + roots: []piHistoryTreeNode{{ + Entry: piHistoryEntry{Type: "message", ID: "root"}, + Children: []piHistoryTreeNode{{Entry: piHistoryEntry{Type: "message", ID: "root", ParentID: &duplicateParent}}}, + }}, + leaf: "root", + }, + "nested parent mismatch": { + roots: []piHistoryTreeNode{{ + Entry: piHistoryEntry{Type: "message", ID: "root"}, + Children: []piHistoryTreeNode{{Entry: piHistoryEntry{Type: "message", ID: "child", ParentID: &orphanParent}}}, + }}, + leaf: "child", + }, + "missing leaf": {roots: []piHistoryTreeNode{root}, leaf: "missing"}, + } { + t.Run(name, func(t *testing.T) { + if _, _, err := projectPiTree("native", "revision", testCase.roots, testCase.leaf); err == nil { + t.Fatal("expected malformed tree rejection") + } + }) + } +} + +func TestValidatePiNavigationMarker(t *testing.T) { + parentID := "parent" + target := piHistoryEntry{Type: "message", ID: "target-user", ParentID: &parentID, Message: piHistoryMessage{Role: "user"}} + data, _ := json.Marshal(piBridgeMarkerData{TargetID: target.ID, Nonce: "nonce"}) + marker := piHistoryEntry{Type: "custom", ID: "marker", ParentID: &parentID, CustomType: piBridgeMarkerType, Data: data} + if err := validatePiNavigationMarker(target, marker, []piHistoryEntry{target, marker}, false); err != nil { + t.Fatalf("valid navigation marker: %v", err) + } + wrongParent := "wrong" + marker.ParentID = &wrongParent + if err := validatePiNavigationMarker(target, marker, []piHistoryEntry{target, marker}, false); err == nil || !strings.Contains(err.Error(), "parent") { + t.Fatalf("expected parent mismatch, got %v", err) + } +} + +func TestPiTreeEditorContentTextPreservesOriginalWhitespace(t *testing.T) { + for name, testCase := range map[string]struct { + raw json.RawMessage + want string + }{ + "string": { + raw: json.RawMessage(`" first line\nsecond line "`), + want: " first line\nsecond line ", + }, + "text blocks": { + raw: json.RawMessage(`[{"type":"text","text":" first "},{"type":"image","data":"ignored"},{"type":"text","text":"second "}]`), + want: " first second ", + }, + } { + t.Run(name, func(t *testing.T) { + if got := piTreeEditorContentText(testCase.raw); got != testCase.want { + t.Fatalf("editor text changed whitespace: got=%q want=%q", got, testCase.want) + } + }) + } +} + +func TestClassifyPiTreeErrorDoesNotExposeInternalDetails(t *testing.T) { + for name, testCase := range map[string]struct { + err error + code string + messagePart string + }{ + "revision": {ErrPiTreeRevisionConflict, "conflict", "refresh"}, + "active": {errors.New("cannot navigate an active Pi web session"), "invalid_state", "current session state"}, + "input": {errors.New("Pi tree revision is required"), "bad_req", "Invalid Pi session tree request"}, + "internal": {errors.New("native-secret-response"), "internal", "operation failed"}, + } { + t.Run(name, func(t *testing.T) { + classified := ClassifyPiTreeError(testCase.err) + if classified.Code != testCase.code || !strings.Contains(classified.Message, testCase.messagePart) { + t.Fatalf("unexpected classification: %#v", classified) + } + frame := newPiTreeErrorFrame(wireCommandFrame{RequestID: "request", SessionID: "session"}, testCase.err) + if frame.Code != testCase.code || frame.Message != classified.Message { + t.Fatalf("wire classification drifted: %#v", frame) + } + if strings.Contains(frame.Message, "native-secret-response") { + t.Fatalf("internal error leaked through wire frame: %#v", frame) + } + }) + } +} + +func TestNavigatePiSessionTreeRequiresManager(t *testing.T) { + var manager *Manager + if _, err := manager.NavigatePiSessionTree(context.Background(), "session", PiTreeNavigateInput{}); err == nil { + t.Fatal("expected nil manager rejection") + } +} + +func stringPointer(value string) *string { + return &value +} diff --git a/service/websession/pi_trust.go b/service/websession/pi_trust.go new file mode 100644 index 00000000..78b46332 --- /dev/null +++ b/service/websession/pi_trust.go @@ -0,0 +1,183 @@ +package websession + +import ( + "context" + "errors" + "fmt" + "os/exec" + "strings" + + "code-kanban/service" +) + +type piRuntimeTerminator struct { + projectID string + terminate func() +} + +func (m *Manager) GetProjectPiTrust( + ctx context.Context, + projectID string, +) (service.ProjectAgentTrustStatus, error) { + if m == nil || m.agentTrustSvc == nil { + return service.ProjectAgentTrustStatus{}, errors.New("project agent trust service is not configured") + } + return m.agentTrustSvc.GetStatus(ctx, projectID, service.ProjectAgentPi) +} + +func (m *Manager) TrustProjectForPi( + ctx context.Context, + projectID string, +) (service.ProjectAgentTrustStatus, error) { + if m == nil || m.agentTrustSvc == nil { + return service.ProjectAgentTrustStatus{}, errors.New("project agent trust service is not configured") + } + return m.agentTrustSvc.Trust(ctx, projectID, service.ProjectAgentPi) +} + +func (m *Manager) RevokeProjectPiTrust( + ctx context.Context, + projectID string, +) (service.ProjectAgentTrustStatus, error) { + if m == nil || m.agentTrustSvc == nil { + return service.ProjectAgentTrustStatus{}, errors.New("project agent trust service is not configured") + } + status, err := m.agentTrustSvc.Revoke(ctx, projectID, service.ProjectAgentPi) + if err != nil { + return service.ProjectAgentTrustStatus{}, err + } + m.StopProjectPiRuntimes(projectID) + return status, nil +} + +func (m *Manager) EnsureProjectPiTrust(ctx context.Context, projectID, cwd string) error { + if m == nil || m.agentTrustSvc == nil { + return errors.New("project agent trust service is not configured") + } + return m.agentTrustSvc.EnsureTrustedPath(ctx, projectID, service.ProjectAgentPi, cwd) +} + +// buildTrustedPiRPCCommand is the only launch path for persistent Pi RPC +// processes. The no-session capability probe intentionally does not use it. +func (m *Manager) buildTrustedPiRPCCommand( + ctx context.Context, + projectID string, + cwd string, + args ...string, +) (*exec.Cmd, error) { + if err := m.EnsureProjectPiTrust(ctx, projectID, cwd); err != nil { + return nil, err + } + for _, arg := range args { + normalized := strings.ToLower(strings.TrimSpace(arg)) + if normalized == "--approve" || normalized == "--no-approve" || + strings.HasPrefix(normalized, "--approve=") || strings.HasPrefix(normalized, "--no-approve=") { + return nil, fmt.Errorf("Pi approval flags are managed by CodeKanban") + } + if normalized == "--mode" || strings.HasPrefix(normalized, "--mode=") { + return nil, fmt.Errorf("Pi mode is managed by CodeKanban") + } + if normalized == "--extension" || normalized == "-e" || strings.HasPrefix(normalized, "--extension=") { + return nil, fmt.Errorf("Pi extensions are managed by CodeKanban") + } + } + bridgePath, err := m.materializePiBridge() + if err != nil { + return nil, err + } + launchArgs := []string{"--mode", "rpc", "--approve", "--extension", bridgePath} + launchArgs = append(launchArgs, args...) + // The caller context governs trust lookup, not the lifetime of the reusable + // process. Runtime shutdown is owned by the Pi registry. + cmd, err := buildPiCommand(context.Background(), m.cfg.PiPath, launchArgs...) + if err != nil { + return nil, err + } + cmd.Dir = cwd + return cmd, nil +} + +func (m *Manager) registerPiRuntimeTerminator( + sessionID string, + projectID string, + terminate func(), +) { + if m == nil || strings.TrimSpace(sessionID) == "" || terminate == nil { + return + } + m.piRuntimeMu.Lock() + m.piRuntimeTerminators[strings.TrimSpace(sessionID)] = piRuntimeTerminator{ + projectID: strings.TrimSpace(projectID), + terminate: terminate, + } + m.piRuntimeMu.Unlock() +} + +func (m *Manager) unregisterPiRuntimeTerminator(sessionID string) { + if m == nil { + return + } + m.piRuntimeMu.Lock() + delete(m.piRuntimeTerminators, strings.TrimSpace(sessionID)) + m.piRuntimeMu.Unlock() +} + +// StopSessionPiRuntime terminates a registered Pi process, including an idle one. +func (m *Manager) StopSessionPiRuntime(sessionID string) { + if m == nil { + return + } + m.piRuntimeMu.Lock() + runtime, ok := m.piRuntimeTerminators[strings.TrimSpace(sessionID)] + if ok { + delete(m.piRuntimeTerminators, strings.TrimSpace(sessionID)) + delete(m.piRuntimes, strings.TrimSpace(sessionID)) + } + m.piRuntimeMu.Unlock() + if ok && runtime.terminate != nil { + runtime.terminate() + } +} + +// StopAllPiRuntimes terminates every persistent Pi process owned by the manager. +func (m *Manager) StopAllPiRuntimes() { + if m == nil { + return + } + m.piRuntimeMu.Lock() + terminators := make([]func(), 0, len(m.piRuntimeTerminators)) + for sessionID, runtime := range m.piRuntimeTerminators { + delete(m.piRuntimeTerminators, sessionID) + delete(m.piRuntimes, sessionID) + if runtime.terminate != nil { + terminators = append(terminators, runtime.terminate) + } + } + m.piRuntimeMu.Unlock() + for _, terminate := range terminators { + terminate() + } +} + +// StopProjectPiRuntimes terminates registered Pi processes for one project. +func (m *Manager) StopProjectPiRuntimes(projectID string) { + if m == nil { + return + } + projectID = strings.TrimSpace(projectID) + terminators := make([]func(), 0) + m.piRuntimeMu.Lock() + for sessionID, runtime := range m.piRuntimeTerminators { + if runtime.projectID != projectID { + continue + } + delete(m.piRuntimeTerminators, sessionID) + if runtime.terminate != nil { + terminators = append(terminators, runtime.terminate) + } + } + m.piRuntimeMu.Unlock() + for _, terminate := range terminators { + terminate() + } +} diff --git a/service/websession/pi_trust_test.go b/service/websession/pi_trust_test.go new file mode 100644 index 00000000..3b9f45a3 --- /dev/null +++ b/service/websession/pi_trust_test.go @@ -0,0 +1,139 @@ +package websession + +import ( + "context" + "errors" + "os" + "strings" + "sync/atomic" + "testing" + + "code-kanban/model" + "code-kanban/model/tables" + "code-kanban/service" + + "go.uber.org/zap" +) + +func TestBuildTrustedPiRPCCommandRequiresTrustAndAddsApprove(t *testing.T) { + cleanup := initTestDB(t) + defer cleanup() + project := seedProject(t) + manager, err := NewManager(Config{ + DataDir: t.TempDir(), + PiPath: `"` + os.Args[0] + `"`, + }, zap.NewNop()) + if err != nil { + t.Fatalf("NewManager returned error: %v", err) + } + + _, err = manager.buildTrustedPiRPCCommand( + context.Background(), + project.ID, + project.Path, + "--name", + "test", + ) + if !errors.Is(err, service.ErrProjectAgentTrustRequired) { + t.Fatalf("untrusted launch error = %v, want trust required", err) + } + if _, err := manager.TrustProjectForPi(context.Background(), project.ID); err != nil { + t.Fatalf("TrustProjectForPi returned error: %v", err) + } + cmd, err := manager.buildTrustedPiRPCCommand( + context.Background(), + project.ID, + project.Path, + "--name", + "test", + ) + if err != nil { + t.Fatalf("trusted launch returned error: %v", err) + } + if cmd.Dir != project.Path { + t.Fatalf("command dir = %q, want %q", cmd.Dir, project.Path) + } + args := strings.Join(cmd.Args[1:], " ") + if !strings.Contains(args, "--mode rpc") || !strings.Contains(args, "--approve") || !strings.Contains(args, "--extension") { + t.Fatalf("trusted command args = %#v", cmd.Args) + } + if strings.Contains(args, "--no-approve") { + t.Fatalf("trusted command contains --no-approve: %#v", cmd.Args) + } + for _, managedFlag := range []string{"--no-approve", "--approve=false", "--mode=interactive", "--extension=untrusted.ts", "-e"} { + if _, err := manager.buildTrustedPiRPCCommand( + context.Background(), + project.ID, + project.Path, + managedFlag, + ); err == nil { + t.Fatalf("expected caller-supplied managed flag %q to be rejected", managedFlag) + } + } +} + +func TestRevokeProjectPiTrustStopsOnlyProjectRuntimes(t *testing.T) { + cleanup := initTestDB(t) + defer cleanup() + project := seedProject(t) + otherProject := seedProject(t) + manager, err := NewManager(Config{DataDir: t.TempDir()}, zap.NewNop()) + if err != nil { + t.Fatalf("NewManager returned error: %v", err) + } + if _, err := manager.TrustProjectForPi(context.Background(), project.ID); err != nil { + t.Fatalf("TrustProjectForPi returned error: %v", err) + } + + var stopped atomic.Int32 + var otherStopped atomic.Int32 + manager.registerPiRuntimeTerminator("session-1", project.ID, func() { stopped.Add(1) }) + manager.registerPiRuntimeTerminator("session-2", project.ID, func() { stopped.Add(1) }) + manager.registerPiRuntimeTerminator("session-other", otherProject.ID, func() { otherStopped.Add(1) }) + status, err := manager.RevokeProjectPiTrust(context.Background(), project.ID) + if err != nil { + t.Fatalf("RevokeProjectPiTrust returned error: %v", err) + } + if status.Trusted || status.RevokedAt == nil { + t.Fatalf("unexpected revoked status: %#v", status) + } + if stopped.Load() != 2 || otherStopped.Load() != 0 { + t.Fatalf("stopped=%d otherStopped=%d", stopped.Load(), otherStopped.Load()) + } + manager.piRuntimeMu.Lock() + _, otherExists := manager.piRuntimeTerminators["session-other"] + remaining := len(manager.piRuntimeTerminators) + manager.piRuntimeMu.Unlock() + if !otherExists || remaining != 1 { + t.Fatalf("remaining terminators = %d, other exists=%v", remaining, otherExists) + } +} + +func TestEnsureProjectPiTrustAcceptsManagedWorktree(t *testing.T) { + cleanup := initTestDB(t) + defer cleanup() + project := seedProject(t) + worktree := &tables.WorktreeTable{ + ProjectID: project.ID, + BranchName: "feature/pi-trust", + Path: t.TempDir(), + } + worktree.Init() + if err := model.GetDB().Create(worktree).Error; err != nil { + t.Fatalf("seed worktree: %v", err) + } + manager, err := NewManager(Config{DataDir: t.TempDir()}, zap.NewNop()) + if err != nil { + t.Fatalf("NewManager returned error: %v", err) + } + if _, err := manager.TrustProjectForPi(context.Background(), project.ID); err != nil { + t.Fatalf("TrustProjectForPi returned error: %v", err) + } + if err := manager.EnsureProjectPiTrust(context.Background(), project.ID, worktree.Path); err != nil { + t.Fatalf("managed worktree rejected: %v", err) + } + if err := manager.EnsureProjectPiTrust(context.Background(), project.ID, t.TempDir()); !errors.Is(err, service.ErrProjectAgentPathNotAllowed) { + t.Fatalf("unmanaged cwd result = %v, want path not allowed", err) + } + +} diff --git a/service/websession/types.go b/service/websession/types.go index aac8e6d9..f7ed4974 100644 --- a/service/websession/types.go +++ b/service/websession/types.go @@ -7,6 +7,7 @@ type Agent string const ( AgentClaude Agent = "claude" AgentCodex Agent = "codex" + AgentPi Agent = "pi" ) type ClaudeRuntime string @@ -21,6 +22,7 @@ type SessionBackend string const ( SessionBackendLegacyExec SessionBackend = "legacy_exec" SessionBackendCodexAppServer SessionBackend = "codex_app_server" + SessionBackendPiRPC SessionBackend = "pi_rpc" ) type WorkflowMode string @@ -226,6 +228,8 @@ type SessionSummary struct { AutoRetryDispatchPendingOnFailure bool `json:"autoRetryDispatchPendingOnFailure"` Cwd string `json:"cwd"` NativeSessionID *string `json:"nativeSessionId,omitempty"` + NativeLeafID *string `json:"nativeLeafId,omitempty"` + SourceRevision *string `json:"sourceRevision,omitempty"` CyberPolicyFlagged bool `json:"cyberPolicyFlagged"` HasScheduledPlanExecution bool `json:"hasScheduledPlanExecution,omitempty"` Status Status `json:"status"` @@ -392,6 +396,7 @@ type PendingInput struct { AttachmentIDs []string `json:"attachmentIds"` ReadyAt *time.Time `json:"readyAt,omitempty"` Paused bool `json:"paused,omitempty"` + NativeQueued bool `json:"nativeQueued,omitempty"` CreatedAt time.Time `json:"createdAt"` } @@ -507,6 +512,8 @@ type ImportResult struct { } type ImportSourceSummary struct { + Agent Agent `json:"agent"` + Importable bool `json:"importable"` AISessionID string `json:"aiSessionId"` SessionID string `json:"sessionId"` Model string `json:"model,omitempty"` @@ -521,8 +528,9 @@ type ImportSourceSummary struct { } type ImportSourceList struct { - Items []ImportSourceSummary `json:"items"` - ScanPhase string `json:"scanPhase,omitempty"` + Items []ImportSourceSummary `json:"items"` + ScanPhase string `json:"scanPhase,omitempty"` + BeforeCursor string `json:"beforeCursor,omitempty"` } type Event struct { diff --git a/service/websession/wire.go b/service/websession/wire.go index 21e5bec7..da98b381 100644 --- a/service/websession/wire.go +++ b/service/websession/wire.go @@ -73,6 +73,8 @@ type wireSess struct { Title string `json:"ttl"` Cwd string `json:"cwd"` NativeSessionID *string `json:"nsid,omitempty"` + NativeLeafID *string `json:"nlid,omitempty"` + SourceRevision *string `json:"srev,omitempty"` CyberPolicyFlagged bool `json:"cpf,omitempty"` HasScheduledPlanExecution bool `json:"spe,omitempty"` Status string `json:"st"` @@ -218,6 +220,7 @@ type wirePendingInput struct { AttachmentIDs []string `json:"atts,omitempty"` ReadyAt *int64 `json:"ra,omitempty"` Paused bool `json:"ps,omitempty"` + NativeQueued bool `json:"nq,omitempty"` CreatedAt int64 `json:"ca"` } @@ -468,6 +471,8 @@ func mapWireSession(session SessionSummary) *wireSess { Title: session.Title, Cwd: session.Cwd, NativeSessionID: session.NativeSessionID, + NativeLeafID: session.NativeLeafID, + SourceRevision: session.SourceRevision, CyberPolicyFlagged: session.CyberPolicyFlagged, HasScheduledPlanExecution: session.HasScheduledPlanExecution, Status: string(session.Status), @@ -557,6 +562,7 @@ func mapWirePendingInputs(items []PendingInput) []wirePendingInput { AttachmentIDs: append([]string(nil), item.AttachmentIDs...), ReadyAt: readyAt, Paused: item.Paused, + NativeQueued: item.NativeQueued, CreatedAt: item.CreatedAt.UnixMilli(), }) } diff --git a/service/websession/wire_pending_test.go b/service/websession/wire_pending_test.go new file mode 100644 index 00000000..41017c42 --- /dev/null +++ b/service/websession/wire_pending_test.go @@ -0,0 +1,36 @@ +package websession + +import ( + "encoding/json" + "testing" + "time" +) + +func TestPendingWireMarksPiNativeQueueReadOnly(t *testing.T) { + createdAt := time.Date(2026, 8, 1, 12, 0, 0, 0, time.UTC) + encoded, err := json.Marshal(newPendingFrame("session_1", []PendingInput{ + { + ID: "pi-native-1", + Mode: PendingInputModeQueue, + Text: "Accepted by Pi", + NativeQueued: true, + CreatedAt: createdAt, + }, + })) + if err != nil { + t.Fatalf("marshal pending frame: %v", err) + } + + var payload map[string]any + if err := json.Unmarshal(encoded, &payload); err != nil { + t.Fatalf("decode pending frame: %v", err) + } + items, ok := payload["pi"].([]any) + if !ok || len(items) != 1 { + t.Fatalf("expected one pending item, got %s", encoded) + } + item, _ := items[0].(map[string]any) + if item["nq"] != true || item["id"] != "pi-native-1" || item["m"] != string(PendingInputModeQueue) { + t.Fatalf("expected compact native queue marker, got %#v", item) + } +} diff --git a/tools/logwatcher/main.go b/tools/logwatcher/main.go index a24ae977..9fbb2833 100644 --- a/tools/logwatcher/main.go +++ b/tools/logwatcher/main.go @@ -78,7 +78,7 @@ func main() { case "watch": // watch [mode] - directly start watching with specified type and path if len(parts) < 3 { - fmt.Println("Usage: watch [mode]") + fmt.Println("Usage: watch [mode]") fmt.Println(" mode: both (default), ctime, mtime") fmt.Println("Example: watch claude D:\\codes\\2025\\aicode-kanban") fmt.Println("Example: watch claude D:\\codes\\2025\\aicode-kanban mtime") @@ -91,8 +91,10 @@ func main() { aType = types.AssistantTypeCodex case "claude", "claudecode": aType = types.AssistantTypeClaudeCode + case "pi": + aType = types.AssistantTypePi default: - fmt.Printf("Unknown type: %s (use 'codex' or 'claude')\n", parts[1]) + fmt.Printf("Unknown type: %s (use 'codex', 'claude', or 'pi')\n", parts[1]) continue } @@ -126,7 +128,7 @@ func main() { case "session": // session - Watch a specific session by ID if len(parts) < 4 { - fmt.Println("Usage: session ") + fmt.Println("Usage: session ") fmt.Println("Example: session claude D:\\codes\\2025\\aicode-kanban 8a874861-cbd9-4c66-964d-0b9311c68598") continue } @@ -137,8 +139,10 @@ func main() { aType = types.AssistantTypeCodex case "claude", "claudecode": aType = types.AssistantTypeClaudeCode + case "pi": + aType = types.AssistantTypePi default: - fmt.Printf("Unknown type: %s (use 'codex' or 'claude')\n", parts[1]) + fmt.Printf("Unknown type: %s (use 'codex', 'claude', or 'pi')\n", parts[1]) continue } @@ -314,7 +318,7 @@ func main() { func printHelp() { fmt.Println(`Available commands: watch [mode] - Start watching with specified type and working directory - type: codex, claude + type: codex, claude, pi mode: both (default), ctime, mtime Example: watch claude D:\codes\2025\aicode-kanban Example: watch claude D:\codes\2025\aicode-kanban mtime @@ -552,7 +556,8 @@ func startWatcherBySessionID(aType types.AssistantType, workingDir string, sessi // Find the session file by ID var filePath string - if aType == types.AssistantTypeClaudeCode { + switch aType { + case types.AssistantTypeClaudeCode: searcher, err := log_watcher.NewClaudeCodeFileSearcher(workingDir) if err != nil { fmt.Printf("Failed to create searcher: %v\n", err) @@ -563,8 +568,19 @@ func startWatcherBySessionID(aType types.AssistantType, workingDir string, sessi fmt.Printf("Error finding session file: %v\n", err) return nil, nil, nil } - } else { - fmt.Println("Session ID lookup is only supported for Claude Code currently.") + case types.AssistantTypePi: + searcher, err := log_watcher.NewPiFileSearcherWithWorkingDir(workingDir) + if err != nil { + fmt.Printf("Failed to create searcher: %v\n", err) + return nil, nil, nil + } + filePath, err = searcher.FindBySessionID(context.Background(), sessionID) + if err != nil { + fmt.Printf("Error finding session file: %v\n", err) + return nil, nil, nil + } + default: + fmt.Println("Session ID lookup is supported for Claude Code and Pi.") return nil, nil, nil } diff --git a/ui/src/api/__tests__/webSession.test.ts b/ui/src/api/__tests__/webSession.test.ts index 389ea7a7..269436f8 100644 --- a/ui/src/api/__tests__/webSession.test.ts +++ b/ui/src/api/__tests__/webSession.test.ts @@ -1,18 +1,22 @@ import { beforeEach, describe, expect, it, vi } from 'vitest'; -const { postMethodMock, postSendMock, postAbortMock, fetchMock } = vi.hoisted(() => { - const postSendMock = vi.fn(); - const postAbortMock = vi.fn(); - return { - postMethodMock: vi.fn(() => ({ - send: postSendMock, - abort: postAbortMock, - })), - postSendMock, - postAbortMock, - fetchMock: vi.fn(), - }; -}); +const { getMethodMock, getSendMock, postMethodMock, postSendMock, postAbortMock, fetchMock } = + vi.hoisted(() => { + const getSendMock = vi.fn(); + const postSendMock = vi.fn(); + const postAbortMock = vi.fn(); + return { + getMethodMock: vi.fn(() => ({ send: getSendMock })), + getSendMock, + postMethodMock: vi.fn(() => ({ + send: postSendMock, + abort: postAbortMock, + })), + postSendMock, + postAbortMock, + fetchMock: vi.fn(), + }; + }); vi.mock('@/api', () => ({ urlBase: '', @@ -32,7 +36,7 @@ vi.mock('@/api', () => ({ vi.mock('@/api/http', () => ({ http: { - Get: vi.fn(), + Get: getMethodMock, Post: postMethodMock, Patch: vi.fn(), Delete: vi.fn(), @@ -230,6 +234,90 @@ describe('webSessionApi.searchConversation', () => { }); }); +describe('webSessionApi Pi tree', () => { + const tree = { + sessionId: 'session-1', + leafId: 'leaf-1', + revision: 'rev-1', + nodes: [ + { + id: 'leaf-1', + type: 'message', + role: 'user', + active: true, + children: [], + }, + ], + }; + + beforeEach(() => { + getMethodMock.mockClear(); + getSendMock.mockReset(); + postMethodMock.mockClear(); + postSendMock.mockReset(); + }); + + it('loads the encoded project session tree', async () => { + getSendMock.mockResolvedValueOnce({ item: tree }); + + await expect(webSessionApi.tree('project/1', 'session 2')).resolves.toEqual(tree); + expect(getMethodMock).toHaveBeenCalledWith( + '/projects/project%2F1/web-sessions/session%202/tree' + ); + }); + + it('navigates with a revision and optional summary request', async () => { + postSendMock.mockResolvedValueOnce({ item: { tree, editorText: 'rewrite this' } }); + + await expect( + webSessionApi.navigateTree('project-1', 'session-1', { + targetId: 'leaf-1', + revision: 'rev-1', + summarize: true, + }) + ).resolves.toMatchObject({ tree, editorText: 'rewrite this' }); + expect(postMethodMock).toHaveBeenCalledWith( + '/projects/project-1/web-sessions/session-1/tree/navigate', + { targetId: 'leaf-1', revision: 'rev-1', summarize: true } + ); + }); + + it('forks and clones into a returned target session', async () => { + postSendMock + .mockResolvedValueOnce({ + item: { + session: { id: 'forked' }, + tree: { ...tree, sessionId: 'forked' }, + editorText: 'u', + }, + }) + .mockResolvedValueOnce({ + item: { session: { id: 'cloned' }, tree: { ...tree, sessionId: 'cloned' } }, + }); + + await expect( + webSessionApi.forkTree('project-1', 'session-1', { + targetId: 'leaf-1', + revision: 'rev-1', + }) + ).resolves.toMatchObject({ session: { id: 'forked' }, editorText: 'u' }); + expect(postMethodMock).toHaveBeenNthCalledWith( + 1, + '/projects/project-1/web-sessions/session-1/tree/fork', + { targetId: 'leaf-1', revision: 'rev-1' } + ); + + await expect( + webSessionApi.cloneTree('project-1', 'session-1', { revision: 'rev-1' }) + ).resolves.toMatchObject({ session: { id: 'cloned' } }); + expect(postMethodMock).toHaveBeenNthCalledWith( + 2, + '/projects/project-1/web-sessions/session-1/tree/clone', + { revision: 'rev-1' } + ); + }); +}); + describe('webSessionApi local files', () => { beforeEach(() => { postMethodMock.mockClear(); diff --git a/ui/src/api/project.ts b/ui/src/api/project.ts index ced56012..f52b8207 100644 --- a/ui/src/api/project.ts +++ b/ui/src/api/project.ts @@ -2,6 +2,7 @@ import type { CodexSkillSummary, GitCapabilityResult, Project, + ProjectAgentTrustStatus, Worktree, } from '@/types/models'; import { http } from './http'; @@ -72,7 +73,8 @@ export const projectApi = { }, async markAccess(id: string): Promise { - const body = (await http.Post>(`/projects/${id}/access`, {}).send()) ?? {}; + const body = + (await http.Post>(`/projects/${id}/access`, {}).send()) ?? {}; if (!body.item) { throw new Error('failed to record project access'); } @@ -88,6 +90,41 @@ export const projectApi = { return body.item; }, + async getPiTrust(id: string): Promise { + const body = + (await http + .Get>(`/projects/${id}/agent-trust/pi`, { + cacheFor: 0, + }) + .send(true)) ?? {}; + if (!body.item) { + throw new Error('failed to load Pi project access'); + } + return body.item; + }, + + async trustForPi(id: string): Promise { + const body = + (await http + .Post>(`/projects/${id}/agent-trust/pi`, {}) + .send()) ?? {}; + if (!body.item) { + throw new Error('failed to authorize Pi project access'); + } + return body.item; + }, + + async revokePiTrust(id: string): Promise { + const body = + (await http + .Delete>(`/projects/${id}/agent-trust/pi`) + .send()) ?? {}; + if (!body.item) { + throw new Error('failed to revoke Pi project access'); + } + return body.item; + }, + async gitCapabilities(id: string): Promise { const body = (await http diff --git a/ui/src/api/webSession.ts b/ui/src/api/webSession.ts index 65f8ba37..f3c32f99 100644 --- a/ui/src/api/webSession.ts +++ b/ui/src/api/webSession.ts @@ -1,7 +1,8 @@ import type { + WebSessionAgent, WebSessionAttachment, - WebSessionCodexRuntimeConfig, WebSessionReasoningEffort, + WebSessionRuntimeConfig, WebSessionSummary, } from '@/types/models'; import { ApiError, urlBase } from '@/api'; @@ -118,6 +119,36 @@ export type WebSessionHistoryWindow = { total: number; }; +export type WebSessionPiTreeNode = { + id: string; + parentId?: string; + type: string; + timestamp?: string; + role?: string; + label?: string; + preview?: string; + active: boolean; + children: string[]; +}; + +export type WebSessionPiTree = { + sessionId: string; + leafId?: string; + revision: string; + nodes: WebSessionPiTreeNode[]; +}; + +export type WebSessionPiTreeNavigateResult = { + tree: WebSessionPiTree; + editorText?: string; +}; + +export type WebSessionPiTreeMutationResult = { + session: WebSessionSummary; + tree: WebSessionPiTree; + editorText?: string; +}; + export type WebSessionPendingInputRecord = { id?: string; mode?: 'redirect' | 'queue' | string; @@ -125,6 +156,7 @@ export type WebSessionPendingInputRecord = { attachmentIds?: string[]; readyAt?: string | number | null; paused?: boolean; + nativeQueued?: boolean; createdAt?: string | number | null; }; @@ -204,10 +236,10 @@ export type WebSessionImportResult = Omit< }; export const webSessionApi = { - async runtimeConfig(): Promise { - const config = extractItem( + async runtimeConfig(): Promise { + const config = extractItem( await http - .Get>('/web-sessions/runtime-config', { + .Get>('/web-sessions/runtime-config', { cacheFor: 0, }) .send(true) @@ -235,7 +267,7 @@ export const webSessionApi = { projectId: string, data: { worktreeId?: string; - agent: 'claude' | 'codex'; + agent: WebSessionAgent; claudeRuntime?: 'claude' | 'ccr'; model?: string; reasoningEffort?: WebSessionReasoningEffort; @@ -298,6 +330,7 @@ export const webSessionApi = { async importSession( projectId: string, data: { + agent?: 'codex' | 'pi'; aiSessionId?: string; sessionId?: string; mode?: 'fast' | 'deep'; @@ -306,6 +339,7 @@ export const webSessionApi = { const body = (await http .Post>(`/projects/${projectId}/web-sessions/import`, { + agent: data.agent ?? 'codex', aiSessionId: data.aiSessionId ?? '', sessionId: data.sessionId ?? '', ...(data.mode ? { mode: data.mode } : {}), @@ -382,6 +416,70 @@ export const webSessionApi = { return body.item; }, + async tree(projectId: string, sessionId: string): Promise { + const body = + (await http + .Get< + ItemResponse + >(`/projects/${encodeURIComponent(projectId)}/web-sessions/${encodeURIComponent(sessionId)}/tree`) + .send(true)) ?? {}; + if (!body.item) { + throw new Error('failed to load Pi session tree'); + } + return body.item; + }, + + async navigateTree( + projectId: string, + sessionId: string, + data: { targetId: string; revision: string; summarize?: boolean } + ): Promise { + const body = + (await http + .Post< + ItemResponse + >(`/projects/${encodeURIComponent(projectId)}/web-sessions/${encodeURIComponent(sessionId)}/tree/navigate`, data) + .send()) ?? {}; + if (!body.item?.tree) { + throw new Error('failed to navigate Pi session tree'); + } + return body.item; + }, + + async forkTree( + projectId: string, + sessionId: string, + data: { targetId: string; revision: string } + ): Promise { + const body = + (await http + .Post< + ItemResponse + >(`/projects/${encodeURIComponent(projectId)}/web-sessions/${encodeURIComponent(sessionId)}/tree/fork`, data) + .send()) ?? {}; + if (!body.item?.session?.id) { + throw new Error('failed to fork Pi session tree'); + } + return body.item; + }, + + async cloneTree( + projectId: string, + sessionId: string, + data: { revision: string } + ): Promise { + const body = + (await http + .Post< + ItemResponse + >(`/projects/${encodeURIComponent(projectId)}/web-sessions/${encodeURIComponent(sessionId)}/tree/clone`, data) + .send()) ?? {}; + if (!body.item?.session?.id) { + throw new Error('failed to clone Pi session tree'); + } + return body.item; + }, + async history( projectId: string, sessionId: string, diff --git a/ui/src/components/common/ConversationViewer.vue b/ui/src/components/common/ConversationViewer.vue index 00a529c0..7d793ad8 100644 --- a/ui/src/components/common/ConversationViewer.vue +++ b/ui/src/components/common/ConversationViewer.vue @@ -171,26 +171,11 @@
- + - {{ sessionInfo.type === 'claude_code' ? 'Claude Code' : 'Codex' }} + {{ sessionAssistantLabel }} {{ sessionInfo.sessionId }} @@ -282,11 +267,12 @@ import { computed, h, nextTick, onBeforeUnmount, ref, watch } from 'vue'; import { useTimeAgo } from '@vueuse/core'; import { useDialog, useMessage } from 'naive-ui'; -import { CopyOutline, ImageOutline, LogoGithub, RefreshOutline } from '@vicons/ionicons5'; +import { CopyOutline, ImageOutline, RefreshOutline } from '@vicons/ionicons5'; import { useLocale } from '@/composables/useLocale'; import { useAppClipboard } from '@/composables/useAppClipboard'; import { useConversationVirtualizer } from '@/composables/useConversationVirtualizer'; import { renderMarkdown } from '@/utils/markdown'; +import { getAssistantIconByType } from '@/utils/assistantIcon'; import { getClickedMarkdownCodeCopyText, getClickedMarkdownLink, @@ -402,6 +388,22 @@ const emit = defineEmits<{ const { t } = useLocale(); const dialog = useDialog(); const message = useMessage(); +const sessionAssistantType = computed(() => { + if (props.sessionInfo?.type === 'claude_code') return 'claude-code'; + if (props.sessionInfo?.type === 'pi') return 'pi'; + return 'codex'; +}); +const sessionAssistantIcon = computed(() => getAssistantIconByType(sessionAssistantType.value)); +const sessionAssistantLabel = computed(() => { + if (props.sessionInfo?.type === 'claude_code') return 'Claude Code'; + if (props.sessionInfo?.type === 'pi') return 'Pi'; + return 'Codex'; +}); +const sessionAssistantTagType = computed(() => { + if (props.sessionInfo?.type === 'claude_code') return 'info'; + if (props.sessionInfo?.type === 'pi') return 'warning'; + return 'success'; +}); const { copyText } = useAppClipboard(); const showUserOnly = ref(false); @@ -1481,6 +1483,14 @@ onBeforeUnmount(() => { gap: 8px; } +.session-assistant-icon { + display: inline-flex; + width: 12px; + height: 12px; + align-items: center; + justify-content: center; +} + .session-id-code { font-size: 12px; font-family: monospace; diff --git a/ui/src/components/kanban/TaskDetailDrawer.vue b/ui/src/components/kanban/TaskDetailDrawer.vue index aea67ea0..4d99fb14 100644 --- a/ui/src/components/kanban/TaskDetailDrawer.vue +++ b/ui/src/components/kanban/TaskDetailDrawer.vue @@ -80,9 +80,21 @@
- {{ session.type === 'claude_code' ? 'Claude' : 'Codex' }} + {{ + session.type === 'claude_code' + ? 'Claude' + : session.type === 'pi' + ? 'Pi' + : 'Codex' + }} {{ formatDate(session.sessionStartedAt) @@ -240,8 +252,23 @@ {{ session.title || t('terminal.untitledSession') }}
- - {{ session.type === 'claude_code' ? 'Claude' : 'Codex' }} + + {{ + session.type === 'claude_code' + ? 'Claude' + : session.type === 'pi' + ? 'Pi' + : 'Codex' + }} {{ session.model || '-' }} {{ formatDate(session.sessionStartedAt) }} @@ -368,7 +395,11 @@ const { loadEarlier: loadEarlierConversationWindow, reset: resetConversationWindow, } = useAiConversationWindow( - options => aiSessionApi.conversationWindowByID(currentConversationSession.value?.aiSessionDbId || '', options), + options => + aiSessionApi.conversationWindowByID( + currentConversationSession.value?.aiSessionDbId || '', + options + ), null ); @@ -376,7 +407,7 @@ const currentSessionInfo = computed(() => { if (!currentConversationSession.value) return null; return { sessionId: currentConversationSession.value.sessionId, - type: currentConversationSession.value.type as 'claude_code' | 'codex', + type: currentConversationSession.value.type as 'claude_code' | 'codex' | 'pi', }; }); @@ -553,11 +584,12 @@ async function openLinkSessionModal() { .send(); if (response?.item) { - // 合并 Claude 和 Codex sessions,排除已关联的 + // 合并各 Agent sessions,排除已关联的 const linkedIds = new Set(linkedAiSessions.value.map(s => s.aiSessionDbId)); const allSessions = [ ...(response.item.claudeSessions || []), ...(response.item.codexSessions || []), + ...(response.item.piSessions || []), ].filter(s => !linkedIds.has(s.id)); // 按时间排序,最新的在前 diff --git a/ui/src/components/project/PiProjectTrustDialog.vue b/ui/src/components/project/PiProjectTrustDialog.vue new file mode 100644 index 00000000..a7b35798 --- /dev/null +++ b/ui/src/components/project/PiProjectTrustDialog.vue @@ -0,0 +1,99 @@ + + + + + diff --git a/ui/src/components/project/ProjectEditDialog.vue b/ui/src/components/project/ProjectEditDialog.vue index 6c892431..1e5edd02 100644 --- a/ui/src/components/project/ProjectEditDialog.vue +++ b/ui/src/components/project/ProjectEditDialog.vue @@ -29,15 +29,72 @@ {{ t('project.hidePathHint') }} + +
+
+ + + {{ + piTrustStatus.trusted ? t('project.piAccessTrusted') : t('project.piAccessNotTrusted') + }} + +
+ +
+ + {{ t('project.piAccessNeedsRenewal') }} + + + {{ t('project.piTrustRevoke') }} + + + {{ t('project.piTrustConfirm') }} + +
+
+
+ + + diff --git a/ui/src/components/project/piProjectTrust.test.ts b/ui/src/components/project/piProjectTrust.test.ts new file mode 100644 index 00000000..5a370fd7 --- /dev/null +++ b/ui/src/components/project/piProjectTrust.test.ts @@ -0,0 +1,30 @@ +import { describe, expect, it } from 'vitest'; +import { isPiProjectTrusted, piProjectTrustNeedsRenewal } from './piProjectTrust'; + +const trustedStatus = { + projectId: 'project-1', + agent: 'pi' as const, + projectPath: 'D:/repo', + trustedPath: 'D:/repo', + trusted: true, +}; + +describe('Pi project trust helpers', () => { + it('accepts only an explicit Pi trust for the active project', () => { + expect(isPiProjectTrusted(trustedStatus, 'project-1')).toBe(true); + expect(isPiProjectTrusted(trustedStatus, 'project-2')).toBe(false); + expect(isPiProjectTrusted({ ...trustedStatus, trusted: false }, 'project-1')).toBe(false); + expect(isPiProjectTrusted({ ...trustedStatus, agent: 'codex' as never }, 'project-1')).toBe( + false + ); + expect(isPiProjectTrusted(null, 'project-1')).toBe(false); + }); + + it('detects a stale or revoked path that needs renewed confirmation', () => { + expect(piProjectTrustNeedsRenewal({ ...trustedStatus, trusted: false })).toBe(true); + expect(piProjectTrustNeedsRenewal({ ...trustedStatus, trusted: false, trustedPath: '' })).toBe( + false + ); + expect(piProjectTrustNeedsRenewal(trustedStatus)).toBe(false); + }); +}); diff --git a/ui/src/components/project/piProjectTrust.ts b/ui/src/components/project/piProjectTrust.ts new file mode 100644 index 00000000..0de2ba6b --- /dev/null +++ b/ui/src/components/project/piProjectTrust.ts @@ -0,0 +1,17 @@ +import type { ProjectAgentTrustStatus } from '@/types/models'; + +export function isPiProjectTrusted( + status: ProjectAgentTrustStatus | null | undefined, + projectId: string +) { + return Boolean( + status?.agent === 'pi' && + status.trusted === true && + status.projectId.trim() === projectId.trim() && + projectId.trim() + ); +} + +export function piProjectTrustNeedsRenewal(status: ProjectAgentTrustStatus | null | undefined) { + return Boolean(status && !status.trusted && status.trustedPath?.trim()); +} diff --git a/ui/src/components/terminal/AISessionHistoryDialog.vue b/ui/src/components/terminal/AISessionHistoryDialog.vue index 793fda7e..89c701b6 100644 --- a/ui/src/components/terminal/AISessionHistoryDialog.vue +++ b/ui/src/components/terminal/AISessionHistoryDialog.vue @@ -214,6 +214,75 @@
+ +
+ + +
+
@@ -332,6 +401,10 @@ import ConversationViewer, { } from '@/components/common/ConversationViewer.vue'; import DirectoryPickerDialog from '@/components/common/DirectoryPickerDialog.vue'; import { useProjectStore } from '@/stores/project'; +import { + isProjectAISessionScanning, + resolvePreferredAISessionType, +} from '@/components/terminal/aiSessionHistory'; type ScanPhase = 'recent' | 'extended' | 'complete'; @@ -351,10 +424,14 @@ interface AISessionSummary { interface ProjectAISessions { hasClaudeCode: boolean; hasCodex: boolean; + hasPi: boolean; claudeSessions: AISessionSummary[]; codexSessions: AISessionSummary[]; + piSessions: AISessionSummary[]; claudeScanPhase?: ScanPhase; codexScanPhase?: ScanPhase; + piScanPhase?: ScanPhase; + piBeforeCursor?: string; } interface ItemResponse { @@ -384,12 +461,14 @@ const currentProjectPath = computed(() => { }); const loading = ref(false); -const activeType = ref<'claude_code' | 'codex'>('claude_code'); +const activeType = ref<'claude_code' | 'codex' | 'pi'>('claude_code'); const expandedSessionId = ref(null); const claudeSessions = ref([]); const codexSessions = ref([]); +const piSessions = ref([]); const claudeScanPhase = ref('complete'); const codexScanPhase = ref('complete'); +const piScanPhase = ref('complete'); const searchQuery = ref(''); const customPath = ref(''); const showDirectoryPicker = ref(false); @@ -420,6 +499,17 @@ const filteredCodexSessions = computed(() => { ); }); +const filteredPiSessions = computed(() => { + if (!searchQuery.value.trim()) return piSessions.value; + const query = searchQuery.value.toLowerCase(); + return piSessions.value.filter( + s => + (s.title && s.title.toLowerCase().includes(query)) || + s.sessionId.toLowerCase().includes(query) || + s.model?.toLowerCase().includes(query) + ); +}); + const showConversationModal = ref(false); const currentSessionTitle = ref(''); const currentSession = ref(null); @@ -454,12 +544,16 @@ const currentSessionInfo = computed(() => { if (!currentSession.value) return null; return { sessionId: currentSession.value.sessionId, - type: currentSession.value.type as 'claude_code' | 'codex', + type: currentSession.value.type as 'claude_code' | 'codex' | 'pi', }; }); -const isScanning = computed( - () => claudeScanPhase.value !== 'complete' || codexScanPhase.value !== 'complete' +const isScanning = computed(() => + isProjectAISessionScanning({ + claudeScanPhase: claudeScanPhase.value, + codexScanPhase: codexScanPhase.value, + piScanPhase: piScanPhase.value, + }) ); const claudeCodeTabLabel = computed(() => { @@ -474,6 +568,12 @@ const codexTabLabel = computed(() => { return `Codex${count > 0 ? ` (${count}${scanIndicator})` : scanIndicator}`; }); +const piTabLabel = computed(() => { + const count = piSessions.value.length; + const scanIndicator = piScanPhase.value !== 'complete' ? ' ...' : ''; + return `Pi${count > 0 ? ` (${count}${scanIndicator})` : scanIndicator}`; +}); + watch(showModal, async show => { if (show && props.projectId) { await loadSessions(); @@ -506,23 +606,18 @@ async function loadSessions(isRefresh = false) { if (data) { claudeSessions.value = data.claudeSessions || []; codexSessions.value = data.codexSessions || []; + piSessions.value = data.piSessions || []; claudeScanPhase.value = data.claudeScanPhase || 'complete'; codexScanPhase.value = data.codexScanPhase || 'complete'; + piScanPhase.value = data.piScanPhase || 'complete'; // Auto-select tab with sessions (only on first load) if (!isRefresh) { - if (data.hasCodex && !data.hasClaudeCode) { - activeType.value = 'codex'; - } else { - activeType.value = 'claude_code'; - } + activeType.value = resolvePreferredAISessionType(data); } // Schedule refresh if still scanning - if ( - showModal.value && - (claudeScanPhase.value !== 'complete' || codexScanPhase.value !== 'complete') - ) { + if (showModal.value && isScanning.value) { if (refreshTimer) { clearTimeout(refreshTimer); } @@ -623,7 +718,6 @@ function handleDirectorySelected(path: string) { async function loadSessionsByPath(path?: string) { const targetPath = path || customPath.value.trim(); - console.log('[AISessionHistoryDialog] loadSessionsByPath:', targetPath); if (!targetPath) { message.warning(t('terminal.pleaseEnterPath')); return; @@ -639,31 +733,21 @@ async function loadSessionsByPath(path?: string) { }) .send(); - console.log('[AISessionHistoryDialog] response:', response); const data = response?.item; - console.log('[AISessionHistoryDialog] data:', data); - console.log('[AISessionHistoryDialog] claudeSessions:', data?.claudeSessions); - console.log('[AISessionHistoryDialog] codexSessions:', data?.codexSessions); if (data) { claudeSessions.value = data.claudeSessions || []; codexSessions.value = data.codexSessions || []; + piSessions.value = data.piSessions || []; claudeScanPhase.value = data.claudeScanPhase || 'complete'; codexScanPhase.value = data.codexScanPhase || 'complete'; + piScanPhase.value = data.piScanPhase || 'complete'; currentViewPath.value = targetPath; isCustomPath.value = true; - // Auto-select tab with sessions - if (data.hasCodex && !data.hasClaudeCode) { - activeType.value = 'codex'; - } else { - activeType.value = 'claude_code'; - } + activeType.value = resolvePreferredAISessionType(data); // Schedule refresh if still scanning - if ( - showModal.value && - (claudeScanPhase.value !== 'complete' || codexScanPhase.value !== 'complete') - ) { + if (showModal.value && isScanning.value) { if (refreshTimer) { clearTimeout(refreshTimer); } diff --git a/ui/src/components/terminal/__tests__/aiSessionHistory.test.ts b/ui/src/components/terminal/__tests__/aiSessionHistory.test.ts new file mode 100644 index 00000000..c4cae39c --- /dev/null +++ b/ui/src/components/terminal/__tests__/aiSessionHistory.test.ts @@ -0,0 +1,49 @@ +import { describe, expect, it } from 'vitest'; + +import { isProjectAISessionScanning, resolvePreferredAISessionType } from '../aiSessionHistory'; + +describe('AI session history provider selection', () => { + it('selects Pi when it is the only provider with history', () => { + expect( + resolvePreferredAISessionType({ + hasClaudeCode: false, + hasCodex: false, + hasPi: true, + }) + ).toBe('pi'); + }); + + it('preserves the existing Claude then Codex preference order', () => { + expect( + resolvePreferredAISessionType({ + hasClaudeCode: true, + hasCodex: true, + hasPi: true, + }) + ).toBe('claude_code'); + expect( + resolvePreferredAISessionType({ + hasClaudeCode: false, + hasCodex: true, + hasPi: true, + }) + ).toBe('codex'); + }); + + it('keeps polling while Pi history discovery is incomplete', () => { + expect( + isProjectAISessionScanning({ + claudeScanPhase: 'complete', + codexScanPhase: 'complete', + piScanPhase: 'extended', + }) + ).toBe(true); + expect( + isProjectAISessionScanning({ + claudeScanPhase: 'complete', + codexScanPhase: 'complete', + piScanPhase: 'complete', + }) + ).toBe(false); + }); +}); diff --git a/ui/src/components/terminal/aiSessionHistory.ts b/ui/src/components/terminal/aiSessionHistory.ts new file mode 100644 index 00000000..69d9d18a --- /dev/null +++ b/ui/src/components/terminal/aiSessionHistory.ts @@ -0,0 +1,22 @@ +import type { AISessionType, ProjectAISessions } from '@/types/models'; + +type AISessionAvailability = Pick; +type AISessionScanState = Pick< + ProjectAISessions, + 'claudeScanPhase' | 'codexScanPhase' | 'piScanPhase' +>; + +export function resolvePreferredAISessionType(data: AISessionAvailability): AISessionType { + if (data.hasClaudeCode) return 'claude_code'; + if (data.hasCodex) return 'codex'; + if (data.hasPi) return 'pi'; + return 'claude_code'; +} + +export function isProjectAISessionScanning(data: AISessionScanState): boolean { + return ( + (data.claudeScanPhase !== undefined && data.claudeScanPhase !== 'complete') || + (data.codexScanPhase !== undefined && data.codexScanPhase !== 'complete') || + (data.piScanPhase !== undefined && data.piScanPhase !== 'complete') + ); +} diff --git a/ui/src/components/web-session/WebSessionImportDialog.vue b/ui/src/components/web-session/WebSessionImportDialog.vue index c54c1e34..e61ffae7 100644 --- a/ui/src/components/web-session/WebSessionImportDialog.vue +++ b/ui/src/components/web-session/WebSessionImportDialog.vue @@ -51,7 +51,7 @@ @@ -245,6 +252,12 @@ import { useAiConversationWindow, } from '@/composables/useAiConversationWindow'; import type { WebSessionSummary } from '@/types/models'; +import { + countImportableWebSessionSources, + normalizeWebSessionImportSources, + type WebSessionImportSourceSummary, + type WebSessionImportSourceWire, +} from '@/components/web-session/webSessionImportSources'; import ConversationViewer, { type ConversationViewerNavState, type SessionInfo, @@ -252,22 +265,8 @@ import ConversationViewer, { type ScanPhase = 'recent' | 'extended' | 'complete'; -type ImportSourceSummary = { - aiSessionId: string; - sessionId: string; - model: string; - title: string | null; - sessionStartedAt: string; - lastMessageAt: string | null; - messageCount: number; - assistantMessageCount: number; - filePath: string; - duplicate: boolean; - existingSession?: WebSessionSummary | null; -}; - type ImportSourceList = { - items?: ImportSourceSummary[]; + items?: WebSessionImportSourceWire[]; scanPhase?: ScanPhase; }; @@ -283,7 +282,7 @@ const props = defineProps<{ const showModal = defineModel('show', { default: false }); const emit = defineEmits<{ - (e: 'import-session', sessionId: string): void; + (e: 'import-session', source: Pick): void; (e: 'open-existing-session', session: WebSessionSummary): void; }>(); @@ -293,11 +292,11 @@ const message = useMessage(); const loading = ref(false); const searchQuery = ref(''); const hideImported = ref(false); -const importSources = ref([]); +const importSources = ref([]); const scanPhase = ref('complete'); const showPreviewModal = ref(false); const previewingSourceId = ref(''); -const previewSource = ref(null); +const previewSource = ref(null); const conversationViewerRef = ref<{ goToPrevUserMessage: () => void; goToNextUserMessage: () => void; @@ -318,10 +317,13 @@ const { load: loadConversationWindow, loadEarlier: loadEarlierConversationWindow, reset: resetConversationWindow, -} = useAiConversationWindow( - options => aiSessionApi.conversationWindowBySessionID(previewSource.value?.sessionId || '', options), - null -); +} = useAiConversationWindow(options => { + const source = previewSource.value; + if (source?.aiSessionId) { + return aiSessionApi.conversationWindowByID(source.aiSessionId, options); + } + return aiSessionApi.conversationWindowBySessionID(source?.sessionId || '', options); +}, null); let refreshTimer: ReturnType | null = null; const filteredSources = computed(() => { @@ -346,7 +348,7 @@ const filteredSources = computed(() => { const duplicateCount = computed( () => importSources.value.filter(source => source.duplicate).length ); -const importableCount = computed(() => importSources.value.length - duplicateCount.value); +const importableCount = computed(() => countImportableWebSessionSources(importSources.value)); const previewSessionTitle = computed(() => { return ( @@ -360,7 +362,7 @@ const previewSessionInfo = computed(() => { } return { sessionId: previewSource.value.sessionId, - type: 'codex', + type: previewSource.value.agent, }; }); @@ -404,7 +406,7 @@ function updateConversationNavState(state: ConversationViewerNavState) { conversationNavState.value = state; } -function formatSessionTime(source: ImportSourceSummary) { +function formatSessionTime(source: WebSessionImportSourceSummary) { const raw = source.lastMessageAt || source.sessionStartedAt; if (!raw) { return '-'; @@ -418,20 +420,23 @@ function formatSessionTime(source: ImportSourceSummary) { function emitImportFromPreview() { const sessionId = previewSource.value?.sessionId || ''; - if (!sessionId || props.pendingSessionId) { + if (!sessionId || !previewSource.value?.importable || props.pendingSessionId) { return; } - emit('import-session', sessionId); + emit('import-session', { + agent: previewSource.value.agent, + sessionId, + }); } -function openExistingSession(source: ImportSourceSummary) { +function openExistingSession(source: WebSessionImportSourceSummary) { if (!source.existingSession) { return; } emit('open-existing-session', source.existingSession); } -async function openPreview(source: ImportSourceSummary) { +async function openPreview(source: WebSessionImportSourceSummary) { previewSource.value = source; previewingSourceId.value = source.sessionId; showPreviewModal.value = true; @@ -480,7 +485,7 @@ async function loadSources(isRefresh = false) { return; } - importSources.value = Array.isArray(data.items) ? data.items : []; + importSources.value = normalizeWebSessionImportSources(data.items); scanPhase.value = data.scanPhase || 'complete'; if (showModal.value && scanPhase.value !== 'complete') { @@ -489,7 +494,7 @@ async function loadSources(isRefresh = false) { }, 2000); } } catch (error) { - console.error('Failed to load codex import sources:', error); + console.error('Failed to load AI import sources:', error); if (!isRefresh) { message.error(t('common.loadFailed')); } diff --git a/ui/src/components/web-session/WebSessionPanel.vue b/ui/src/components/web-session/WebSessionPanel.vue index 7c5eaf8a..32e0ea00 100644 --- a/ui/src/components/web-session/WebSessionPanel.vue +++ b/ui/src/components/web-session/WebSessionPanel.vue @@ -14,6 +14,13 @@ + + {{ t('webSession.importCodexSession') }} + + + {{ t('webSession.treeOpen') }} +
@@ -608,7 +646,7 @@ class="timeline-intro" > - {{ currentSession.agent === 'codex' ? 'Codex' : 'Claude' }} + {{ getAgentDisplayName(currentSession.agent) }}
{{ t('webSession.readyTitle') }}
{{ t('webSession.readyDescription') }}
@@ -1756,7 +1794,7 @@ :options="claudeRuntimeOptions" /> @@ -2095,7 +2134,10 @@
{{ pendingInputPreview(item) }}
-
+
+ {{ t('webSession.pendingNativeQueued') }} +
+