diff --git a/pkg/config/config.go b/pkg/config/config.go index 7d2b90d32..6b9997937 100644 --- a/pkg/config/config.go +++ b/pkg/config/config.go @@ -271,7 +271,6 @@ type PicoLMProviderConfig struct { Model string `json:"model" env:"PICOCLAW_PROVIDERS_PICOLM_MODEL"` MaxTokens int `json:"max_tokens" env:"PICOCLAW_PROVIDERS_PICOLM_MAX_TOKENS"` Threads int `json:"threads" env:"PICOCLAW_PROVIDERS_PICOLM_THREADS"` - Template string `json:"template" env:"PICOCLAW_PROVIDERS_PICOLM_TEMPLATE"` } type ProviderConfig struct { diff --git a/pkg/providers/factory.go b/pkg/providers/factory.go index bb1695479..c04fb540c 100644 --- a/pkg/providers/factory.go +++ b/pkg/providers/factory.go @@ -361,7 +361,7 @@ func CreateProvider(cfg *config.Config) (LLMProvider, error) { case providerTypeGitHubCopilot: return NewGitHubCopilotProvider(sel.apiBase, sel.connectMode, sel.model) case providerTypePicoLM: - return NewPicoLMProvider(cfg.Providers.PicoLM), nil + return NewPicoLMProvider(cfg.Providers.PicoLM) default: return NewHTTPProvider(sel.apiKey, sel.apiBase, sel.proxy), nil } diff --git a/pkg/providers/picolm_provider.go b/pkg/providers/picolm_provider.go index 45ca2cbda..84da5ab0e 100644 --- a/pkg/providers/picolm_provider.go +++ b/pkg/providers/picolm_provider.go @@ -3,6 +3,7 @@ package providers import ( "bytes" "context" + "encoding/json" "fmt" "os" "os/exec" @@ -17,13 +18,33 @@ type PicoLMProvider struct { model string maxTokens int threads int - template string } // NewPicoLMProvider creates a new PicoLM provider from config. -func NewPicoLMProvider(cfg config.PicoLMProviderConfig) *PicoLMProvider { - binary := expandHome(cfg.Binary) - model := expandHome(cfg.Model) +// Returns an error if the binary path does not point to an executable file. +func NewPicoLMProvider(cfg config.PicoLMProviderConfig) (*PicoLMProvider, error) { + binary, err := expandHome(cfg.Binary) + if err != nil { + return nil, fmt.Errorf("picolm: failed to expand binary path: %w", err) + } + model, err := expandHome(cfg.Model) + if err != nil { + return nil, fmt.Errorf("picolm: failed to expand model path: %w", err) + } + + if binary != "" { + info, err := os.Stat(binary) + if err != nil { + return nil, fmt.Errorf("picolm: binary not found at %q: %w", binary, err) + } + if info.IsDir() { + return nil, fmt.Errorf("picolm: binary path %q is a directory, not an executable", binary) + } + if info.Mode()&0111 == 0 { + return nil, fmt.Errorf("picolm: binary %q is not executable", binary) + } + } + maxTokens := cfg.MaxTokens if maxTokens <= 0 { maxTokens = 256 @@ -32,17 +53,12 @@ func NewPicoLMProvider(cfg config.PicoLMProviderConfig) *PicoLMProvider { if threads <= 0 { threads = 4 } - template := cfg.Template - if template == "" { - template = "chatml" - } return &PicoLMProvider{ binary: binary, model: model, maxTokens: maxTokens, threads: threads, - template: template, - } + }, nil } // Chat implements LLMProvider.Chat by executing the picolm binary. @@ -74,23 +90,6 @@ func (p *PicoLMProvider) Chat(ctx context.Context, messages []Message, tools []T cmd.Stderr = &stderr err := cmd.Run() - - // Try to parse stdout even on non-zero exit, as picolm may print diagnostics to stderr. - if output := strings.TrimSpace(stdout.String()); output != "" { - toolCalls := extractToolCallsFromText(output) - finishReason := "stop" - content := output - if len(toolCalls) > 0 { - finishReason = "tool_calls" - content = stripToolCallsFromText(output) - } - return &LLMResponse{ - Content: strings.TrimSpace(content), - ToolCalls: toolCalls, - FinishReason: finishReason, - }, nil - } - if err != nil { if ctx.Err() == context.Canceled { return nil, ctx.Err() @@ -101,9 +100,25 @@ func (p *PicoLMProvider) Chat(ctx context.Context, messages []Message, tools []T return nil, fmt.Errorf("picolm error: %w", err) } + output := strings.TrimSpace(stdout.String()) + if output == "" { + return &LLMResponse{ + Content: "", + FinishReason: "stop", + }, nil + } + + toolCalls := extractToolCallsFromText(output) + finishReason := "stop" + content := output + if len(toolCalls) > 0 { + finishReason = "tool_calls" + content = stripToolCallsFromText(output) + } return &LLMResponse{ - Content: "", - FinishReason: "stop", + Content: strings.TrimSpace(content), + ToolCalls: toolCalls, + FinishReason: finishReason, }, nil } @@ -116,12 +131,16 @@ func (p *PicoLMProvider) GetDefaultModel() string { func (p *PicoLMProvider) buildPrompt(messages []Message, tools []ToolDefinition) string { var sb strings.Builder + // Collect system message parts and append tool definitions. var systemParts []string for _, msg := range messages { if msg.Role == "system" { systemParts = append(systemParts, msg.Content) } } + if len(tools) > 0 { + systemParts = append(systemParts, buildToolsPrompt(tools)) + } if len(systemParts) > 0 { sb.WriteString("<|system|>\n") @@ -152,17 +171,50 @@ func (p *PicoLMProvider) buildPrompt(messages []Message, tools []ToolDefinition) return sb.String() } +// buildToolsPrompt creates the tool definitions section for injection into the system prompt. +func buildToolsPrompt(tools []ToolDefinition) string { + var sb strings.Builder + + sb.WriteString("## Available Tools\n\n") + sb.WriteString("When you need to use a tool, respond with ONLY a JSON object:\n\n") + sb.WriteString("```json\n") + sb.WriteString(`{"tool_calls":[{"id":"call_xxx","type":"function","function":{"name":"tool_name","arguments":"{...}"}}]}`) + sb.WriteString("\n```\n\n") + sb.WriteString("CRITICAL: The 'arguments' field MUST be a JSON-encoded STRING.\n\n") + sb.WriteString("### Tool Definitions:\n\n") + + for _, tool := range tools { + if tool.Type != "function" { + continue + } + sb.WriteString(fmt.Sprintf("#### %s\n", tool.Function.Name)) + if tool.Function.Description != "" { + sb.WriteString(fmt.Sprintf("Description: %s\n", tool.Function.Description)) + } + if len(tool.Function.Parameters) > 0 { + paramsJSON, _ := json.Marshal(tool.Function.Parameters) + sb.WriteString(fmt.Sprintf("Parameters:\n```json\n%s\n```\n", string(paramsJSON))) + } + sb.WriteString("\n") + } + + return sb.String() +} + // expandHome expands ~ to the user's home directory. -func expandHome(path string) string { +func expandHome(path string) (string, error) { if path == "" { - return path + return path, nil } if path[0] == '~' { - home, _ := os.UserHomeDir() - if len(path) > 1 && path[1] == '/' { - return home + path[1:] + home, err := os.UserHomeDir() + if err != nil { + return "", fmt.Errorf("failed to resolve home directory: %w", err) } - return home + if len(path) > 1 && path[1] == '/' { + return home + path[1:], nil + } + return home, nil } - return path + return path, nil } diff --git a/pkg/providers/picolm_provider_test.go b/pkg/providers/picolm_provider_test.go new file mode 100644 index 000000000..1377c1b6d --- /dev/null +++ b/pkg/providers/picolm_provider_test.go @@ -0,0 +1,676 @@ +package providers + +import ( + "context" + "fmt" + "os" + "path/filepath" + "runtime" + "strings" + "testing" + "time" + + "github.com/sipeed/picoclaw/pkg/config" +) + +// --- Compile-time interface check --- + +var _ LLMProvider = (*PicoLMProvider)(nil) + +// --- Helper: create mock picolm binary --- + +// createMockPicoLM creates a temporary script that simulates the picolm binary. +// It reads stdin (the prompt) and writes stdout/stderr, exiting with the given code. +func createMockPicoLM(t *testing.T, stdout, stderr string, exitCode int) string { + t.Helper() + if runtime.GOOS == "windows" { + t.Skip("mock CLI scripts not supported on Windows") + } + + dir := t.TempDir() + + if stdout != "" { + if err := os.WriteFile(filepath.Join(dir, "stdout.txt"), []byte(stdout), 0644); err != nil { + t.Fatal(err) + } + } + if stderr != "" { + if err := os.WriteFile(filepath.Join(dir, "stderr.txt"), []byte(stderr), 0644); err != nil { + t.Fatal(err) + } + } + + var sb strings.Builder + sb.WriteString("#!/bin/sh\n") + // Consume stdin to avoid broken pipe + sb.WriteString("cat > /dev/null\n") + if stderr != "" { + sb.WriteString(fmt.Sprintf("cat '%s/stderr.txt' >&2\n", dir)) + } + if stdout != "" { + sb.WriteString(fmt.Sprintf("cat '%s/stdout.txt'\n", dir)) + } + sb.WriteString(fmt.Sprintf("exit %d\n", exitCode)) + + script := filepath.Join(dir, "picolm") + if err := os.WriteFile(script, []byte(sb.String()), 0755); err != nil { + t.Fatal(err) + } + return script +} + +// createStdinCaptureMockPicoLM creates a mock that captures stdin to a file. +func createStdinCaptureMockPicoLM(t *testing.T, captureFile string, stdout string) string { + t.Helper() + if runtime.GOOS == "windows" { + t.Skip("mock CLI scripts not supported on Windows") + } + + dir := t.TempDir() + if stdout != "" { + if err := os.WriteFile(filepath.Join(dir, "stdout.txt"), []byte(stdout), 0644); err != nil { + t.Fatal(err) + } + } + + var sb strings.Builder + sb.WriteString("#!/bin/sh\n") + sb.WriteString(fmt.Sprintf("cat > '%s'\n", captureFile)) + if stdout != "" { + sb.WriteString(fmt.Sprintf("cat '%s/stdout.txt'\n", dir)) + } + sb.WriteString("exit 0\n") + + script := filepath.Join(dir, "picolm") + if err := os.WriteFile(script, []byte(sb.String()), 0755); err != nil { + t.Fatal(err) + } + return script +} + +// createSlowMockPicoLM creates a script that sleeps before responding. +func createSlowMockPicoLM(t *testing.T, sleepSeconds int) string { + t.Helper() + if runtime.GOOS == "windows" { + t.Skip("mock CLI scripts not supported on Windows") + } + + dir := t.TempDir() + script := filepath.Join(dir, "picolm") + content := fmt.Sprintf("#!/bin/sh\ncat > /dev/null\nsleep %d\necho 'late response'\n", sleepSeconds) + if err := os.WriteFile(script, []byte(content), 0755); err != nil { + t.Fatal(err) + } + return script +} + +// --- Constructor tests --- + +func TestNewPicoLMProvider(t *testing.T) { + binary := createMockPicoLM(t, "", "", 0) + p, err := NewPicoLMProvider(config.PicoLMProviderConfig{ + Binary: binary, + Model: "/tmp/model.gguf", + MaxTokens: 128, + Threads: 2, + }) + if err != nil { + t.Fatalf("NewPicoLMProvider() error = %v", err) + } + if p.binary != binary { + t.Errorf("binary = %q, want %q", p.binary, binary) + } + if p.maxTokens != 128 { + t.Errorf("maxTokens = %d, want 128", p.maxTokens) + } + if p.threads != 2 { + t.Errorf("threads = %d, want 2", p.threads) + } +} + +func TestNewPicoLMProvider_Defaults(t *testing.T) { + binary := createMockPicoLM(t, "", "", 0) + p, err := NewPicoLMProvider(config.PicoLMProviderConfig{ + Binary: binary, + Model: "/tmp/model.gguf", + }) + if err != nil { + t.Fatalf("NewPicoLMProvider() error = %v", err) + } + if p.maxTokens != 256 { + t.Errorf("maxTokens = %d, want 256 (default)", p.maxTokens) + } + if p.threads != 4 { + t.Errorf("threads = %d, want 4 (default)", p.threads) + } +} + +func TestNewPicoLMProvider_BinaryNotFound(t *testing.T) { + _, err := NewPicoLMProvider(config.PicoLMProviderConfig{ + Binary: "/nonexistent/path/picolm", + Model: "/tmp/model.gguf", + }) + if err == nil { + t.Fatal("expected error for nonexistent binary") + } + if !strings.Contains(err.Error(), "binary not found") { + t.Errorf("error = %q, want to contain 'binary not found'", err.Error()) + } +} + +func TestNewPicoLMProvider_BinaryIsDirectory(t *testing.T) { + dir := t.TempDir() + _, err := NewPicoLMProvider(config.PicoLMProviderConfig{ + Binary: dir, + Model: "/tmp/model.gguf", + }) + if err == nil { + t.Fatal("expected error when binary is a directory") + } + if !strings.Contains(err.Error(), "is a directory") { + t.Errorf("error = %q, want to contain 'is a directory'", err.Error()) + } +} + +func TestNewPicoLMProvider_BinaryNotExecutable(t *testing.T) { + if runtime.GOOS == "windows" { + t.Skip("executable bit not meaningful on Windows") + } + dir := t.TempDir() + notExec := filepath.Join(dir, "picolm") + if err := os.WriteFile(notExec, []byte("not executable"), 0644); err != nil { + t.Fatal(err) + } + _, err := NewPicoLMProvider(config.PicoLMProviderConfig{ + Binary: notExec, + Model: "/tmp/model.gguf", + }) + if err == nil { + t.Fatal("expected error for non-executable binary") + } + if !strings.Contains(err.Error(), "not executable") { + t.Errorf("error = %q, want to contain 'not executable'", err.Error()) + } +} + +func TestNewPicoLMProvider_EmptyBinaryAllowed(t *testing.T) { + // Empty binary is allowed at construction time (caught at Chat time) + p, err := NewPicoLMProvider(config.PicoLMProviderConfig{ + Model: "/tmp/model.gguf", + }) + if err != nil { + t.Fatalf("NewPicoLMProvider() error = %v", err) + } + if p.binary != "" { + t.Errorf("binary = %q, want empty", p.binary) + } +} + +// --- GetDefaultModel tests --- + +func TestPicoLMProvider_GetDefaultModel(t *testing.T) { + binary := createMockPicoLM(t, "", "", 0) + p, _ := NewPicoLMProvider(config.PicoLMProviderConfig{Binary: binary}) + if got := p.GetDefaultModel(); got != "picolm-local" { + t.Errorf("GetDefaultModel() = %q, want %q", got, "picolm-local") + } +} + +// --- Chat() tests --- + +func TestPicoLMChat_Success(t *testing.T) { + script := createMockPicoLM(t, "Photosynthesis is the process by which plants convert sunlight.", "", 0) + + p := &PicoLMProvider{ + binary: script, + model: "/tmp/model.gguf", + maxTokens: 256, + threads: 4, + } + + resp, err := p.Chat(context.Background(), []Message{ + {Role: "user", Content: "What is photosynthesis?"}, + }, nil, "", nil) + + if err != nil { + t.Fatalf("Chat() error = %v", err) + } + if resp.Content != "Photosynthesis is the process by which plants convert sunlight." { + t.Errorf("Content = %q", resp.Content) + } + if resp.FinishReason != "stop" { + t.Errorf("FinishReason = %q, want %q", resp.FinishReason, "stop") + } + if len(resp.ToolCalls) != 0 { + t.Errorf("ToolCalls len = %d, want 0", len(resp.ToolCalls)) + } +} + +func TestPicoLMChat_WithToolCallsInResponse(t *testing.T) { + mockOutput := `Checking weather. +{"tool_calls":[{"id":"call_1","type":"function","function":{"name":"get_weather","arguments":"{\"location\":\"NYC\"}"}}]}` + script := createMockPicoLM(t, mockOutput, "", 0) + + p := &PicoLMProvider{ + binary: script, + model: "/tmp/model.gguf", + maxTokens: 256, + threads: 4, + } + + resp, err := p.Chat(context.Background(), []Message{ + {Role: "user", Content: "What's the weather?"}, + }, []ToolDefinition{{ + Type: "function", + Function: ToolFunctionDefinition{ + Name: "get_weather", + Description: "Get weather", + Parameters: map[string]interface{}{"type": "object"}, + }, + }}, "", nil) + + if err != nil { + t.Fatalf("Chat() error = %v", err) + } + if resp.FinishReason != "tool_calls" { + t.Errorf("FinishReason = %q, want %q", resp.FinishReason, "tool_calls") + } + if len(resp.ToolCalls) != 1 { + t.Fatalf("ToolCalls len = %d, want 1", len(resp.ToolCalls)) + } + if resp.ToolCalls[0].Name != "get_weather" { + t.Errorf("ToolCalls[0].Name = %q, want %q", resp.ToolCalls[0].Name, "get_weather") + } + if resp.ToolCalls[0].Arguments["location"] != "NYC" { + t.Errorf("ToolCalls[0].Arguments[location] = %v, want NYC", resp.ToolCalls[0].Arguments["location"]) + } +} + +func TestPicoLMChat_EmptyOutput(t *testing.T) { + script := createMockPicoLM(t, "", "", 0) + + p := &PicoLMProvider{ + binary: script, + model: "/tmp/model.gguf", + maxTokens: 256, + threads: 4, + } + + resp, err := p.Chat(context.Background(), []Message{ + {Role: "user", Content: "Hello"}, + }, nil, "", nil) + + if err != nil { + t.Fatalf("Chat() error = %v", err) + } + if resp.Content != "" { + t.Errorf("Content = %q, want empty", resp.Content) + } + if resp.FinishReason != "stop" { + t.Errorf("FinishReason = %q, want %q", resp.FinishReason, "stop") + } +} + +func TestPicoLMChat_StderrError(t *testing.T) { + script := createMockPicoLM(t, "", "model load failed: out of memory", 1) + + p := &PicoLMProvider{ + binary: script, + model: "/tmp/model.gguf", + maxTokens: 256, + threads: 4, + } + + _, err := p.Chat(context.Background(), []Message{ + {Role: "user", Content: "Hello"}, + }, nil, "", nil) + + if err == nil { + t.Fatal("Chat() expected error") + } + if !strings.Contains(err.Error(), "out of memory") { + t.Errorf("error = %q, want to contain 'out of memory'", err.Error()) + } +} + +func TestPicoLMChat_NonZeroExitNoStderr(t *testing.T) { + script := createMockPicoLM(t, "", "", 1) + + p := &PicoLMProvider{ + binary: script, + model: "/tmp/model.gguf", + maxTokens: 256, + threads: 4, + } + + _, err := p.Chat(context.Background(), []Message{ + {Role: "user", Content: "Hello"}, + }, nil, "", nil) + + if err == nil { + t.Fatal("Chat() expected error for non-zero exit") + } + if !strings.Contains(err.Error(), "picolm error") { + t.Errorf("error = %q, want to contain 'picolm error'", err.Error()) + } +} + +func TestPicoLMChat_NonZeroExitWithStdout(t *testing.T) { + // When the process fails with non-zero exit, stdout should NOT be treated as valid output. + script := createMockPicoLM(t, "partial garbage output", "segfault", 139) + + p := &PicoLMProvider{ + binary: script, + model: "/tmp/model.gguf", + maxTokens: 256, + threads: 4, + } + + _, err := p.Chat(context.Background(), []Message{ + {Role: "user", Content: "Hello"}, + }, nil, "", nil) + + if err == nil { + t.Fatal("Chat() expected error when process exits non-zero, even with stdout") + } +} + +func TestPicoLMChat_ContextCancellation(t *testing.T) { + script := createSlowMockPicoLM(t, 30) + + p := &PicoLMProvider{ + binary: script, + model: "/tmp/model.gguf", + maxTokens: 256, + threads: 4, + } + + ctx, cancel := context.WithTimeout(context.Background(), 200*time.Millisecond) + defer cancel() + + _, err := p.Chat(ctx, []Message{ + {Role: "user", Content: "Hello"}, + }, nil, "", nil) + + if err == nil { + t.Fatal("Chat() expected error on context cancellation") + } +} + +func TestPicoLMChat_NoBinary(t *testing.T) { + p := &PicoLMProvider{model: "/tmp/model.gguf"} + _, err := p.Chat(context.Background(), []Message{ + {Role: "user", Content: "Hello"}, + }, nil, "", nil) + if err == nil || !strings.Contains(err.Error(), "binary path not configured") { + t.Errorf("expected 'binary path not configured' error, got: %v", err) + } +} + +func TestPicoLMChat_NoModel(t *testing.T) { + p := &PicoLMProvider{binary: "/usr/bin/echo"} + _, err := p.Chat(context.Background(), []Message{ + {Role: "user", Content: "Hello"}, + }, nil, "", nil) + if err == nil || !strings.Contains(err.Error(), "model path not configured") { + t.Errorf("expected 'model path not configured' error, got: %v", err) + } +} + +// --- buildPrompt tests --- + +func TestPicoLMBuildPrompt_SimpleUserMessage(t *testing.T) { + p := &PicoLMProvider{} + prompt := p.buildPrompt([]Message{ + {Role: "user", Content: "Hello"}, + }, nil) + + if !strings.Contains(prompt, "<|user|>\nHello") { + t.Errorf("prompt missing user message, got:\n%s", prompt) + } + if !strings.HasSuffix(prompt, "<|assistant|>\n") { + t.Errorf("prompt should end with assistant turn, got:\n%s", prompt) + } +} + +func TestPicoLMBuildPrompt_SystemAndUser(t *testing.T) { + p := &PicoLMProvider{} + prompt := p.buildPrompt([]Message{ + {Role: "system", Content: "You are helpful."}, + {Role: "user", Content: "Hello"}, + }, nil) + + if !strings.Contains(prompt, "<|system|>\nYou are helpful.") { + t.Errorf("prompt missing system message, got:\n%s", prompt) + } + if !strings.Contains(prompt, "<|user|>\nHello") { + t.Errorf("prompt missing user message, got:\n%s", prompt) + } +} + +func TestPicoLMBuildPrompt_MultipleSystemMessages(t *testing.T) { + p := &PicoLMProvider{} + prompt := p.buildPrompt([]Message{ + {Role: "system", Content: "Part one."}, + {Role: "system", Content: "Part two."}, + {Role: "user", Content: "Hello"}, + }, nil) + + if !strings.Contains(prompt, "Part one.\n\nPart two.") { + t.Errorf("system parts should be joined, got:\n%s", prompt) + } + // Should have exactly one <|system|> block + if strings.Count(prompt, "<|system|>") != 1 { + t.Errorf("expected exactly 1 system block, got %d", strings.Count(prompt, "<|system|>")) + } +} + +func TestPicoLMBuildPrompt_AssistantMessage(t *testing.T) { + p := &PicoLMProvider{} + prompt := p.buildPrompt([]Message{ + {Role: "user", Content: "Hi"}, + {Role: "assistant", Content: "Hello!"}, + {Role: "user", Content: "How are you?"}, + }, nil) + + if !strings.Contains(prompt, "<|assistant|>\nHello!") { + t.Errorf("prompt missing assistant message, got:\n%s", prompt) + } +} + +func TestPicoLMBuildPrompt_ToolResult(t *testing.T) { + p := &PicoLMProvider{} + prompt := p.buildPrompt([]Message{ + {Role: "user", Content: "What's the weather?"}, + {Role: "tool", ToolCallID: "call_1", Content: `{"temp": 72}`}, + }, nil) + + if !strings.Contains(prompt, "[Tool Result for call_1]: {\"temp\": 72}") { + t.Errorf("prompt missing tool result, got:\n%s", prompt) + } + // Tool results are wrapped in user tags + if !strings.Contains(prompt, "<|user|>\n[Tool Result for call_1]") { + t.Errorf("tool result should be in user block, got:\n%s", prompt) + } +} + +func TestPicoLMBuildPrompt_WithTools(t *testing.T) { + p := &PicoLMProvider{} + tools := []ToolDefinition{ + { + Type: "function", + Function: ToolFunctionDefinition{ + Name: "get_weather", + Description: "Get weather for a city", + Parameters: map[string]interface{}{ + "type": "object", + "properties": map[string]interface{}{ + "city": map[string]interface{}{"type": "string"}, + }, + "required": []interface{}{"city"}, + }, + }, + }, + } + + prompt := p.buildPrompt([]Message{ + {Role: "system", Content: "You are helpful."}, + {Role: "user", Content: "What's the weather in NYC?"}, + }, tools) + + // Tool definitions should be in the system block + if !strings.Contains(prompt, "## Available Tools") { + t.Errorf("prompt missing tool definitions header, got:\n%s", prompt) + } + if !strings.Contains(prompt, "#### get_weather") { + t.Errorf("prompt missing tool name, got:\n%s", prompt) + } + if !strings.Contains(prompt, "Get weather for a city") { + t.Errorf("prompt missing tool description, got:\n%s", prompt) + } + if !strings.Contains(prompt, `"tool_calls"`) { + t.Errorf("prompt missing tool call format example, got:\n%s", prompt) + } + // Tool definitions should be part of the system block + systemIdx := strings.Index(prompt, "<|system|>") + systemEnd := strings.Index(prompt, "") + toolsIdx := strings.Index(prompt, "## Available Tools") + if toolsIdx < systemIdx || toolsIdx > systemEnd { + t.Errorf("tool definitions should be inside system block") + } +} + +func TestPicoLMBuildPrompt_WithTools_NoSystem(t *testing.T) { + p := &PicoLMProvider{} + tools := []ToolDefinition{ + { + Type: "function", + Function: ToolFunctionDefinition{ + Name: "search", + Description: "Search the web", + }, + }, + } + + prompt := p.buildPrompt([]Message{ + {Role: "user", Content: "Search for cats"}, + }, tools) + + // Should still create a system block for tools even without system messages + if !strings.Contains(prompt, "<|system|>") { + t.Errorf("expected system block for tools, got:\n%s", prompt) + } + if !strings.Contains(prompt, "#### search") { + t.Errorf("prompt missing tool name, got:\n%s", prompt) + } +} + +func TestPicoLMBuildPrompt_ToolsFilterNonFunction(t *testing.T) { + p := &PicoLMProvider{} + tools := []ToolDefinition{ + { + Type: "not_a_function", + Function: ToolFunctionDefinition{ + Name: "should_be_skipped", + }, + }, + { + Type: "function", + Function: ToolFunctionDefinition{ + Name: "included", + Description: "This one is included", + }, + }, + } + + prompt := p.buildPrompt([]Message{ + {Role: "user", Content: "Hello"}, + }, tools) + + if strings.Contains(prompt, "should_be_skipped") { + t.Errorf("non-function tools should be filtered out") + } + if !strings.Contains(prompt, "#### included") { + t.Errorf("function tools should be included") + } +} + +// --- Stdin prompt verification --- + +func TestPicoLMChat_PromptSentViaStdin(t *testing.T) { + captureFile := filepath.Join(t.TempDir(), "captured_stdin.txt") + script := createStdinCaptureMockPicoLM(t, captureFile, "Response from model") + + p := &PicoLMProvider{ + binary: script, + model: "/tmp/model.gguf", + maxTokens: 256, + threads: 4, + } + + _, err := p.Chat(context.Background(), []Message{ + {Role: "user", Content: "What is 2+2?"}, + }, nil, "", nil) + + if err != nil { + t.Fatalf("Chat() error = %v", err) + } + + captured, err := os.ReadFile(captureFile) + if err != nil { + t.Fatalf("failed to read captured stdin: %v", err) + } + + stdinContent := string(captured) + if !strings.Contains(stdinContent, "What is 2+2?") { + t.Errorf("stdin should contain user message, got:\n%s", stdinContent) + } + if !strings.Contains(stdinContent, "<|user|>") { + t.Errorf("stdin should contain ChatML tags, got:\n%s", stdinContent) + } +} + +// --- expandHome tests --- + +func TestExpandHome_Empty(t *testing.T) { + result, err := expandHome("") + if err != nil { + t.Fatalf("expandHome(\"\") error = %v", err) + } + if result != "" { + t.Errorf("expandHome(\"\") = %q, want empty", result) + } +} + +func TestExpandHome_NoTilde(t *testing.T) { + result, err := expandHome("/usr/bin/picolm") + if err != nil { + t.Fatalf("expandHome error = %v", err) + } + if result != "/usr/bin/picolm" { + t.Errorf("expandHome = %q, want %q", result, "/usr/bin/picolm") + } +} + +func TestExpandHome_WithTilde(t *testing.T) { + result, err := expandHome("~/bin/picolm") + if err != nil { + t.Fatalf("expandHome error = %v", err) + } + home, _ := os.UserHomeDir() + expected := home + "/bin/picolm" + if result != expected { + t.Errorf("expandHome = %q, want %q", result, expected) + } +} + +func TestExpandHome_TildeOnly(t *testing.T) { + result, err := expandHome("~") + if err != nil { + t.Fatalf("expandHome error = %v", err) + } + home, _ := os.UserHomeDir() + if result != home { + t.Errorf("expandHome = %q, want %q", result, home) + } +}