From 8e3b931fc266718493643d3d46c07c1d9ea2354e Mon Sep 17 00:00:00 2001 From: dj-oyu <68707227+dj-oyu@users.noreply.github.com> Date: Fri, 20 Mar 2026 09:24:27 +0900 Subject: [PATCH] refactor: simplify miniapp and helper packages --- pkg/git/worktree.go | 116 +++++++++++++++----------------- pkg/mediacache/cache.go | 28 ++++---- pkg/miniapp/api.go | 60 ++++++----------- pkg/miniapp/api_research.go | 66 +++++++++--------- pkg/miniapp/auth.go | 6 +- pkg/miniapp/dev.go | 129 ++++++++---------------------------- pkg/miniapp/dev_helpers.go | 92 +++++++++++++++++++++++++ pkg/miniapp/http_helpers.go | 27 ++++++++ pkg/miniapp/logs.go | 12 ++-- pkg/miniapp/miniapp.go | 44 ++++++------ pkg/miniapp/ws.go | 2 +- pkg/orch/broadcaster.go | 48 +++++++++----- pkg/research/store.go | 31 +++------ 13 files changed, 338 insertions(+), 323 deletions(-) create mode 100644 pkg/miniapp/dev_helpers.go create mode 100644 pkg/miniapp/http_helpers.go diff --git a/pkg/git/worktree.go b/pkg/git/worktree.go index 9e75f3c5f..a92a32fa4 100644 --- a/pkg/git/worktree.go +++ b/pkg/git/worktree.go @@ -49,26 +49,45 @@ type DisposeResult struct { CommitsAhead int // unique commits on branch (0 = safe to delete) } -// FindRepoRoot returns the git repository root for dir, or "" if not a git repo. -func FindRepoRoot(dir string) string { - cmd := exec.Command("git", "rev-parse", "--show-toplevel") +func gitOutputTrim(dir string, args ...string) (string, error) { + cmd := exec.Command("git", args...) cmd.Dir = dir out, err := cmd.Output() + if err != nil { + return "", err + } + return strings.TrimSpace(string(out)), nil +} + +func gitCombinedOutputTrim(dir string, args ...string) (string, error) { + cmd := exec.Command("git", args...) + cmd.Dir = dir + out, err := cmd.CombinedOutput() + return strings.TrimSpace(string(out)), err +} + +func gitRunOK(dir string, args ...string) bool { + cmd := exec.Command("git", args...) + cmd.Dir = dir + return cmd.Run() == nil +} + +// FindRepoRoot returns the git repository root for dir, or "" if not a git repo. +func FindRepoRoot(dir string) string { + out, err := gitOutputTrim(dir, "rev-parse", "--show-toplevel") if err != nil { return "" } - return strings.TrimSpace(string(out)) + return out } // CurrentBranch returns the current branch name, or "" on error. func CurrentBranch(dir string) string { - cmd := exec.Command("git", "rev-parse", "--abbrev-ref", "HEAD") - cmd.Dir = dir - out, err := cmd.Output() + out, err := gitOutputTrim(dir, "rev-parse", "--abbrev-ref", "HEAD") if err != nil { return "" } - return strings.TrimSpace(string(out)) + return out } var unsafeBranchRe = regexp.MustCompile(`[^a-z0-9-]`) @@ -111,9 +130,7 @@ func CreateWorktree(repoDir, worktreePath, branchName string) (*WorktreeInfo, er } // Check if branch already exists - checkCmd := exec.Command("git", "rev-parse", "--verify", branchName) - checkCmd.Dir = repoDir - branchExists := checkCmd.Run() == nil + branchExists := gitRunOK(repoDir, "rev-parse", "--verify", branchName) var cmd *exec.Cmd if branchExists { @@ -138,44 +155,36 @@ func CreateWorktree(repoDir, worktreePath, branchName string) (*WorktreeInfo, er // HasUncommittedChanges returns true if the working tree has staged or unstaged changes. func HasUncommittedChanges(dir string) bool { - cmd := exec.Command("git", "status", "--porcelain") - cmd.Dir = dir - out, err := cmd.Output() + out, err := gitOutputTrim(dir, "status", "--porcelain") if err != nil { return false } - return len(strings.TrimSpace(string(out))) > 0 + return out != "" } // AutoCommit stages all changes and commits with the given message. func AutoCommit(worktreePath, message string) error { - addCmd := exec.Command("git", "add", "-A") - addCmd.Dir = worktreePath - if out, err := addCmd.CombinedOutput(); err != nil { - return fmt.Errorf("git add: %s: %w", strings.TrimSpace(string(out)), err) + if out, err := gitCombinedOutputTrim(worktreePath, "add", "-A"); err != nil { + return fmt.Errorf("git add: %s: %w", out, err) } - commitCmd := exec.Command("git", "commit", "-m", message, "--allow-empty-message") - commitCmd.Dir = worktreePath - if out, err := commitCmd.CombinedOutput(); err != nil { + if out, err := gitCombinedOutputTrim(worktreePath, "commit", "-m", message, "--allow-empty-message"); err != nil { // "nothing to commit" is not a real error - if strings.Contains(string(out), "nothing to commit") { + if strings.Contains(out, "nothing to commit") { return nil } - return fmt.Errorf("git commit: %s: %w", strings.TrimSpace(string(out)), err) + return fmt.Errorf("git commit: %s: %w", out, err) } return nil } // CommitsAhead returns the number of commits on branch that are not on base. func CommitsAhead(repoDir, base, branch string) int { - cmd := exec.Command("git", "rev-list", "--count", base+".."+branch) - cmd.Dir = repoDir - out, err := cmd.Output() + out, err := gitOutputTrim(repoDir, "rev-list", "--count", base+".."+branch) if err != nil { return 0 } - n, _ := strconv.Atoi(strings.TrimSpace(string(out))) + n, _ := strconv.Atoi(out) return n } @@ -360,13 +369,11 @@ func DisposeManagedWorktree(repoDir, worktreesDir, name, baseBranch string) (Dis // WorktreeStatusShort returns "git status --short" output for a worktree. func WorktreeStatusShort(worktreePath string) (string, error) { - cmd := exec.Command("git", "status", "--short") - cmd.Dir = worktreePath - out, err := cmd.CombinedOutput() + out, err := gitCombinedOutputTrim(worktreePath, "status", "--short") if err != nil { - return "", fmt.Errorf("git status --short: %s: %w", strings.TrimSpace(string(out)), err) + return "", fmt.Errorf("git status --short: %s: %w", out, err) } - return strings.TrimSpace(string(out)), nil + return out, nil } // WorktreeRecentLog returns recent oneline commits for a worktree branch. @@ -374,24 +381,20 @@ func WorktreeRecentLog(worktreePath string, n int) (string, error) { if n <= 0 { n = 10 } - cmd := exec.Command("git", "log", "--oneline", fmt.Sprintf("-%d", n)) - cmd.Dir = worktreePath - out, err := cmd.CombinedOutput() + out, err := gitCombinedOutputTrim(worktreePath, "log", "--oneline", fmt.Sprintf("-%d", n)) if err != nil { - return "", fmt.Errorf("git log --oneline: %s: %w", strings.TrimSpace(string(out)), err) + return "", fmt.Errorf("git log --oneline: %s: %w", out, err) } - return strings.TrimSpace(string(out)), nil + return out, nil } // WorktreeDiffStat returns a compact diff stat for a worktree. func WorktreeDiffStat(worktreePath string) (string, error) { - cmd := exec.Command("git", "diff", "--stat") - cmd.Dir = worktreePath - out, err := cmd.CombinedOutput() + out, err := gitCombinedOutputTrim(worktreePath, "diff", "--stat") if err != nil { - return "", fmt.Errorf("git diff --stat: %s: %w", strings.TrimSpace(string(out)), err) + return "", fmt.Errorf("git diff --stat: %s: %w", out, err) } - return strings.TrimSpace(string(out)), nil + return out, nil } // DetectDefaultBranch returns the preferred base branch name for merges/dispose. @@ -404,10 +407,7 @@ func DetectDefaultBranch(repoDir string) string { } // Try origin/HEAD -> origin/ - cmd := exec.Command("git", "symbolic-ref", "refs/remotes/origin/HEAD") - cmd.Dir = repoDir - if out, err := cmd.Output(); err == nil { - ref := strings.TrimSpace(string(out)) + if ref, err := gitOutputTrim(repoDir, "symbolic-ref", "refs/remotes/origin/HEAD"); err == nil { if idx := strings.LastIndex(ref, "/"); idx >= 0 && idx < len(ref)-1 { return ref[idx+1:] } @@ -559,13 +559,11 @@ func buildManagedWorktree(repoDir, name, wtPath string) ManagedWorktree { } func lastCommitInfo(dir string) (hash, subject, age string) { - cmd := exec.Command("git", "log", "-1", "--pretty=format:%h\x1f%s\x1f%cr") - cmd.Dir = dir - out, err := cmd.Output() + out, err := gitOutputTrim(dir, "log", "-1", "--pretty=format:%h\x1f%s\x1f%cr") if err != nil { return "", "", "" } - parts := strings.SplitN(strings.TrimSpace(string(out)), "\x1f", 3) + parts := strings.SplitN(out, "\x1f", 3) if len(parts) > 0 { hash = parts[0] } @@ -579,28 +577,22 @@ func lastCommitInfo(dir string) (hash, subject, age string) { } func localBranchExists(repoDir, name string) bool { - cmd := exec.Command("git", "rev-parse", "--verify", "refs/heads/"+name) - cmd.Dir = repoDir - return cmd.Run() == nil + return gitRunOK(repoDir, "rev-parse", "--verify", "refs/heads/"+name) } func checkoutBranch(repoDir, branch string) error { - cmd := exec.Command("git", "checkout", branch) - cmd.Dir = repoDir - if out, err := cmd.CombinedOutput(); err != nil { - return fmt.Errorf("git checkout %s: %s: %w", branch, strings.TrimSpace(string(out)), err) + if out, err := gitCombinedOutputTrim(repoDir, "checkout", branch); err != nil { + return fmt.Errorf("git checkout %s: %s: %w", branch, out, err) } return nil } func listGitWorktreePaths(repoDir string) ([]string, error) { - cmd := exec.Command("git", "worktree", "list", "--porcelain") - cmd.Dir = repoDir - out, err := cmd.Output() + out, err := gitOutputTrim(repoDir, "worktree", "list", "--porcelain") if err != nil { return nil, err } - lines := strings.Split(string(out), "\n") + lines := strings.Split(out, "\n") paths := make([]string, 0) for _, line := range lines { if !strings.HasPrefix(line, "worktree ") { diff --git a/pkg/mediacache/cache.go b/pkg/mediacache/cache.go index ecbfb3bb8..05309e848 100644 --- a/pkg/mediacache/cache.go +++ b/pkg/mediacache/cache.go @@ -37,6 +37,10 @@ type Cache struct { db *sql.DB } +func currentTimestamp() string { + return time.Now().UTC().Format(time.RFC3339) +} + // Open opens (or creates) a media cache database at dbPath. func Open(dbPath string) (*Cache, error) { connStr := "file:" + dbPath + "?_journal_mode=WAL&_busy_timeout=5000" @@ -70,18 +74,13 @@ func (c *Cache) Get(hash, entryType string) (string, bool) { if err != nil { return "", false } - // Update accessed_at - now := time.Now().UTC().Format(time.RFC3339) - _, _ = c.db.Exec( - `UPDATE media_cache SET accessed_at = ? WHERE hash = ? AND type = ?`, - now, hash, entryType, - ) + c.touchAccessed(hash, entryType) return result, true } // Put stores a result in the cache. func (c *Cache) Put(hash, entryType, result string) error { - now := time.Now().UTC().Format(time.RFC3339) + now := currentTimestamp() _, err := c.db.Exec( `INSERT INTO media_cache (hash, type, result, created_at, accessed_at) VALUES (?, ?, ?, ?, ?) @@ -110,17 +109,13 @@ func (c *Cache) GetEntry(hash, entryType string) (Entry, bool) { if err != nil { return Entry{}, false } - now := time.Now().UTC().Format(time.RFC3339) - _, _ = c.db.Exec( - `UPDATE media_cache SET accessed_at = ? WHERE hash = ? AND type = ?`, - now, hash, entryType, - ) + c.touchAccessed(hash, entryType) return e, true } // PutEntry stores a full entry in the cache. func (c *Cache) PutEntry(hash, entryType string, entry Entry) error { - now := time.Now().UTC().Format(time.RFC3339) + now := currentTimestamp() _, err := c.db.Exec( `INSERT INTO media_cache (hash, type, result, file_path, pages, created_at, accessed_at) VALUES (?, ?, ?, ?, ?, ?, ?) @@ -188,6 +183,13 @@ func (c *Cache) Prune(ttl time.Duration) (int64, error) { return res.RowsAffected() } +func (c *Cache) touchAccessed(hash, entryType string) { + _, _ = c.db.Exec( + `UPDATE media_cache SET accessed_at = ? WHERE hash = ? AND type = ?`, + currentTimestamp(), hash, entryType, + ) +} + // HashData computes a fast FNV-1a 64-bit hash of the given data, // returned as a hex string. func HashData(data []byte) string { diff --git a/pkg/miniapp/api.go b/pkg/miniapp/api.go index 6854aa9c4..b0a170e76 100644 --- a/pkg/miniapp/api.go +++ b/pkg/miniapp/api.go @@ -5,7 +5,6 @@ import ( "encoding/json" "errors" "fmt" - "io" "net/http" "path/filepath" "strings" @@ -61,7 +60,7 @@ func (h *Handler) apiGit(w http.ResponseWriter, r *http.Request) { func (h *Handler) apiWorktrees(w http.ResponseWriter, r *http.Request) { repoRoot := git.FindRepoRoot(h.workspace) if repoRoot == "" { - http.Error(w, `{"error":"workspace is not a git repository"}`, http.StatusBadRequest) + writeJSONError(w, http.StatusBadRequest, "workspace is not a git repository") return } worktreesDir := filepath.Join(h.workspace, ".worktrees") @@ -70,26 +69,20 @@ func (h *Handler) apiWorktrees(w http.ResponseWriter, r *http.Request) { case http.MethodGet: items, err := git.ListManagedWorktrees(repoRoot, worktreesDir) if err != nil { - http.Error(w, `{"error":"failed to list worktrees"}`, http.StatusInternalServerError) + writeJSONError(w, http.StatusInternalServerError, "failed to list worktrees") return } writeJSON(w, items) case http.MethodPost: - body, err := io.ReadAll(io.LimitReader(r.Body, 4096)) - if err != nil { - http.Error(w, `{"error":"bad request"}`, http.StatusBadRequest) - return - } - var req struct { Action string `json:"action"` Name string `json:"name"` Force bool `json:"force"` BaseBranch string `json:"base_branch"` } - if err := json.Unmarshal(body, &req); err != nil { - http.Error(w, `{"error":"invalid JSON"}`, http.StatusBadRequest) + if err := decodeJSONBody(r, 4096, &req); err != nil { + writeJSONError(w, http.StatusBadRequest, "invalid JSON") return } @@ -97,7 +90,7 @@ func (h *Handler) apiWorktrees(w http.ResponseWriter, r *http.Request) { req.Name = strings.TrimSpace(req.Name) req.BaseBranch = strings.TrimSpace(req.BaseBranch) if req.Action == "" || req.Name == "" { - http.Error(w, `{"error":"action and name are required"}`, http.StatusBadRequest) + writeJSONError(w, http.StatusBadRequest, "action and name are required") return } @@ -108,7 +101,7 @@ func (h *Handler) apiWorktrees(w http.ResponseWriter, r *http.Request) { if writeWorktreeAPIError(w, err) { return } - http.Error(w, `{"error":"merge failed"}`, http.StatusInternalServerError) + writeJSONError(w, http.StatusInternalServerError, "merge failed") return } writeJSON(w, map[string]any{ @@ -125,15 +118,11 @@ func (h *Handler) apiWorktrees(w http.ResponseWriter, r *http.Request) { if writeWorktreeAPIError(w, err) { return } - http.Error(w, `{"error":"failed to inspect worktree"}`, http.StatusInternalServerError) + writeJSONError(w, http.StatusInternalServerError, "failed to inspect worktree") return } if wt.HasUncommitted && !req.Force { - http.Error( - w, - `{"error":"worktree has uncommitted changes; retry with force=true"}`, - http.StatusConflict, - ) + writeJSONError(w, http.StatusConflict, "worktree has uncommitted changes; retry with force=true") return } @@ -142,7 +131,7 @@ func (h *Handler) apiWorktrees(w http.ResponseWriter, r *http.Request) { if writeWorktreeAPIError(w, err) { return } - http.Error(w, `{"error":"dispose failed"}`, http.StatusInternalServerError) + writeJSONError(w, http.StatusInternalServerError, "dispose failed") return } writeJSON(w, map[string]any{ @@ -153,11 +142,11 @@ func (h *Handler) apiWorktrees(w http.ResponseWriter, r *http.Request) { }) default: - http.Error(w, `{"error":"unknown action"}`, http.StatusBadRequest) + writeJSONError(w, http.StatusBadRequest, "unknown action") } default: - http.Error(w, `{"error":"method not allowed"}`, http.StatusMethodNotAllowed) + writeMethodNotAllowed(w) } } @@ -174,26 +163,20 @@ func (h *Handler) apiSessionGraph(w http.ResponseWriter, r *http.Request) { func (h *Handler) apiCommand(w http.ResponseWriter, r *http.Request) { if r.Method != http.MethodPost { - http.Error(w, `{"error":"method not allowed"}`, http.StatusMethodNotAllowed) - return - } - - body, err := io.ReadAll(io.LimitReader(r.Body, 4096)) - if err != nil { - http.Error(w, `{"error":"bad request"}`, http.StatusBadRequest) + writeMethodNotAllowed(w) return } var req struct { Command string `json:"command"` } - if err := json.Unmarshal(body, &req); err != nil || req.Command == "" { - http.Error(w, `{"error":"missing command"}`, http.StatusBadRequest) + if err := decodeJSONBody(r, 4096, &req); err != nil || req.Command == "" { + writeJSONError(w, http.StatusBadRequest, "missing command") return } if !strings.HasPrefix(req.Command, "/") { - http.Error(w, `{"error":"command must start with /"}`, http.StatusBadRequest) + writeJSONError(w, http.StatusBadRequest, "command must start with /") return } @@ -201,7 +184,7 @@ func (h *Handler) apiCommand(w http.ResponseWriter, r *http.Request) { initData := r.URL.Query().Get("initData") userID, chatID := extractUserFromInitData(initData) if userID == "" { - http.Error(w, `{"error":"cannot identify user"}`, http.StatusBadRequest) + writeJSONError(w, http.StatusBadRequest, "cannot identify user") return } @@ -212,7 +195,7 @@ func (h *Handler) apiCommand(w http.ResponseWriter, r *http.Request) { func (h *Handler) apiEvents(w http.ResponseWriter, r *http.Request) { flusher, ok := w.(http.Flusher) if !ok { - http.Error(w, `{"error":"streaming not supported"}`, http.StatusInternalServerError) + writeJSONError(w, http.StatusInternalServerError, "streaming not supported") return } rc := http.NewResponseController(w) @@ -288,18 +271,13 @@ func sendSSEIfChanged(w http.ResponseWriter, f http.Flusher, event string, v any } } -func writeJSON(w http.ResponseWriter, v any) { - w.Header().Set("Content-Type", "application/json") - json.NewEncoder(w).Encode(v) -} - func writeWorktreeAPIError(w http.ResponseWriter, err error) bool { switch { case errors.Is(err, git.ErrInvalidWorktreeName): - http.Error(w, `{"error":"invalid worktree name"}`, http.StatusBadRequest) + writeJSONError(w, http.StatusBadRequest, "invalid worktree name") return true case errors.Is(err, git.ErrWorktreeNotFound): - http.Error(w, `{"error":"worktree not found"}`, http.StatusNotFound) + writeJSONError(w, http.StatusNotFound, "worktree not found") return true default: return false diff --git a/pkg/miniapp/api_research.go b/pkg/miniapp/api_research.go index cebfbdaab..8b348471a 100644 --- a/pkg/miniapp/api_research.go +++ b/pkg/miniapp/api_research.go @@ -1,8 +1,6 @@ package miniapp import ( - "encoding/json" - "io" "net/http" "os" "path/filepath" @@ -24,7 +22,7 @@ func (h *Handler) SetResearchFocus(ft *research.FocusTracker) { // apiResearch handles GET /miniapp/api/research (list) and POST (create). func (h *Handler) apiResearch(w http.ResponseWriter, r *http.Request) { if h.researchStore == nil { - http.Error(w, `{"error":"research store not available"}`, http.StatusServiceUnavailable) + writeJSONError(w, http.StatusServiceUnavailable, "research store not available") return } @@ -34,7 +32,7 @@ func (h *Handler) apiResearch(w http.ResponseWriter, r *http.Request) { case http.MethodPost: h.apiResearchCreate(w, r) default: - http.Error(w, `{"error":"method not allowed"}`, http.StatusMethodNotAllowed) + writeMethodNotAllowed(w) } } @@ -80,7 +78,7 @@ func (h *Handler) apiResearchList(w http.ResponseWriter, r *http.Request) { statusFilter := r.URL.Query().Get("status") tasks, err := h.researchStore.ListTasks(research.TaskStatus(statusFilter)) if err != nil { - http.Error(w, `{"error":"`+err.Error()+`"}`, http.StatusInternalServerError) + writeJSONError(w, http.StatusInternalServerError, err.Error()) return } @@ -102,12 +100,12 @@ func (h *Handler) apiResearchCreate(w http.ResponseWriter, r *http.Request) { Description string `json:"description"` Interval string `json:"interval"` } - if err := json.NewDecoder(io.LimitReader(r.Body, 1<<16)).Decode(&req); err != nil { - http.Error(w, `{"error":"invalid request body"}`, http.StatusBadRequest) + if err := decodeJSONBody(r, 1<<16, &req); err != nil { + writeJSONError(w, http.StatusBadRequest, "invalid request body") return } if strings.TrimSpace(req.Title) == "" { - http.Error(w, `{"error":"title is required"}`, http.StatusBadRequest) + writeJSONError(w, http.StatusBadRequest, "title is required") return } @@ -117,7 +115,7 @@ func (h *Handler) apiResearchCreate(w http.ResponseWriter, r *http.Request) { strings.TrimSpace(req.Interval), ) if err != nil { - http.Error(w, `{"error":"`+err.Error()+`"}`, http.StatusInternalServerError) + writeJSONError(w, http.StatusInternalServerError, err.Error()) return } @@ -128,7 +126,7 @@ func (h *Handler) apiResearchCreate(w http.ResponseWriter, r *http.Request) { // apiResearchDetail handles /miniapp/api/research/{id} and /miniapp/api/research/{id}/doc/{docId}. func (h *Handler) apiResearchDetail(w http.ResponseWriter, r *http.Request) { if h.researchStore == nil { - http.Error(w, `{"error":"research store not available"}`, http.StatusServiceUnavailable) + writeJSONError(w, http.StatusServiceUnavailable, "research store not available") return } @@ -145,7 +143,7 @@ func (h *Handler) apiResearchDetail(w http.ResponseWriter, r *http.Request) { case http.MethodPost: h.apiResearchTaskAction(w, r, taskID) default: - http.Error(w, `{"error":"method not allowed"}`, http.StatusMethodNotAllowed) + writeMethodNotAllowed(w) } return } @@ -153,7 +151,7 @@ func (h *Handler) apiResearchDetail(w http.ResponseWriter, r *http.Request) { if len(parts) == 3 && parts[1] == "doc" && parts[2] != "" { // /miniapp/api/research/{id}/doc/{docId} if r.Method != http.MethodGet { - http.Error(w, `{"error":"method not allowed"}`, http.StatusMethodNotAllowed) + writeMethodNotAllowed(w) return } h.apiResearchGetDoc(w, parts[0], parts[2]) @@ -182,7 +180,7 @@ type researchTaskDetailResponse struct { func (h *Handler) apiResearchGetTask(w http.ResponseWriter, taskID string) { task, err := h.researchStore.GetTask(taskID) if err != nil { - http.Error(w, `{"error":"task not found"}`, http.StatusNotFound) + writeJSONError(w, http.StatusNotFound, "task not found") return } @@ -219,36 +217,36 @@ func (h *Handler) apiResearchTaskAction(w http.ResponseWriter, r *http.Request, Description string `json:"description"` Interval string `json:"interval"` } - if err := json.NewDecoder(io.LimitReader(r.Body, 1<<16)).Decode(&req); err != nil { - http.Error(w, `{"error":"invalid request body"}`, http.StatusBadRequest) + if err := decodeJSONBody(r, 1<<16, &req); err != nil { + writeJSONError(w, http.StatusBadRequest, "invalid request body") return } switch req.Action { case "cancel": if err := h.researchStore.SetTaskStatus(taskID, research.StatusCanceled); err != nil { - http.Error(w, `{"error":"`+err.Error()+`"}`, http.StatusBadRequest) + writeJSONError(w, http.StatusBadRequest, err.Error()) return } case "reopen": if err := h.researchStore.SetTaskStatus(taskID, research.StatusPending); err != nil { - http.Error(w, `{"error":"`+err.Error()+`"}`, http.StatusBadRequest) + writeJSONError(w, http.StatusBadRequest, err.Error()) return } case "activate": if err := h.researchStore.SetTaskStatus(taskID, research.StatusActive); err != nil { - http.Error(w, `{"error":"`+err.Error()+`"}`, http.StatusBadRequest) + writeJSONError(w, http.StatusBadRequest, err.Error()) return } case "complete": if err := h.researchStore.SetTaskStatus(taskID, research.StatusCompleted); err != nil { - http.Error(w, `{"error":"`+err.Error()+`"}`, http.StatusBadRequest) + writeJSONError(w, http.StatusBadRequest, err.Error()) return } case "update": task, err := h.researchStore.GetTask(taskID) if err != nil { - http.Error(w, `{"error":"task not found"}`, http.StatusNotFound) + writeJSONError(w, http.StatusNotFound, "task not found") return } title := req.Title @@ -260,20 +258,20 @@ func (h *Handler) apiResearchTaskAction(w http.ResponseWriter, r *http.Request, desc = task.Description } if err := h.researchStore.UpdateTask(taskID, title, desc); err != nil { - http.Error(w, `{"error":"`+err.Error()+`"}`, http.StatusInternalServerError) + writeJSONError(w, http.StatusInternalServerError, err.Error()) return } case "set_interval": if req.Interval == "" { - http.Error(w, `{"error":"interval is required"}`, http.StatusBadRequest) + writeJSONError(w, http.StatusBadRequest, "interval is required") return } if err := h.researchStore.SetInterval(taskID, req.Interval); err != nil { - http.Error(w, `{"error":"`+err.Error()+`"}`, http.StatusBadRequest) + writeJSONError(w, http.StatusBadRequest, err.Error()) return } default: - http.Error(w, `{"error":"unknown action"}`, http.StatusBadRequest) + writeJSONError(w, http.StatusBadRequest, "unknown action") return } @@ -284,7 +282,7 @@ func (h *Handler) apiResearchTaskAction(w http.ResponseWriter, r *http.Request, func (h *Handler) apiResearchGetDoc(w http.ResponseWriter, taskID, docID string) { docs, err := h.researchStore.ListDocuments(taskID) if err != nil { - http.Error(w, `{"error":"failed to list documents"}`, http.StatusInternalServerError) + writeJSONError(w, http.StatusInternalServerError, "failed to list documents") return } @@ -296,7 +294,7 @@ func (h *Handler) apiResearchGetDoc(w http.ResponseWriter, taskID, docID string) } } if found == nil { - http.Error(w, `{"error":"document not found"}`, http.StatusNotFound) + writeJSONError(w, http.StatusNotFound, "document not found") return } @@ -308,7 +306,7 @@ func (h *Handler) apiResearchGetDoc(w http.ResponseWriter, taskID, docID string) content, err := os.ReadFile(absPath) if err != nil { - http.Error(w, `{"error":"failed to read document"}`, http.StatusInternalServerError) + writeJSONError(w, http.StatusInternalServerError, "failed to read document") return } @@ -324,7 +322,7 @@ func (h *Handler) apiResearchGetDoc(w http.ResponseWriter, taskID, docID string) // GET returns the current focus state; POST sets focus/unfocus for a task. func (h *Handler) apiResearchFocus(w http.ResponseWriter, r *http.Request) { if h.researchFocus == nil { - http.Error(w, `{"error":"research focus not available"}`, http.StatusServiceUnavailable) + writeJSONError(w, http.StatusServiceUnavailable, "research focus not available") return } @@ -334,7 +332,7 @@ func (h *Handler) apiResearchFocus(w http.ResponseWriter, r *http.Request) { case http.MethodPost: h.apiResearchFocusSet(w, r) default: - http.Error(w, `{"error":"method not allowed"}`, http.StatusMethodNotAllowed) + writeMethodNotAllowed(w) } } @@ -351,20 +349,20 @@ func (h *Handler) apiResearchFocusSet(w http.ResponseWriter, r *http.Request) { Action string `json:"action"` // "recall" or "forget" TaskID string `json:"task_id"` } - if err := json.NewDecoder(io.LimitReader(r.Body, 1<<14)).Decode(&req); err != nil { - http.Error(w, `{"error":"invalid request body"}`, http.StatusBadRequest) + if err := decodeJSONBody(r, 1<<14, &req); err != nil { + writeJSONError(w, http.StatusBadRequest, "invalid request body") return } switch req.Action { case "recall": if req.TaskID == "" { - http.Error(w, `{"error":"task_id is required"}`, http.StatusBadRequest) + writeJSONError(w, http.StatusBadRequest, "task_id is required") return } task, err := h.researchStore.GetTask(req.TaskID) if err != nil { - http.Error(w, `{"error":"task not found"}`, http.StatusNotFound) + writeJSONError(w, http.StatusNotFound, "task not found") return } h.researchFocus.Focus(task.ID, task.Title) @@ -375,7 +373,7 @@ func (h *Handler) apiResearchFocusSet(w http.ResponseWriter, r *http.Request) { h.researchFocus.Unfocus(req.TaskID) } default: - http.Error(w, `{"error":"action must be recall or forget"}`, http.StatusBadRequest) + writeJSONError(w, http.StatusBadRequest, "action must be recall or forget") return } diff --git a/pkg/miniapp/auth.go b/pkg/miniapp/auth.go index cc2b3406f..180e02c37 100644 --- a/pkg/miniapp/auth.go +++ b/pkg/miniapp/auth.go @@ -21,17 +21,17 @@ func (h *Handler) requireAuth(next http.HandlerFunc) http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { initData := r.URL.Query().Get("initData") if initData == "" { - http.Error(w, `{"error":"missing initData"}`, http.StatusUnauthorized) + writeJSONError(w, http.StatusUnauthorized, "missing initData") return } if !ValidateInitData(initData, h.botToken) { - http.Error(w, `{"error":"invalid initData"}`, http.StatusUnauthorized) + writeJSONError(w, http.StatusUnauthorized, "invalid initData") return } if len(h.allowList) > 0 { userID, _ := extractUserFromInitData(initData) if userID == "" || !isAllowed(userID, h.allowList) { - http.Error(w, `{"error":"forbidden"}`, http.StatusForbidden) + writeJSONError(w, http.StatusForbidden, "forbidden") return } } diff --git a/pkg/miniapp/dev.go b/pkg/miniapp/dev.go index 67bfcd898..32d720260 100644 --- a/pkg/miniapp/dev.go +++ b/pkg/miniapp/dev.go @@ -5,11 +5,8 @@ import ( "encoding/json" "fmt" "io" - "net" "net/http" - "net/http/httputil" "net/url" - "sort" "strconv" "strings" "time" @@ -45,9 +42,7 @@ func (h *Handler) RegisterDevTarget(name, target string) (string, error) { id := strconv.Itoa(h.devNextID) h.devTargets[id] = &DevTarget{ID: id, Name: name, Target: target} - if h.notifier != nil { - h.notifier.Notify() - } + h.notifyStateChanged() return id, nil } @@ -64,13 +59,9 @@ func (h *Handler) UnregisterDevTarget(id string) error { delete(h.devTargets, id) if h.devActiveID == id { - h.devActiveID = "" - h.devTarget = nil - h.devProxy = nil - } - if h.notifier != nil { - h.notifier.Notify() + h.clearActiveDevTargetLocked() } + h.notifyStateChanged() return nil } @@ -86,59 +77,13 @@ func (h *Handler) ActivateDevTarget(id string) error { return fmt.Errorf("target %q not found", id) } - u, err := url.Parse(dt.Target) + u, proxy, err := buildDevProxy(dt) if err != nil { - return fmt.Errorf("invalid URL: %w", err) + return err } - // Fix IPv6: resolve "localhost" to 127.0.0.1 to avoid connection refused on systems - // where localhost resolves to [::1] but the dev server only listens on IPv4. - if u.Hostname() == "localhost" { - u.Host = net.JoinHostPort("127.0.0.1", u.Port()) - } - - proxy := httputil.NewSingleHostReverseProxy(u) - proxy.ModifyResponse = func(resp *http.Response) error { - // Prevent browser/WebView from caching dev proxy responses (CSS, JS, etc.) - resp.Header.Set("Cache-Control", "no-cache, no-store, must-revalidate") - resp.Header.Del("ETag") - resp.Header.Del("Last-Modified") - - ct := resp.Header.Get("Content-Type") - if !strings.Contains(ct, "text/html") { - return nil - } - body, err := io.ReadAll(resp.Body) - if err != nil { - return err - } - resp.Body.Close() - modified := injectDevProxyScript(body) - resp.Body = io.NopCloser(bytes.NewReader(modified)) - resp.ContentLength = int64(len(modified)) - resp.Header.Set("Content-Length", strconv.Itoa(len(modified))) - resp.Header.Del("Content-Encoding") - return nil - } - proxy.ErrorHandler = func(w http.ResponseWriter, r *http.Request, err error) { - w.Header().Set("Content-Type", "text/html; charset=utf-8") - w.WriteHeader(http.StatusBadGateway) - fmt.Fprintf(w, ` -

Cannot connect

%s

Target: %s

`, - escapeHTMLString(err.Error()), escapeHTMLString(dt.Target)) - } - - h.devTarget = u - h.devProxy = proxy - h.devActiveID = id - if h.notifier != nil { - h.notifier.Notify() - } + h.setActiveDevTargetLocked(id, u, proxy) + h.notifyStateChanged() return nil } @@ -149,12 +94,8 @@ func (h *Handler) DeactivateDevTarget() error { h.devMu.Lock() defer h.devMu.Unlock() - h.devActiveID = "" - h.devTarget = nil - h.devProxy = nil - if h.notifier != nil { - h.notifier.Notify() - } + h.clearActiveDevTargetLocked() + h.notifyStateChanged() return nil } @@ -176,14 +117,7 @@ func (h *Handler) GetDevTarget() string { func (h *Handler) ListDevTargets() []DevTarget { h.devMu.RLock() defer h.devMu.RUnlock() - - targets := make([]DevTarget, 0, len(h.devTargets)) - for _, dt := range h.devTargets { - targets = append(targets, *dt) - } - // Sort by ID for stable order - sort.Slice(targets, func(i, j int) bool { return targets[i].ID < targets[j].ID }) - return targets + return h.sortedDevTargetsLocked() } // devProxyScript is the JavaScript injected into HTML responses from the dev proxy. @@ -301,50 +235,45 @@ func (h *Handler) apiDev(w http.ResponseWriter, r *http.Request) { case http.MethodGet: writeJSON(w, h.devStatus()) case http.MethodPost: - body, err := io.ReadAll(io.LimitReader(r.Body, 4096)) - if err != nil { - http.Error(w, `{"error":"bad request"}`, http.StatusBadRequest) - return - } var req struct { Action string `json:"action"` ID string `json:"id"` } - if err := json.Unmarshal(body, &req); err != nil { - http.Error(w, `{"error":"invalid JSON"}`, http.StatusBadRequest) + if err := decodeJSONBody(r, 4096, &req); err != nil { + writeJSONError(w, http.StatusBadRequest, "invalid JSON") return } switch req.Action { case "activate": if req.ID == "" { - writeJSON(w, map[string]any{"error": "id is required"}) + writeJSONError(w, http.StatusBadRequest, "id is required") return } if err := h.ActivateDevTarget(req.ID); err != nil { - writeJSON(w, map[string]any{"error": err.Error()}) + writeJSONError(w, http.StatusBadRequest, err.Error()) return } case "deactivate": if err := h.DeactivateDevTarget(); err != nil { - writeJSON(w, map[string]any{"error": err.Error()}) + writeJSONError(w, http.StatusBadRequest, err.Error()) return } case "unregister": if req.ID == "" { - writeJSON(w, map[string]any{"error": "id is required"}) + writeJSONError(w, http.StatusBadRequest, "id is required") return } if err := h.UnregisterDevTarget(req.ID); err != nil { - writeJSON(w, map[string]any{"error": err.Error()}) + writeJSONError(w, http.StatusBadRequest, err.Error()) return } default: - writeJSON(w, map[string]any{"error": "unknown action"}) + writeJSONError(w, http.StatusBadRequest, "unknown action") return } writeJSON(w, h.devStatus()) default: - http.Error(w, `{"error":"method not allowed"}`, http.StatusMethodNotAllowed) + writeMethodNotAllowed(w) } } @@ -354,7 +283,7 @@ func (h *Handler) serveDevProxy(w http.ResponseWriter, r *http.Request) { h.devMu.RUnlock() if proxy == nil { - http.Error(w, "dev proxy not configured", http.StatusServiceUnavailable) + writeJSONError(w, http.StatusServiceUnavailable, "dev proxy not configured") return } @@ -379,30 +308,24 @@ func (h *Handler) devStatus() map[string]any { target = h.devTargets[h.devActiveID].Target // original URL before IPv6 rewrite } - targets := make([]DevTarget, 0, len(h.devTargets)) - for _, dt := range h.devTargets { - targets = append(targets, *dt) - } - sort.Slice(targets, func(i, j int) bool { return targets[i].ID < targets[j].ID }) - return map[string]any{ "active": active, "active_id": h.devActiveID, "target": target, - "targets": targets, + "targets": h.sortedDevTargetsLocked(), } } // apiDevConsole receives console output from dev preview iframes. func (h *Handler) apiDevConsole(w http.ResponseWriter, r *http.Request) { if r.Method != http.MethodPost { - http.Error(w, `{"error":"method not allowed"}`, http.StatusMethodNotAllowed) + writeMethodNotAllowed(w) return } // Only accept console posts when dev proxy is active if h.GetDevTarget() == "" { - http.Error(w, `{"error":"not available"}`, http.StatusNotFound) + writeJSONError(w, http.StatusNotFound, "not available") return } @@ -417,13 +340,13 @@ func (h *Handler) apiDevConsole(w http.ResponseWriter, r *http.Request) { over := h.consoleReqCount > 10 h.consoleMu.Unlock() if over { - http.Error(w, `{"error":"rate limit"}`, http.StatusTooManyRequests) + writeJSONError(w, http.StatusTooManyRequests, "rate limit") return } body, err := io.ReadAll(io.LimitReader(r.Body, 32*1024)) if err != nil { - http.Error(w, `{"error":"bad request"}`, http.StatusBadRequest) + writeJSONError(w, http.StatusBadRequest, "bad request") return } @@ -432,7 +355,7 @@ func (h *Handler) apiDevConsole(w http.ResponseWriter, r *http.Request) { Message string `json:"message"` } if err := json.Unmarshal(body, &entries); err != nil { - http.Error(w, `{"error":"invalid JSON"}`, http.StatusBadRequest) + writeJSONError(w, http.StatusBadRequest, "invalid JSON") return } diff --git a/pkg/miniapp/dev_helpers.go b/pkg/miniapp/dev_helpers.go new file mode 100644 index 000000000..9eccf63e3 --- /dev/null +++ b/pkg/miniapp/dev_helpers.go @@ -0,0 +1,92 @@ +package miniapp + +import ( + "bytes" + "fmt" + "io" + "net" + "net/http" + "net/http/httputil" + "net/url" + "sort" + "strconv" + "strings" +) + +func (h *Handler) notifyStateChanged() { + if h.notifier != nil { + h.notifier.Notify() + } +} + +func (h *Handler) clearActiveDevTargetLocked() { + h.devActiveID = "" + h.devTarget = nil + h.devProxy = nil +} + +func (h *Handler) setActiveDevTargetLocked(id string, target *url.URL, proxy *httputil.ReverseProxy) { + h.devActiveID = id + h.devTarget = target + h.devProxy = proxy +} + +func (h *Handler) sortedDevTargetsLocked() []DevTarget { + targets := make([]DevTarget, 0, len(h.devTargets)) + for _, dt := range h.devTargets { + targets = append(targets, *dt) + } + sort.Slice(targets, func(i, j int) bool { return targets[i].ID < targets[j].ID }) + return targets +} + +func buildDevProxy(target *DevTarget) (*url.URL, *httputil.ReverseProxy, error) { + u, err := url.Parse(target.Target) + if err != nil { + return nil, nil, fmt.Errorf("invalid URL: %w", err) + } + + // Fix IPv6: resolve "localhost" to 127.0.0.1 to avoid connection refused on systems + // where localhost resolves to [::1] but the dev server only listens on IPv4. + if u.Hostname() == "localhost" { + u.Host = net.JoinHostPort("127.0.0.1", u.Port()) + } + + proxy := httputil.NewSingleHostReverseProxy(u) + proxy.ModifyResponse = func(resp *http.Response) error { + // Prevent browser/WebView from caching dev proxy responses (CSS, JS, etc.) + resp.Header.Set("Cache-Control", "no-cache, no-store, must-revalidate") + resp.Header.Del("ETag") + resp.Header.Del("Last-Modified") + + ct := resp.Header.Get("Content-Type") + if !strings.Contains(ct, "text/html") { + return nil + } + body, err := io.ReadAll(resp.Body) + if err != nil { + return err + } + resp.Body.Close() + modified := injectDevProxyScript(body) + resp.Body = io.NopCloser(bytes.NewReader(modified)) + resp.ContentLength = int64(len(modified)) + resp.Header.Set("Content-Length", strconv.Itoa(len(modified))) + resp.Header.Del("Content-Encoding") + return nil + } + proxy.ErrorHandler = func(w http.ResponseWriter, r *http.Request, err error) { + w.Header().Set("Content-Type", "text/html; charset=utf-8") + w.WriteHeader(http.StatusBadGateway) + fmt.Fprintf(w, ` +

Cannot connect

%s

Target: %s

`, + escapeHTMLString(err.Error()), escapeHTMLString(target.Target)) + } + + return u, proxy, nil +} diff --git a/pkg/miniapp/http_helpers.go b/pkg/miniapp/http_helpers.go new file mode 100644 index 000000000..25d6d61db --- /dev/null +++ b/pkg/miniapp/http_helpers.go @@ -0,0 +1,27 @@ +package miniapp + +import ( + "encoding/json" + "io" + "net/http" +) + +func writeJSON(w http.ResponseWriter, v any) { + w.Header().Set("Content-Type", "application/json") + _ = json.NewEncoder(w).Encode(v) +} + +func writeJSONError(w http.ResponseWriter, status int, message string) { + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(status) + _ = json.NewEncoder(w).Encode(map[string]string{"error": message}) +} + +func writeMethodNotAllowed(w http.ResponseWriter) { + writeJSONError(w, http.StatusMethodNotAllowed, "method not allowed") +} + +func decodeJSONBody(r *http.Request, limit int64, dst any) error { + dec := json.NewDecoder(io.LimitReader(r.Body, limit)) + return dec.Decode(dst) +} diff --git a/pkg/miniapp/logs.go b/pkg/miniapp/logs.go index 3e013aca7..d300fcb02 100644 --- a/pkg/miniapp/logs.go +++ b/pkg/miniapp/logs.go @@ -17,7 +17,7 @@ import ( // apiLogsSnapshot creates a tar.gz snapshot of the current log buffer. func (h *Handler) apiLogsSnapshot(w http.ResponseWriter, r *http.Request) { if r.Method != http.MethodPost { - http.Error(w, `{"error":"method not allowed"}`, http.StatusMethodNotAllowed) + writeMethodNotAllowed(w) return } @@ -25,7 +25,7 @@ func (h *Handler) apiLogsSnapshot(w http.ResponseWriter, r *http.Request) { snapshotDir := filepath.Join(h.workspace, "logs", "snapshots") if err := os.MkdirAll(snapshotDir, 0o755); err != nil { - http.Error(w, `{"error":"cannot create snapshot dir"}`, http.StatusInternalServerError) + writeJSONError(w, http.StatusInternalServerError, "cannot create snapshot dir") return } @@ -36,7 +36,7 @@ func (h *Handler) apiLogsSnapshot(w http.ResponseWriter, r *http.Request) { // Create tar.gz f, err := os.Create(snapshotPath) if err != nil { - http.Error(w, `{"error":"cannot create snapshot file"}`, http.StatusInternalServerError) + writeJSONError(w, http.StatusInternalServerError, "cannot create snapshot file") return } @@ -88,7 +88,7 @@ func (h *Handler) apiLogsSnapshot(w http.ResponseWriter, r *http.Request) { // apiLogsSnapshotDownload serves a snapshot tar.gz file. func (h *Handler) apiLogsSnapshotDownload(w http.ResponseWriter, r *http.Request) { if r.Method != http.MethodGet { - http.Error(w, `{"error":"method not allowed"}`, http.StatusMethodNotAllowed) + writeMethodNotAllowed(w) return } @@ -96,7 +96,7 @@ func (h *Handler) apiLogsSnapshotDownload(w http.ResponseWriter, r *http.Request id = filepath.Base(id) // path traversal prevention if id == "" || id == "." || id == ".." { - http.Error(w, `{"error":"invalid id"}`, http.StatusBadRequest) + writeJSONError(w, http.StatusBadRequest, "invalid id") return } @@ -104,7 +104,7 @@ func (h *Handler) apiLogsSnapshotDownload(w http.ResponseWriter, r *http.Request snapshotPath := filepath.Join(h.workspace, "logs", "snapshots", filename) if _, err := os.Stat(snapshotPath); os.IsNotExist(err) { - http.Error(w, `{"error":"not found"}`, http.StatusNotFound) + writeJSONError(w, http.StatusNotFound, "not found") return } diff --git a/pkg/miniapp/miniapp.go b/pkg/miniapp/miniapp.go index ff136ed7f..289292125 100644 --- a/pkg/miniapp/miniapp.go +++ b/pkg/miniapp/miniapp.go @@ -84,33 +84,37 @@ func (h *Handler) SetOrchBroadcaster(b *orch.Broadcaster) { h.orchBroadcaster = b } +func (h *Handler) handleProtectedFunc(mux *http.ServeMux, pattern string, handler http.HandlerFunc) { + mux.HandleFunc(pattern, h.requireAuth(handler)) +} + // RegisterRoutes registers Mini App routes on the given mux. func (h *Handler) RegisterRoutes(mux *http.ServeMux) { mux.HandleFunc("/miniapp", h.serveIndex) mux.HandleFunc("/miniapp/index.html", h.serveIndex) mux.HandleFunc("/miniapp/", h.serveStatic) - mux.HandleFunc("/miniapp/api/skills", h.requireAuth(h.apiSkills)) - mux.HandleFunc("/miniapp/api/plan", h.requireAuth(h.apiPlan)) - mux.HandleFunc("/miniapp/api/session", h.requireAuth(h.apiSession)) - mux.HandleFunc("/miniapp/api/sessions", h.requireAuth(h.apiSessions)) - mux.HandleFunc("/miniapp/api/sessions/graph", h.requireAuth(h.apiSessionGraph)) - mux.HandleFunc("/miniapp/api/command", h.requireAuth(h.apiCommand)) - mux.HandleFunc("/miniapp/api/context", h.requireAuth(h.apiContext)) - mux.HandleFunc("/miniapp/api/prompt", h.requireAuth(h.apiPrompt)) - mux.HandleFunc("/miniapp/api/git", h.requireAuth(h.apiGit)) - mux.HandleFunc("/miniapp/api/worktrees", h.requireAuth(h.apiWorktrees)) - mux.HandleFunc("/miniapp/api/dev", h.requireAuth(h.apiDev)) - mux.HandleFunc("/miniapp/api/events", h.requireAuth(h.apiEvents)) - mux.HandleFunc("/miniapp/api/logs/ws", h.requireAuth(h.wsLogs)) - mux.HandleFunc("/miniapp/api/logs/snapshot", h.requireAuth(h.apiLogsSnapshot)) - mux.HandleFunc("/miniapp/api/logs/snapshot/", h.requireAuth(h.apiLogsSnapshotDownload)) - mux.HandleFunc("/miniapp/api/orchestration/ws", h.requireAuth(h.wsOrchestration)) + h.handleProtectedFunc(mux, "/miniapp/api/skills", h.apiSkills) + h.handleProtectedFunc(mux, "/miniapp/api/plan", h.apiPlan) + h.handleProtectedFunc(mux, "/miniapp/api/session", h.apiSession) + h.handleProtectedFunc(mux, "/miniapp/api/sessions", h.apiSessions) + h.handleProtectedFunc(mux, "/miniapp/api/sessions/graph", h.apiSessionGraph) + h.handleProtectedFunc(mux, "/miniapp/api/command", h.apiCommand) + h.handleProtectedFunc(mux, "/miniapp/api/context", h.apiContext) + h.handleProtectedFunc(mux, "/miniapp/api/prompt", h.apiPrompt) + h.handleProtectedFunc(mux, "/miniapp/api/git", h.apiGit) + h.handleProtectedFunc(mux, "/miniapp/api/worktrees", h.apiWorktrees) + h.handleProtectedFunc(mux, "/miniapp/api/dev", h.apiDev) + h.handleProtectedFunc(mux, "/miniapp/api/events", h.apiEvents) + h.handleProtectedFunc(mux, "/miniapp/api/logs/ws", h.wsLogs) + h.handleProtectedFunc(mux, "/miniapp/api/logs/snapshot", h.apiLogsSnapshot) + h.handleProtectedFunc(mux, "/miniapp/api/logs/snapshot/", h.apiLogsSnapshotDownload) + h.handleProtectedFunc(mux, "/miniapp/api/orchestration/ws", h.wsOrchestration) mux.HandleFunc("/miniapp/dev/console", h.apiDevConsole) mux.HandleFunc("/miniapp/dev/", h.serveDevProxy) - mux.HandleFunc("/miniapp/api/cache", h.requireAuth(h.apiCache)) - mux.HandleFunc("/miniapp/api/research", h.requireAuth(h.apiResearch)) - mux.HandleFunc("/miniapp/api/research/focus", h.requireAuth(h.apiResearchFocus)) - mux.HandleFunc("/miniapp/api/research/", h.requireAuth(h.apiResearchDetail)) + h.handleProtectedFunc(mux, "/miniapp/api/cache", h.apiCache) + h.handleProtectedFunc(mux, "/miniapp/api/research", h.apiResearch) + h.handleProtectedFunc(mux, "/miniapp/api/research/focus", h.apiResearchFocus) + h.handleProtectedFunc(mux, "/miniapp/api/research/", h.apiResearchDetail) } func (h *Handler) serveIndex(w http.ResponseWriter, r *http.Request) { diff --git a/pkg/miniapp/ws.go b/pkg/miniapp/ws.go index 5cfa09196..53b47c50a 100644 --- a/pkg/miniapp/ws.go +++ b/pkg/miniapp/ws.go @@ -160,7 +160,7 @@ func (h *Handler) wsLogs(w http.ResponseWriter, r *http.Request) { // {"type":"event","event":{...orch.Event}} -- pushed on each state change func (h *Handler) wsOrchestration(w http.ResponseWriter, r *http.Request) { if h.orchBroadcaster == nil { - http.Error(w, `{"error":"orchestration not enabled"}`, http.StatusServiceUnavailable) + writeJSONError(w, http.StatusServiceUnavailable, "orchestration not enabled") return } diff --git a/pkg/orch/broadcaster.go b/pkg/orch/broadcaster.go index 1597a7eca..69a30c524 100644 --- a/pkg/orch/broadcaster.go +++ b/pkg/orch/broadcaster.go @@ -76,11 +76,7 @@ func (b *Broadcaster) Unsubscribe(sub *Subscriber) { func (b *Broadcaster) Snapshot() []AgentInfo { b.mu.Lock() defer b.mu.Unlock() - out := make([]AgentInfo, 0, len(b.agents)) - for _, a := range b.agents { - out = append(out, *a) - } - return out + return b.snapshotLocked() } // ReportSpawn implements AgentReporter. @@ -110,6 +106,35 @@ func (b *Broadcaster) Publish(ev Event) { } b.mu.Lock() + b.applyEventLocked(ev) + subs := b.subscribersLocked() + b.mu.Unlock() + + for _, sub := range subs { + select { + case sub.Ch <- ev: + default: // subscriber slow — drop (non-blocking) + } + } +} + +func (b *Broadcaster) snapshotLocked() []AgentInfo { + out := make([]AgentInfo, 0, len(b.agents)) + for _, a := range b.agents { + out = append(out, *a) + } + return out +} + +func (b *Broadcaster) subscribersLocked() []*Subscriber { + subs := make([]*Subscriber, 0, len(b.subs)) + for sub := range b.subs { + subs = append(subs, sub) + } + return subs +} + +func (b *Broadcaster) applyEventLocked(ev Event) { switch ev.Type { case "agent_spawn": b.agents[ev.ID] = &AgentInfo{ @@ -127,17 +152,4 @@ func (b *Broadcaster) Publish(ev Event) { case "agent_gc": delete(b.agents, ev.ID) } - // snapshot subs while holding lock, then release before sending - subs := make([]*Subscriber, 0, len(b.subs)) - for sub := range b.subs { - subs = append(subs, sub) - } - b.mu.Unlock() - - for _, sub := range subs { - select { - case sub.Ch <- ev: - default: // subscriber slow — drop (non-blocking) - } - } } diff --git a/pkg/research/store.go b/pkg/research/store.go index 073051eba..7bc102d7a 100644 --- a/pkg/research/store.go +++ b/pkg/research/store.go @@ -49,6 +49,10 @@ type ResearchStore struct { workspace string } +func nowRFC3339() string { + return time.Now().UTC().Format(time.RFC3339) +} + // OpenResearchStore opens (or creates) the research SQLite database. func OpenResearchStore(dbPath, workspace string) (*ResearchStore, error) { connStr := "file:" + dbPath + "?_journal_mode=WAL&_foreign_keys=on&_busy_timeout=5000" @@ -254,7 +258,7 @@ func (s *ResearchStore) SetTaskStatus(id string, status TaskStatus) error { return fmt.Errorf("invalid transition: %s → %s", task.Status, status) } - now := time.Now().UTC().Format(time.RFC3339) + now := nowRFC3339() completedAt := "" if status == StatusCompleted || status == StatusFailed { completedAt = now @@ -267,7 +271,7 @@ func (s *ResearchStore) SetTaskStatus(id string, status TaskStatus) error { // UpdateTask updates a task's title and/or description. func (s *ResearchStore) UpdateTask(id, title, description string) error { - now := time.Now().UTC().Format(time.RFC3339) + now := nowRFC3339() _, err := s.db.Exec( `UPDATE research_tasks SET title = ?, description = ?, updated_at = ? WHERE id = ?`, title, description, now, id) @@ -373,7 +377,7 @@ func (s *ResearchStore) SearchTasks(query string) ([]*Task, error) { // TouchLastResearched updates the last_researched_at timestamp for a task. func (s *ResearchStore) TouchLastResearched(taskID string) error { - now := time.Now().UTC().Format(time.RFC3339) + now := nowRFC3339() _, err := s.db.Exec( `UPDATE research_tasks SET last_researched_at = ?, updated_at = ? WHERE id = ?`, now, now, taskID) @@ -385,7 +389,7 @@ func (s *ResearchStore) SetInterval(taskID, interval string) error { if _, err := ParseInterval(interval); err != nil { return fmt.Errorf("invalid interval %q: %w", interval, err) } - now := time.Now().UTC().Format(time.RFC3339) + now := nowRFC3339() _, err := s.db.Exec( `UPDATE research_tasks SET interval = ?, updated_at = ? WHERE id = ?`, interval, now, taskID) @@ -395,24 +399,7 @@ func (s *ResearchStore) SetInterval(taskID, interval string) error { // --- helpers --- func scanTask(row *sql.Row) (*Task, error) { - var t Task - var statusStr, lastResearchedStr, createdStr, updatedStr, completedStr string - err := row.Scan(&t.ID, &t.Title, &t.Slug, &t.Description, &statusStr, - &t.OutputDir, &t.Interval, &lastResearchedStr, &createdStr, &updatedStr, &completedStr) - if err != nil { - return nil, err - } - t.Status = TaskStatus(statusStr) - if t.Interval == "" { - t.Interval = DefaultResearchInterval - } - t.LastResearchedAt, _ = time.Parse(time.RFC3339, lastResearchedStr) - t.CreatedAt, _ = time.Parse(time.RFC3339, createdStr) - t.UpdatedAt, _ = time.Parse(time.RFC3339, updatedStr) - if completedStr != "" { - t.CompletedAt, _ = time.Parse(time.RFC3339, completedStr) - } - return &t, nil + return scanTaskRow(row) } type rowScanner interface {