package main import ( "encoding/base64" "encoding/json" "net/http" "net/http/httptest" "os" "path/filepath" "strings" "testing" ) func TestResearchReqJSONUsesSeparateFields(t *testing.T) { raw := `{"provider":"above","query":"research something","maxNewTok":1024}` var r ResearchReq if err := json.Unmarshal([]byte(raw), &r); err != nil { t.Fatal(err) } if r.Provider != "above" || r.Query != "research something" || r.MaxNewTok != 1024 { t.Fatalf("decoded research request incorrectly: %+v", r) } } func TestParseParallelMarkdownResults(t *testing.T) { text := "# Results\n[OpenAI docs](https://example.com/docs)\n[Parallel docs](https://example.com/parallel)" rows := parseParallelMarkdownResults(text) if len(rows) != 2 { t.Fatalf("expected 2 rows, got %d", len(rows)) } if rows[0]["url"] != "https://example.com/docs" { t.Fatalf("unexpected first URL: %v", rows[0]["url"]) } } func TestTerminalDetailsHasShell(t *testing.T) { d := terminalDetails() if d["shell"] == "" || d["display"] == "" { t.Fatalf("terminal details missing shell/display: %+v", d) } } func TestRunTerminalCommand(t *testing.T) { d := terminalDetails() cwd, _ := d["cwd"].(string) if cwd == "" { t.Fatal("expected default terminal cwd") } res := runTerminalCommand(TerminalReq{Command: "printf 'agentdesk-terminal-ok'", CWD: cwd, TimeoutSecond: 5}) if ok, _ := res["ok"].(bool); !ok { t.Fatalf("terminal command failed: %+v", res) } if out, _ := res["output"].(string); out != "agentdesk-terminal-ok" { t.Fatalf("unexpected terminal output: %q", out) } } func TestSafeArtifactName(t *testing.T) { got := safeArtifactName("../demo:bad?.py", "file.txt") if got == "../demo:bad?.py" || got == "" { t.Fatalf("artifact name was not sanitized: %q", got) } } func TestPrepareProviderMessagesResolvesImageRef(t *testing.T) { root := t.TempDir() id := "test-image" img := []byte{0x89, 'P', 'N', 'G', '\r', '\n', 0x1a, '\n'} if err := os.WriteFile(filepath.Join(root, id+".png"), img, 0600); err != nil { t.Fatal(err) } msgs := []map[string]any{{ "role": "user", "content": []any{ map[string]any{"type": "text", "text": "Inspect this."}, map[string]any{"type": "image_ref", "image_ref": map[string]any{"id": id, "mime": "image/png", "name": "misleading.png"}}, }, }} prepared, err := prepareProviderMessages(msgs, root) if err != nil { t.Fatal(err) } parts, ok := prepared[0]["content"].([]any) if !ok || len(parts) != 2 { t.Fatalf("unexpected prepared content: %#v", prepared[0]["content"]) } part, ok := parts[1].(map[string]any) if !ok || part["type"] != "image_url" { t.Fatalf("image_ref was not converted: %#v", parts[1]) } imgURL, ok := part["image_url"].(map[string]any) if !ok { t.Fatalf("missing image_url payload: %#v", part) } dataURL, _ := imgURL["url"].(string) if !strings.HasPrefix(dataURL, "data:image/png;base64,") { t.Fatalf("unexpected data URL: %q", dataURL) } encoded := strings.TrimPrefix(dataURL, "data:image/png;base64,") decoded, err := base64.StdEncoding.DecodeString(encoded) if err != nil || string(decoded) != string(img) { t.Fatalf("decoded image differs: err=%v decoded=%v", err, decoded) } } func TestConfigGETDoesNotEraseAPIKey(t *testing.T) { orig := loadConfig _ = orig } func TestAboveHostDetection(t *testing.T) { if !isAboveHost("https://api.above.dev/v1") { t.Fatal("expected above.dev host to be detected") } if isAboveHost("https://api.openrouter.ai/api/v1") { t.Fatal("did not expect OpenRouter to be detected as above.dev") } } func TestFilterAboveModelsWithoutKeyLeavesCatalogUntouched(t *testing.T) { models := []map[string]any{{"id": "glm-5.3-flash-modal"}, {"id": "deepseek-v4.1-flash-modal"}} got, err := filterAboveModels("https://api.above.dev/v1", "", models) if err != nil { t.Fatal(err) } if len(got) != 2 { t.Fatalf("expected catalog untouched without key, got %d", len(got)) } } func TestConfigGETPreservesStoredAPIKey(t *testing.T) { t.Setenv("XDG_CONFIG_HOME", t.TempDir()) payload := `{"providers":[{"id":"test","name":"Test Provider","baseUrl":"https://example.com/v1","model":"test-model","apiKey":"secret-key","maxNewTok":1024}]}` req := httptest.NewRequest("POST", "/api/config", strings.NewReader(payload)) req.Header.Set("Content-Type", "application/json") rec := httptest.NewRecorder() config(rec, req) if rec.Code != 200 { t.Fatalf("POST config failed: %d %s", rec.Code, rec.Body.String()) } rec = httptest.NewRecorder() req = httptest.NewRequest("GET", "/api/config", nil) config(rec, req) if rec.Code != 200 { t.Fatalf("GET config failed: %d %s", rec.Code, rec.Body.String()) } stored := loadConfig() if len(stored.Providers) != 1 || stored.Providers[0].APIKey != "secret-key" { t.Fatalf("GET config erased the stored API key: %+v", stored.Providers) } } func TestIncompleteProviderCanBeSavedForEditing(t *testing.T) { t.Setenv("XDG_CONFIG_HOME", t.TempDir()) payload := `{"providers":[{"id":"draft","name":"Draft Provider","baseUrl":"","model":"","apiKey":""}]}` req := httptest.NewRequest("POST", "/api/config", strings.NewReader(payload)) req.Header.Set("Content-Type", "application/json") rec := httptest.NewRecorder() config(rec, req) if rec.Code != 200 { t.Fatalf("draft provider POST failed: %d %s", rec.Code, rec.Body.String()) } stored := loadConfig() if len(stored.Providers) != 1 || stored.Providers[0].ID != "draft" { t.Fatalf("incomplete provider was not persisted for editing: %+v", stored.Providers) } } func TestResponseChoiceMetaLengthAndUsage(t *testing.T) { out := map[string]any{ "choices": []any{map[string]any{ "finish_reason": "length", "message": map[string]any{"content": "partial"}, }}, "usage": map[string]any{"completion_tokens": float64(16384)}, } msg, finish, used := responseChoiceMeta(out) if msg == nil || finish != "length" || used != 16384 { t.Fatalf("unexpected meta: msg=%v finish=%q used=%d", msg, finish, used) } } func TestCitationTrackerStableIDs(t *testing.T) { c := newCitationTracker() if got := c.Register("Example", "https://example.com"); got != "S1" { t.Fatalf("expected S1, got %q", got) } if got := c.Register("Example", "https://example.com"); got != "S1" { t.Fatalf("expected duplicate URL to keep S1, got %q", got) } if got := c.Register("Other", "https://example.org"); got != "S2" { t.Fatalf("expected S2, got %q", got) } if len(c.Sources) != 2 { t.Fatalf("expected 2 sources, got %d", len(c.Sources)) } } func TestManualContinuationDetection(t *testing.T) { for _, text := range []string{"continue", "Please continue.", " keep going! ", "continue the response"} { if !isManualContinueRequest(text) { t.Fatalf("expected %q to be detected as a continuation request", text) } } for _, text := range []string{"continue this project", "tell me more about it", "what about next?"} { if isManualContinueRequest(text) { t.Fatalf("unexpected continuation detection for %q", text) } } } func TestManualContinuationReplacesUserMessage(t *testing.T) { raw := []map[string]any{ {"role": "user", "content": "Write a report."}, {"role": "assistant", "content": "The report begins here and ends mid-sentence"}, {"role": "user", "content": "Please continue."}, } msgs, manual := prepareManualContinuationMessages(raw, nil) if !manual { t.Fatal("expected manual continuation to be detected") } if len(msgs) < 3 { t.Fatalf("expected memory + conversation + continuation messages, got %d", len(msgs)) } last := msgs[len(msgs)-1] if last["role"] != "user" || !strings.Contains(last["content"].(string), "Continue immediately") { t.Fatalf("unexpected hidden continuation message: %#v", last) } for _, m := range msgs { if m["content"] == "Please continue." { t.Fatal("raw manual continuation message leaked into provider context") } } } func TestMessageTextHandlesContentArrays(t *testing.T) { msg := map[string]any{"content": []any{ map[string]any{"type": "text", "text": "Hello "}, map[string]any{"type": "text", "text": "world"}, }} got := messageText(msg["content"]) if got != "Hello world" { t.Fatalf("unexpected message text: %q", got) } } func TestRoutineCodingRequestDoesNotRequireWebResearch(t *testing.T) { if !routineCodingRequest("Write a simple Python script") { t.Fatal("expected coding request to be detected") } if shouldOfferWebTools("Write a simple Python script", true) { t.Fatal("routine coding request should not automatically enable web research") } if !shouldOfferWebTools("What is the current Python version?", true) { t.Fatal("current-information request should enable web research") } } func TestResolveMarketSymbol(t *testing.T) { if got := resolveMarketSymbol("show me amazon stocks"); got != "AMZN" { t.Fatalf("expected AMZN, got %q", got) } if got := resolveMarketSymbol("what is $NVDA doing today?"); got != "NVDA" { t.Fatalf("expected NVDA, got %q", got) } } func TestMarketFollowupResolution(t *testing.T) { if got := resolveMarketSymbol("what about roblox?"); got != "RBLX" { t.Fatalf("expected RBLX, got %q", got) } if !isMarketQuery("what about roblox?") { t.Fatal("expected Roblox follow-up to be treated as a market query") } } func TestRegularResearchSearchBudgetConstant(t *testing.T) { // This mirrors the chatStream guard: regular chat research is capped at five web searches. const maxRegularWebSearches = 5 if maxRegularWebSearches != 5 { t.Fatalf("unexpected regular research search limit: %d", maxRegularWebSearches) } } func TestModelMaxCompletionTokensUsesModelCatalog(t *testing.T) { srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { if r.URL.Path != "/models" { http.NotFound(w, r) return } w.Header().Set("Content-Type", "application/json") _, _ = w.Write([]byte(`{"data":[{"id":"demo-64k","max_completion_tokens":65536},{"id":"demo-16k","max_completion_tokens":16384}]}`)) })) defer srv.Close() if got := modelMaxCompletionTokens(Provider{BaseURL: srv.URL, Model: "demo-64k"}); got != 65536 { t.Fatalf("expected 65536, got %d", got) } if got := modelMaxCompletionTokens(Provider{BaseURL: srv.URL, Model: "demo-16k"}); got != 16384 { t.Fatalf("expected 16384, got %d", got) } } func TestDeepResearchBudgetRespectsDetectedModelMaximum(t *testing.T) { srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { if r.URL.Path != "/models" { http.NotFound(w, r) return } w.Header().Set("Content-Type", "application/json") _, _ = w.Write([]byte(`{"data":[{"id":"demo-16k","max_completion_tokens":16384}]}`)) })) defer srv.Close() p := Provider{BaseURL: srv.URL, Model: "demo-16k", MaxNewTok: 130000} eff, reserve, searchBudget, passes, detected := deepResearchBudget(p, 130000) if detected != 16384 || eff != 16384 || reserve != 4000 || searchBudget != 12384 || passes != 24 { t.Fatalf("unexpected detected-model budget: detected=%d eff=%d reserve=%d search=%d passes=%d", detected, eff, reserve, searchBudget, passes) } } func TestDeepResearchBudgetScalesWithConfiguredMax(t *testing.T) { p := Provider{MaxNewTok: 64000} eff, reserve, searchBudget, passes, _ := deepResearchBudget(p, 64000) if eff != 64000 { t.Fatalf("expected effective 64000, got %d", eff) } if reserve != 5000 { t.Fatalf("expected 5000-token report reserve, got %d", reserve) } if searchBudget != 59000 { t.Fatalf("expected 59000 search budget, got %d", searchBudget) } if passes < 100 { t.Fatalf("expected 100+ search passes at 64k, got %d", passes) } } func TestDeepResearchBudgetUsesLowerModelMaximum(t *testing.T) { p := Provider{MaxNewTok: 130000} // A model catalog failure falls back to the configured provider maximum. eff, reserve, _, passes, modelMax := deepResearchBudget(p, 130000) if modelMax != 0 { t.Fatalf("expected unavailable model catalog to return 0, got %d", modelMax) } if eff != 130000 || reserve != 5000 || passes != 250 { t.Fatalf("unexpected fallback budget: eff=%d reserve=%d passes=%d", eff, reserve, passes) } } func TestDeepResearchPlanQueriesCanScalePastTwentyFive(t *testing.T) { qs := deepResearchPlanQueries("MCP ecosystem", 120) if len(qs) != 120 { t.Fatalf("expected 120 queries, got %d", len(qs)) } seen := map[string]bool{} for _, q := range qs { k := strings.ToLower(q) if seen[k] { t.Fatalf("duplicate research query: %q", q) } seen[k] = true } } func TestMemorySearchFindsPersistentContext(t *testing.T) { t.Setenv("XDG_CONFIG_HOME", t.TempDir()) if _, _, _, err := upsertMemories([]Memory{{Text:"Prefers concise TypeScript examples",Kind:"preference"}}); err != nil { t.Fatal(err) } got := memorySearch("TypeScript", 4) if len(got) != 1 || !strings.Contains(got[0].Text, "TypeScript") { t.Fatalf("memory search failed: %+v", got) } }