mirror of
https://github.com/ollama/ollama.git
synced 2026-09-08 12:13:43 -04:00
Compare commits
115
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
b0b7bf26f3 | ||
|
|
90dd5b3a70 | ||
|
|
4855c61358 | ||
|
|
573386c35e | ||
|
|
794a254111 | ||
|
|
714b6fc2a4 | ||
|
|
61e1b1ba5e | ||
|
|
5865a01e48 | ||
|
|
e61c1c73fe | ||
|
|
03d61e1925 | ||
|
|
30c390384e | ||
|
|
d590830091 | ||
|
|
fdcf9efafd | ||
|
|
76188f60cd | ||
|
|
8a0016f826 | ||
|
|
d49b96d9ab | ||
|
|
3bd506bd1c | ||
|
|
123b1f2479 | ||
|
|
556245843a | ||
|
|
c963822dca | ||
|
|
dd49563d55 | ||
|
|
4e96f4dbf2 | ||
|
|
d573a2367b | ||
|
|
4f7786d0ba | ||
|
|
f1a0ffd621 | ||
|
|
cd600e19a3 | ||
|
|
59bd0b49bb | ||
|
|
82f905cd9c | ||
|
|
cb3d98ccb2 | ||
|
|
d47859ce49 | ||
|
|
a6293eb516 | ||
|
|
892e7f6be6 | ||
|
|
f3d69a3dee | ||
|
|
67b6a1c2d4 | ||
|
|
87b64213b4 | ||
|
|
f2d069f6df | ||
|
|
5208ae7500 | ||
|
|
9d779572a7 | ||
|
|
964ea42c09 | ||
|
|
dba1e27fa8 | ||
|
|
e436db25ff | ||
|
|
26acfa42b5 | ||
|
|
7b22ac9683 | ||
|
|
a2b3a5e9a3 | ||
|
|
624cada952 | ||
|
|
cecd265d3a | ||
|
|
2ea95fb059 | ||
|
|
8e7be3aed1 | ||
|
|
710292ff4f | ||
|
|
ada1eb5163 | ||
|
|
1c5ebbf5f4 | ||
|
|
7926b99e0e | ||
|
|
32a97b7493 | ||
|
|
d26a58557d | ||
|
|
2e474c98f9 | ||
|
|
2cb2c5381f | ||
|
|
2a6b50421a | ||
|
|
f22ec2ec49 | ||
|
|
d9075caf1a | ||
|
|
e11eeb3ba0 | ||
|
|
0a408b2225 | ||
|
|
16739dee60 | ||
|
|
d48d790baf | ||
|
|
0463940334 | ||
|
|
570679c9e0 | ||
|
|
89a171cc70 | ||
|
|
33878e671a | ||
|
|
c191a145bb | ||
|
|
479e1cf94e | ||
|
|
836507378b | ||
|
|
46bc1bcb4c | ||
|
|
2a8b31531e | ||
|
|
505e35f2b9 | ||
|
|
114875133b | ||
|
|
42c330283b | ||
|
|
f93efe2809 | ||
|
|
28fbbb06d5 | ||
|
|
340c51bbb7 | ||
|
|
2e9d68dc38 | ||
|
|
fc58544422 | ||
|
|
e434a93884 | ||
|
|
9c02d8e69d | ||
|
|
07ed752353 | ||
|
|
e1f7f9cbdb | ||
|
|
8c432fc88a | ||
|
|
acfb50d9af | ||
|
|
0f047feef5 | ||
|
|
9e4ed74efe | ||
|
|
bbb40a0a6c | ||
|
|
993acc7504 | ||
|
|
7ea692cb2b | ||
|
|
12e04379cd | ||
|
|
f8a48df24d | ||
|
|
82e0ddb6fe | ||
|
|
1abd56b6e6 | ||
|
|
ded2db7d86 | ||
|
|
d00622060f | ||
|
|
177aefb8a9 | ||
|
|
07588c64ee | ||
|
|
4c97a940ca | ||
|
|
74cbf1d2c2 | ||
|
|
5c1e37eb67 | ||
|
|
f0078ae476 | ||
|
|
96201a623a | ||
|
|
9c94c2b11e | ||
|
|
e09b3f9fb5 | ||
|
|
a0099da2d1 | ||
|
|
25e0e81e12 | ||
|
|
87cff95af8 | ||
|
|
3ef69ef784 | ||
|
|
1a7786be14 | ||
|
|
3370ff8b1c | ||
|
|
455f57457d | ||
|
|
1d955ed990 | ||
|
|
d071237131 |
No files matched your search
@@ -39,11 +39,27 @@ jobs:
|
||||
APPLE_ID: ${{ vars.APPLE_ID }}
|
||||
MACOS_SIGNING_KEY: ${{ secrets.MACOS_SIGNING_KEY }}
|
||||
MACOS_SIGNING_KEY_PASSWORD: ${{ secrets.MACOS_SIGNING_KEY_PASSWORD }}
|
||||
DEVELOPER_DIR: /Applications/Xcode_26.4.1.app/Contents/Developer
|
||||
CGO_CFLAGS: '-mmacosx-version-min=14.0 -O3'
|
||||
CGO_CXXFLAGS: '-mmacosx-version-min=14.0 -O3'
|
||||
CGO_LDFLAGS: '-mmacosx-version-min=14.0 -O3'
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
- name: Select Xcode 26.4.1
|
||||
shell: bash
|
||||
run: |
|
||||
set -euo pipefail
|
||||
if [ ! -d "${DEVELOPER_DIR}" ]; then
|
||||
echo "Missing ${DEVELOPER_DIR}"
|
||||
ls -1 /Applications | grep '^Xcode' || true
|
||||
exit 1
|
||||
fi
|
||||
|
||||
sudo xcode-select -s "${DEVELOPER_DIR}"
|
||||
sw_vers
|
||||
xcodebuild -version
|
||||
xcrun --sdk macosx --show-sdk-version
|
||||
xcrun --find metal
|
||||
- run: |
|
||||
echo $MACOS_SIGNING_KEY | base64 --decode > certificate.p12
|
||||
security create-keychain -p password build.keychain
|
||||
|
||||
@@ -0,0 +1,21 @@
|
||||
# AGENTS.md
|
||||
|
||||
## Building
|
||||
|
||||
For a full build from the repository root:
|
||||
|
||||
```sh
|
||||
cmake -B build .
|
||||
cmake --build build --parallel 8
|
||||
./ollama serve
|
||||
```
|
||||
|
||||
For quick Go-only iteration against an existing native payload:
|
||||
|
||||
```sh
|
||||
go build .
|
||||
go run . serve
|
||||
```
|
||||
|
||||
See `docs/development.md` for prerequisites, platform notes, GPU backends, and
|
||||
the full development workflow.
|
||||
@@ -0,0 +1,3 @@
|
||||
# CLAUDE.md
|
||||
|
||||
See `AGENTS.md` for the shared agent instructions for this repository.
|
||||
+1
-1
@@ -1 +1 @@
|
||||
b9493
|
||||
b9888
|
||||
+1
-1
@@ -1 +1 @@
|
||||
2165dc08d7b33258260aa849d39f087d50e62962
|
||||
de7b4ed986b6d6f55b8ace5e73c24d1ca0bea89b
|
||||
@@ -77,10 +77,10 @@ ollama launch openclaw
|
||||
|
||||
### Chat with a model
|
||||
|
||||
Run and chat with [Gemma 3](https://ollama.com/library/gemma3):
|
||||
Run and chat with [Gemma 4](https://ollama.com/library/gemma4):
|
||||
|
||||
```
|
||||
ollama run gemma3
|
||||
ollama run gemma4
|
||||
```
|
||||
|
||||
See [ollama.com/library](https://ollama.com/library) for the full list.
|
||||
@@ -93,7 +93,7 @@ Ollama has a REST API for running and managing models.
|
||||
|
||||
```
|
||||
curl http://localhost:11434/api/chat -d '{
|
||||
"model": "gemma3",
|
||||
"model": "gemma4",
|
||||
"messages": [{
|
||||
"role": "user",
|
||||
"content": "Why is the sky blue?"
|
||||
@@ -113,7 +113,7 @@ pip install ollama
|
||||
```python
|
||||
from ollama import chat
|
||||
|
||||
response = chat(model='gemma3', messages=[
|
||||
response = chat(model='gemma4', messages=[
|
||||
{
|
||||
'role': 'user',
|
||||
'content': 'Why is the sky blue?',
|
||||
@@ -132,7 +132,7 @@ npm i ollama
|
||||
import ollama from "ollama";
|
||||
|
||||
const response = await ollama.chat({
|
||||
model: "gemma3",
|
||||
model: "gemma4",
|
||||
messages: [{ role: "user", content: "Why is the sky blue?" }],
|
||||
});
|
||||
console.log(response.message.content);
|
||||
|
||||
@@ -0,0 +1,198 @@
|
||||
package agent
|
||||
|
||||
import (
|
||||
"context"
|
||||
"strings"
|
||||
"sync"
|
||||
)
|
||||
|
||||
type ApprovalRequest struct {
|
||||
WorkingDir string
|
||||
Calls []ApprovalToolCall
|
||||
}
|
||||
|
||||
func (r *ApprovalRequest) AddToolCall(id, name, scope string, args map[string]any) {
|
||||
r.Calls = append(r.Calls, ApprovalToolCall{
|
||||
ToolCallID: id,
|
||||
ToolName: name,
|
||||
Args: args,
|
||||
ApprovalScope: scope,
|
||||
})
|
||||
}
|
||||
|
||||
type ApprovalToolCall struct {
|
||||
ToolCallID string
|
||||
ToolName string
|
||||
Args map[string]any
|
||||
ApprovalScope string
|
||||
}
|
||||
|
||||
type Approval struct {
|
||||
Allow bool
|
||||
AllowAll bool
|
||||
AllowScopes []string
|
||||
Reason string
|
||||
}
|
||||
|
||||
type ApprovalPrompter interface {
|
||||
PromptApproval(context.Context, ApprovalRequest) (Approval, error)
|
||||
}
|
||||
|
||||
type ApprovalState struct {
|
||||
mu sync.RWMutex
|
||||
allowAll bool
|
||||
scopes map[string]bool
|
||||
}
|
||||
|
||||
func (s *ApprovalState) Set(allowAll bool, scopes map[string]bool) {
|
||||
if s == nil {
|
||||
return
|
||||
}
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
s.allowAll = allowAll
|
||||
s.scopes = cloneApprovalScopes(scopes)
|
||||
}
|
||||
|
||||
// GrantAll grants blanket approval for all future tool calls.
|
||||
func (s *ApprovalState) GrantAll() {
|
||||
if s == nil {
|
||||
return
|
||||
}
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
s.allowAll = true
|
||||
}
|
||||
|
||||
// AllGranted reports whether blanket approval has been granted.
|
||||
func (s *ApprovalState) AllGranted() bool {
|
||||
if s == nil {
|
||||
return false
|
||||
}
|
||||
s.mu.RLock()
|
||||
defer s.mu.RUnlock()
|
||||
return s.allowAll
|
||||
}
|
||||
|
||||
func (s *ApprovalState) Allows(scope string) bool {
|
||||
if s == nil {
|
||||
return false
|
||||
}
|
||||
s.mu.RLock()
|
||||
defer s.mu.RUnlock()
|
||||
return s.allowAll || s.scopes[scope]
|
||||
}
|
||||
|
||||
// Apply merges an approval's scopes and allow-all flag into the state. It
|
||||
// returns true if the approval grants permission (allow-all or at least one
|
||||
// scope). It does not mutate the approval; the caller sets Allow based on the
|
||||
// returned value.
|
||||
func (s *ApprovalState) Apply(result *Approval) bool {
|
||||
if s == nil || result == nil {
|
||||
return false
|
||||
}
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
granted := false
|
||||
if result.AllowAll {
|
||||
s.allowAll = true
|
||||
granted = true
|
||||
}
|
||||
if len(result.AllowScopes) > 0 {
|
||||
granted = true
|
||||
s.grantScopesLocked(result.AllowScopes)
|
||||
}
|
||||
return granted
|
||||
}
|
||||
|
||||
// GrantScopes merges the given scopes into the state.
|
||||
func (s *ApprovalState) GrantScopes(scopes []string) {
|
||||
if s == nil {
|
||||
return
|
||||
}
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
s.grantScopesLocked(scopes)
|
||||
}
|
||||
|
||||
// grantScopesLocked adds trimmed, non-empty scopes to the state. Caller must
|
||||
// hold s.mu.
|
||||
func (s *ApprovalState) grantScopesLocked(scopes []string) {
|
||||
if s.scopes == nil {
|
||||
s.scopes = make(map[string]bool, len(scopes))
|
||||
}
|
||||
for _, scope := range scopes {
|
||||
scope = strings.TrimSpace(scope)
|
||||
if scope != "" {
|
||||
s.scopes[scope] = true
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func cloneApprovalScopes(src map[string]bool) map[string]bool {
|
||||
if len(src) == 0 {
|
||||
return nil
|
||||
}
|
||||
dst := make(map[string]bool, len(src))
|
||||
for scope, allowed := range src {
|
||||
if allowed {
|
||||
dst[scope] = true
|
||||
}
|
||||
}
|
||||
return dst
|
||||
}
|
||||
|
||||
func (s *Session) needsApproval(tool Tool, name string, args map[string]any) bool {
|
||||
return ToolRequiresApproval(tool, args) && !s.allows(toolApprovalScope(tool, name, args))
|
||||
}
|
||||
|
||||
// allows reports whether scope is permitted by the session's accumulated approval state.
|
||||
func (s *Session) allows(scope string) bool {
|
||||
if s == nil || s.ApprovalState == nil {
|
||||
return false
|
||||
}
|
||||
return s.ApprovalState.Allows(scope)
|
||||
}
|
||||
|
||||
// applyApproval merges an approval result into the session's state and marks
|
||||
// the result as allowed when scopes or allow-all were granted.
|
||||
func (s *Session) applyApproval(result *Approval) {
|
||||
if s == nil || result == nil {
|
||||
return
|
||||
}
|
||||
if s.ApprovalState == nil {
|
||||
s.ApprovalState = &ApprovalState{}
|
||||
}
|
||||
if s.ApprovalState.Apply(result) {
|
||||
result.Allow = true
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Session) authorizeToolCalls(ctx context.Context, req ApprovalRequest) (Approval, error) {
|
||||
if s == nil || len(req.Calls) == 0 || (s.ApprovalState != nil && s.ApprovalState.AllGranted()) {
|
||||
return Approval{Allow: true}, nil
|
||||
}
|
||||
if s.ApprovalPrompter == nil {
|
||||
return Approval{
|
||||
Reason: "Tool execution requires approval, but no approval prompter is available.",
|
||||
}, nil
|
||||
}
|
||||
|
||||
result, err := s.ApprovalPrompter.PromptApproval(ctx, req)
|
||||
if err != nil {
|
||||
return Approval{}, err
|
||||
}
|
||||
s.applyApproval(&result)
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// toolApprovalScope returns the approval scope key for a tool invocation.
|
||||
// If the tool implements ScopedTool, its ApprovalScope method determines the
|
||||
// scope (e.g. shell tools scope to "<tool>\x00<command>"). Otherwise the scope
|
||||
// is the trimmed tool name.
|
||||
func toolApprovalScope(tool Tool, toolName string, args map[string]any) string {
|
||||
if scoped, ok := tool.(ScopedTool); ok {
|
||||
return scoped.ApprovalScope(args)
|
||||
}
|
||||
return strings.TrimSpace(toolName)
|
||||
}
|
||||
@@ -0,0 +1,95 @@
|
||||
package agent
|
||||
|
||||
import (
|
||||
"context"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/ollama/ollama/api"
|
||||
)
|
||||
|
||||
type mockTool struct {
|
||||
name string
|
||||
}
|
||||
|
||||
func (m mockTool) Name() string { return m.name }
|
||||
func (m mockTool) Description() string { return "" }
|
||||
func (m mockTool) Schema() api.ToolFunction {
|
||||
return api.ToolFunction{Name: m.name}
|
||||
}
|
||||
|
||||
func (m mockTool) Execute(context.Context, ToolContext, map[string]any) (ToolResult, error) {
|
||||
return ToolResult{}, nil
|
||||
}
|
||||
|
||||
func TestToolApprovalScopeUsesScopedTool(t *testing.T) {
|
||||
shellTool := mockScopedTool{
|
||||
mockTool: mockTool{name: "bash"},
|
||||
scope: func(args map[string]any) string {
|
||||
if cmd, ok := args["command"].(string); ok {
|
||||
cmd = strings.TrimSpace(cmd)
|
||||
if cmd != "" {
|
||||
return "bash\x00" + cmd
|
||||
}
|
||||
}
|
||||
return "bash"
|
||||
},
|
||||
}
|
||||
plainTool := mockTool{name: "edit"}
|
||||
|
||||
tests := []struct {
|
||||
tool Tool
|
||||
name string
|
||||
args map[string]any
|
||||
want string
|
||||
}{
|
||||
{shellTool, "bash", map[string]any{"command": " pwd "}, "bash\x00pwd"},
|
||||
{shellTool, "bash", map[string]any{"command": "Get-ChildItem"}, "bash\x00Get-ChildItem"},
|
||||
{plainTool, "edit", map[string]any{"path": "README.md"}, "edit"},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
if got := toolApprovalScope(tt.tool, tt.name, tt.args); got != tt.want {
|
||||
t.Fatalf("toolApprovalScope(%q) = %q, want %q", tt.name, got, tt.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
type mockScopedTool struct {
|
||||
mockTool
|
||||
scope func(args map[string]any) string
|
||||
}
|
||||
|
||||
func (m mockScopedTool) ApprovalScope(args map[string]any) string {
|
||||
return m.scope(args)
|
||||
}
|
||||
|
||||
func TestSessionApplyApprovalScopes(t *testing.T) {
|
||||
session := &Session{}
|
||||
result := Approval{AllowScopes: []string{"edit", "bash\x00pwd", " "}}
|
||||
|
||||
session.applyApproval(&result)
|
||||
|
||||
if !result.Allow {
|
||||
t.Fatal("scoped approval should allow the current request")
|
||||
}
|
||||
if !session.allows("edit") || !session.allows("bash\x00pwd") {
|
||||
t.Fatal("scoped approval was not saved")
|
||||
}
|
||||
if session.allows("bash") || session.allows("bash\x00ls") {
|
||||
t.Fatal("shell approval was too broad")
|
||||
}
|
||||
if session.ApprovalState.AllGranted() {
|
||||
t.Fatal("allow all = true, want false for scoped approval")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSessionApplyApprovalAllowAll(t *testing.T) {
|
||||
session := &Session{}
|
||||
result := Approval{AllowAll: true}
|
||||
|
||||
session.applyApproval(&result)
|
||||
|
||||
if !result.Allow || !session.allows("anything") {
|
||||
t.Fatalf("allow all = %v result = %#v, want allow all", session.ApprovalState.AllGranted(), result)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,667 @@
|
||||
package agent
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"strings"
|
||||
|
||||
"github.com/ollama/ollama/api"
|
||||
)
|
||||
|
||||
// Compaction wire-format. These constants and helpers are the single canonical
|
||||
// definition of how a compacted turn is represented in message history.
|
||||
const (
|
||||
CompactionSummaryMessagePrefix = "Conversation summary:\n"
|
||||
CompactionToolName = "summary"
|
||||
CompactionToolCallID = "ollama_compaction"
|
||||
CompactionContinueInstruction = "continue the task in progress. the history has been compacted, do not mention compaction to the user"
|
||||
)
|
||||
|
||||
const (
|
||||
defaultCompactionContextWindowTokens = 32768
|
||||
defaultCompactionKeepUserTurns = 3
|
||||
defaultCompactionThreshold = 0.8
|
||||
compactOnlySummaryContextTokens = 16000
|
||||
|
||||
maxCompactionSummaryRunes = 16 * 1024
|
||||
|
||||
compactionSystemPrompt = "Summarize the archived part of an Ollama agent conversation. Preserve user goals, decisions, files, commands, tool results, and unresolved tasks needed to continue. Omit private reasoning and return only the summary."
|
||||
)
|
||||
|
||||
type Compactor interface {
|
||||
MaybeCompact(context.Context, CompactionRequest) (CompactionResult, error)
|
||||
|
||||
// ContextWindowTokens returns the effective context window size in
|
||||
// tokens, resolving runtime options against configured defaults.
|
||||
ContextWindowTokens(options map[string]any) int
|
||||
|
||||
// Threshold returns the compaction threshold as a fraction of the
|
||||
// context window (e.g. 0.8 means compact at 80% capacity).
|
||||
Threshold() float64
|
||||
|
||||
// ShouldCompact reports whether a compaction should run and returns the
|
||||
// trigger reason. An empty trigger means compaction is not needed.
|
||||
ShouldCompact(req CompactionRequest) (trigger string, should bool)
|
||||
}
|
||||
|
||||
type CompactionOptions struct {
|
||||
ContextWindowTokens int
|
||||
KeepUserTurns int
|
||||
Threshold float64
|
||||
}
|
||||
|
||||
type CompactionRequest struct {
|
||||
ChatID string
|
||||
Model string
|
||||
SystemPrompt string
|
||||
Messages []api.Message
|
||||
Tools api.Tools
|
||||
Format string
|
||||
Latest api.ChatResponse
|
||||
Options map[string]any
|
||||
KeepAlive *api.Duration
|
||||
Think *api.ThinkValue
|
||||
Force bool
|
||||
ContinueTask bool
|
||||
KeepUserTurns *int
|
||||
Progress func(CompactionProgress)
|
||||
}
|
||||
|
||||
type CompactionProgress struct {
|
||||
Tokens int
|
||||
}
|
||||
|
||||
type CompactionResult struct {
|
||||
Messages []api.Message
|
||||
Compacted bool
|
||||
Due bool
|
||||
Summary string
|
||||
Reason string
|
||||
}
|
||||
|
||||
type SimpleCompactor struct {
|
||||
Client ChatClient
|
||||
Options CompactionOptions
|
||||
}
|
||||
|
||||
func (c *SimpleCompactor) MaybeCompact(ctx context.Context, req CompactionRequest) (CompactionResult, error) {
|
||||
result := CompactionResult{Messages: req.Messages}
|
||||
if c == nil {
|
||||
return result, nil
|
||||
}
|
||||
|
||||
result.Due = req.Force || c.shouldCompact(req)
|
||||
if !result.Due {
|
||||
return result, nil
|
||||
}
|
||||
if c.Client == nil {
|
||||
result.Reason = "compaction is unavailable"
|
||||
return result, nil
|
||||
}
|
||||
|
||||
keepUserTurns := c.keepUserTurns(req.Options)
|
||||
if req.KeepUserTurns != nil {
|
||||
keepUserTurns = *req.KeepUserTurns
|
||||
}
|
||||
prefix, previousSummary, archive, suffix, _, ok := splitCompactionMessages(req.Messages, keepUserTurns)
|
||||
if !ok || len(archive) == 0 {
|
||||
result.Reason = "nothing to compact"
|
||||
return result, nil
|
||||
}
|
||||
|
||||
summary, err := c.summarize(ctx, req, previousSummary, archive)
|
||||
if err != nil {
|
||||
result.Reason = err.Error()
|
||||
return result, err
|
||||
}
|
||||
summary = truncateCompactionSummary(strings.TrimSpace(summary))
|
||||
if summary == "" {
|
||||
summary, err = c.summarizeEmptyFallback(ctx, req, previousSummary, archive)
|
||||
if err != nil {
|
||||
result.Reason = err.Error()
|
||||
return result, err
|
||||
}
|
||||
summary = truncateCompactionSummary(strings.TrimSpace(summary))
|
||||
}
|
||||
if summary == "" {
|
||||
result.Reason = "summary was empty"
|
||||
return result, nil
|
||||
}
|
||||
|
||||
compacted := make([]api.Message, 0, len(prefix)+len(suffix)+2)
|
||||
compacted = append(compacted, prefix...)
|
||||
compacted = append(compacted, CompactionSummaryMessages(summary, req.ContinueTask)...)
|
||||
compacted = append(compacted, suffix...)
|
||||
result.Messages = compacted
|
||||
result.Compacted = true
|
||||
result.Summary = summary
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func (c *SimpleCompactor) shouldCompact(req CompactionRequest) bool {
|
||||
contextWindow := c.contextWindowTokens(req.Options)
|
||||
threshold := int(float64(contextWindow) * c.threshold())
|
||||
if threshold <= 0 {
|
||||
return false
|
||||
}
|
||||
if req.Latest.PromptEvalCount > 0 && req.Latest.PromptEvalCount >= threshold {
|
||||
return true
|
||||
}
|
||||
return estimateCompactionRequestTokens(req) >= threshold
|
||||
}
|
||||
|
||||
func (c *SimpleCompactor) contextWindowTokens(options map[string]any) int {
|
||||
return ResolveContextWindowTokens(options, c.Options.ContextWindowTokens)
|
||||
}
|
||||
|
||||
// ContextWindowTokens resolves the effective context window from runtime
|
||||
// options or configured defaults. Satisfies the Compactor interface.
|
||||
func (c *SimpleCompactor) ContextWindowTokens(options map[string]any) int {
|
||||
if c == nil {
|
||||
return 0
|
||||
}
|
||||
return c.contextWindowTokens(options)
|
||||
}
|
||||
|
||||
func (c *SimpleCompactor) threshold() float64 {
|
||||
return ResolveCompactionThreshold(c.Options.Threshold)
|
||||
}
|
||||
|
||||
// Threshold returns the configured compaction threshold fraction. Satisfies
|
||||
// the Compactor interface.
|
||||
func (c *SimpleCompactor) Threshold() float64 {
|
||||
if c == nil {
|
||||
return 0
|
||||
}
|
||||
return c.threshold()
|
||||
}
|
||||
|
||||
// ShouldCompact reports whether compaction is due and the trigger reason.
|
||||
// Satisfies the Compactor interface.
|
||||
func (c *SimpleCompactor) ShouldCompact(req CompactionRequest) (string, bool) {
|
||||
if c == nil {
|
||||
return "", false
|
||||
}
|
||||
if req.Force {
|
||||
return "force", true
|
||||
}
|
||||
if c.shouldCompact(req) {
|
||||
contextWindow := c.contextWindowTokens(req.Options)
|
||||
threshold := int(float64(contextWindow) * c.threshold())
|
||||
if req.Latest.PromptEvalCount > 0 && req.Latest.PromptEvalCount >= threshold {
|
||||
return "prompt_eval", true
|
||||
}
|
||||
return "estimate", true
|
||||
}
|
||||
return "", false
|
||||
}
|
||||
|
||||
func (c *SimpleCompactor) keepUserTurns(options map[string]any) int {
|
||||
contextWindow := c.contextWindowTokens(options)
|
||||
if contextWindow > 0 && contextWindow < compactOnlySummaryContextTokens {
|
||||
return 0
|
||||
}
|
||||
if c.Options.KeepUserTurns > 0 {
|
||||
return c.Options.KeepUserTurns
|
||||
}
|
||||
return defaultCompactionKeepUserTurns
|
||||
}
|
||||
|
||||
func ResolveContextWindowTokens(options map[string]any, configured int) int {
|
||||
if n := intOption(options, "num_ctx"); n > 0 {
|
||||
return n
|
||||
}
|
||||
if configured > 0 {
|
||||
return configured
|
||||
}
|
||||
return defaultCompactionContextWindowTokens
|
||||
}
|
||||
|
||||
func ResolveCompactionThreshold(configured float64) float64 {
|
||||
if configured > 0 {
|
||||
return configured
|
||||
}
|
||||
return defaultCompactionThreshold
|
||||
}
|
||||
|
||||
func (c *SimpleCompactor) summarize(ctx context.Context, req CompactionRequest, previousSummary string, archive []api.Message) (string, error) {
|
||||
body, err := compactionPrompt(previousSummary, archive, c.compactionPromptBodyBudgetTokens(req.Options))
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
chatReq := &api.ChatRequest{
|
||||
Model: req.Model,
|
||||
Messages: []api.Message{
|
||||
{
|
||||
Role: "system",
|
||||
Content: compactionSystemPrompt,
|
||||
},
|
||||
{
|
||||
Role: "user",
|
||||
Content: body,
|
||||
},
|
||||
},
|
||||
Options: req.Options,
|
||||
Think: req.Think,
|
||||
}
|
||||
if req.KeepAlive != nil {
|
||||
chatReq.KeepAlive = req.KeepAlive
|
||||
}
|
||||
|
||||
var summary strings.Builder
|
||||
if err := c.Client.Chat(ctx, chatReq, func(response api.ChatResponse) error {
|
||||
summary.WriteString(response.Message.Content)
|
||||
if req.Progress != nil {
|
||||
tokens := response.EvalCount
|
||||
if tokens <= 0 {
|
||||
tokens = estimateCompactionTokens(summary.String())
|
||||
}
|
||||
req.Progress(CompactionProgress{Tokens: tokens})
|
||||
}
|
||||
return nil
|
||||
}); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return summary.String(), nil
|
||||
}
|
||||
|
||||
func (c *SimpleCompactor) summarizeEmptyFallback(ctx context.Context, req CompactionRequest, previousSummary string, archive []api.Message) (string, error) {
|
||||
retry := req
|
||||
retry.Think = &api.ThinkValue{Value: false}
|
||||
summary, err := c.summarize(ctx, retry, previousSummary, archive)
|
||||
if err == nil {
|
||||
return summary, nil
|
||||
}
|
||||
if !isUnsupportedCompactionThinkError(err) {
|
||||
return "", err
|
||||
}
|
||||
if req.Think == nil {
|
||||
return "", nil
|
||||
}
|
||||
retry.Think = nil
|
||||
return c.summarize(ctx, retry, previousSummary, archive)
|
||||
}
|
||||
|
||||
func isUnsupportedCompactionThinkError(err error) bool {
|
||||
if err == nil {
|
||||
return false
|
||||
}
|
||||
text := strings.ToLower(err.Error())
|
||||
if !strings.Contains(text, "think") {
|
||||
return false
|
||||
}
|
||||
var statusErr api.StatusError
|
||||
if errors.As(err, &statusErr) && statusErr.StatusCode != 0 {
|
||||
return statusErr.StatusCode == http.StatusBadRequest
|
||||
}
|
||||
return strings.Contains(text, "does not support") || strings.Contains(text, "not supported") || strings.Contains(text, "unsupported")
|
||||
}
|
||||
|
||||
// compactionSummaryMessageForTask renders a compaction summary as the content
|
||||
// string stored on the synthetic tool-result message.
|
||||
func compactionSummaryMessageForTask(summary string, continueTask bool) string {
|
||||
content := CompactionSummaryMessagePrefix + strings.TrimSpace(summary)
|
||||
if continueTask {
|
||||
content = strings.TrimSpace(content) + "\n\n" + CompactionContinueInstruction
|
||||
}
|
||||
return content
|
||||
}
|
||||
|
||||
// CompactionSummaryMessages renders a compaction summary as the assistant
|
||||
// tool-call plus tool-result pair that represents a compacted turn in the
|
||||
// message history.
|
||||
func CompactionSummaryMessages(summary string, continueTask bool) []api.Message {
|
||||
return []api.Message{
|
||||
{
|
||||
Role: "assistant",
|
||||
ToolCalls: []api.ToolCall{{
|
||||
ID: CompactionToolCallID,
|
||||
Function: api.ToolCallFunction{
|
||||
Name: CompactionToolName,
|
||||
},
|
||||
}},
|
||||
},
|
||||
{
|
||||
Role: "tool",
|
||||
ToolName: CompactionToolName,
|
||||
ToolCallID: CompactionToolCallID,
|
||||
Content: compactionSummaryMessageForTask(summary, continueTask),
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func (c *SimpleCompactor) compactionPromptBodyBudgetTokens(options map[string]any) int {
|
||||
contextWindow := c.contextWindowTokens(options)
|
||||
threshold := int(float64(contextWindow) * c.threshold())
|
||||
if threshold <= 0 {
|
||||
return 0
|
||||
}
|
||||
systemTokens := estimateCompactionTokens("system") + estimateCompactionTokens(compactionSystemPrompt)
|
||||
userRoleTokens := estimateCompactionTokens("user")
|
||||
budget := threshold - systemTokens - userRoleTokens
|
||||
if budget <= 0 {
|
||||
return 0
|
||||
}
|
||||
return budget
|
||||
}
|
||||
|
||||
func truncateCompactionSummary(summary string) string {
|
||||
return Truncate(summary, TruncateConfig{
|
||||
MaxRunes: maxCompactionSummaryRunes,
|
||||
Label: "summary",
|
||||
})
|
||||
}
|
||||
|
||||
func estimateCompactionTokens(text string) int {
|
||||
text = strings.TrimSpace(text)
|
||||
if text == "" {
|
||||
return 0
|
||||
}
|
||||
return ApproximateTokens(len([]rune(text)))
|
||||
}
|
||||
|
||||
func estimateMessagesTokens(messages []api.Message) int {
|
||||
var total int
|
||||
for _, msg := range messages {
|
||||
total += estimateCompactionTokens(msg.Role)
|
||||
total += estimateCompactionTokens(msg.Content)
|
||||
total += estimateCompactionTokens(msg.Thinking)
|
||||
total += estimateCompactionTokens(msg.ToolName)
|
||||
total += estimateCompactionTokens(msg.ToolCallID)
|
||||
for _, call := range msg.ToolCalls {
|
||||
total += estimateCompactionTokens(call.Function.Name)
|
||||
total += estimateCompactionTokens(call.Function.Arguments.String())
|
||||
}
|
||||
}
|
||||
return total
|
||||
}
|
||||
|
||||
func estimateCompactionRequestTokens(req CompactionRequest) int {
|
||||
requestMessages := sanitizeMessagesForEstimate(req.Messages)
|
||||
if strings.TrimSpace(req.SystemPrompt) != "" {
|
||||
requestMessages = make([]api.Message, 0, len(req.Messages)+1)
|
||||
requestMessages = append(requestMessages, api.Message{Role: "system", Content: strings.TrimSpace(req.SystemPrompt)})
|
||||
requestMessages = append(requestMessages, sanitizeMessagesForEstimate(req.Messages)...)
|
||||
}
|
||||
|
||||
payload := struct {
|
||||
Messages []api.Message `json:"messages,omitempty"`
|
||||
Tools api.Tools `json:"tools,omitempty"`
|
||||
Format json.RawMessage `json:"format,omitempty"`
|
||||
}{
|
||||
Messages: requestMessages,
|
||||
Tools: req.Tools,
|
||||
}
|
||||
if rawFormat, ok := compactionFormatForEstimate(req.Format); ok {
|
||||
payload.Format = rawFormat
|
||||
}
|
||||
if data, err := json.Marshal(payload); err == nil {
|
||||
return estimateCompactionTokens(string(data))
|
||||
}
|
||||
|
||||
total := estimateMessagesTokens(requestMessages)
|
||||
total += estimateCompactionTokens(req.Tools.String())
|
||||
total += estimateCompactionTokens(req.Format)
|
||||
return total
|
||||
}
|
||||
|
||||
func (s *Session) estimateRunPromptTokens(opts RunOptions, messages []api.Message) int {
|
||||
return estimateCompactionRequestTokens(CompactionRequest{
|
||||
SystemPrompt: opts.SystemPrompt,
|
||||
Messages: messages,
|
||||
Tools: s.availableTools(),
|
||||
Format: opts.Format,
|
||||
Options: opts.Options,
|
||||
})
|
||||
}
|
||||
|
||||
func (s *Session) checkPreflightPromptBudget(opts RunOptions, messages []api.Message) error {
|
||||
contextWindow := s.contextWindowTokens(opts)
|
||||
if contextWindow <= 0 {
|
||||
return nil
|
||||
}
|
||||
estimated := s.estimateRunPromptTokens(opts, messages)
|
||||
if estimated < contextWindow {
|
||||
return nil
|
||||
}
|
||||
return fmt.Errorf("prompt is too large for the current context (~%d/%d tokens). Reduce the system prompt or message history, compact the conversation, or use a model with a larger context", estimated, contextWindow)
|
||||
}
|
||||
|
||||
func (s *Session) checkPostCompactionPromptBudget(opts RunOptions, messages []api.Message) error {
|
||||
contextWindow := s.contextWindowTokens(opts)
|
||||
if contextWindow <= 0 {
|
||||
return nil
|
||||
}
|
||||
estimated := s.estimateRunPromptTokens(opts, messages)
|
||||
if estimated < contextWindow {
|
||||
return nil
|
||||
}
|
||||
return fmt.Errorf("history is still too large after compaction (~%d/%d tokens). Start a fresh request, reduce the system prompt or history, or use a model with a larger context", estimated, contextWindow)
|
||||
}
|
||||
|
||||
func sanitizeMessagesForEstimate(messages []api.Message) []api.Message {
|
||||
requestMessages := sanitizeMessagesForRequest(messages)
|
||||
for i := range requestMessages {
|
||||
// Image token accounting is model-specific. Without the active model's
|
||||
// tokenizer and vision accounting, raw image bytes/base64 make the
|
||||
// estimate look much larger than the prompt the model actually sees.
|
||||
requestMessages[i].Images = nil
|
||||
}
|
||||
return requestMessages
|
||||
}
|
||||
|
||||
func compactionFormatForEstimate(format string) (json.RawMessage, bool) {
|
||||
format = strings.TrimSpace(format)
|
||||
if format == "" {
|
||||
return nil, false
|
||||
}
|
||||
if format == "json" {
|
||||
return json.RawMessage(`"json"`), true
|
||||
}
|
||||
if !json.Valid([]byte(format)) {
|
||||
return nil, false
|
||||
}
|
||||
return json.RawMessage(format), true
|
||||
}
|
||||
|
||||
func compactionPrompt(previousSummary string, archive []api.Message, maxTokens int) (string, error) {
|
||||
messages := make([]api.Message, 0, len(archive))
|
||||
for _, msg := range archive {
|
||||
msg.Thinking = ""
|
||||
msg.Images = nil
|
||||
messages = append(messages, msg)
|
||||
}
|
||||
return renderCompactionPrompt(previousSummary, fitCompactionMessagesToBudget(previousSummary, messages, maxTokens))
|
||||
}
|
||||
|
||||
func renderCompactionPrompt(previousSummary string, messages []api.Message) (string, error) {
|
||||
payload, err := json.MarshalIndent(messages, "", " ")
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("marshal compaction messages: %w", err)
|
||||
}
|
||||
|
||||
var b strings.Builder
|
||||
if strings.TrimSpace(previousSummary) != "" {
|
||||
b.WriteString("Previous summary:\n")
|
||||
b.WriteString(strings.TrimSpace(previousSummary))
|
||||
b.WriteString("\n\n")
|
||||
}
|
||||
b.WriteString("Messages to archive as JSON:\n")
|
||||
b.Write(payload)
|
||||
return b.String(), nil
|
||||
}
|
||||
|
||||
func fitCompactionMessagesToBudget(previousSummary string, messages []api.Message, maxTokens int) []api.Message {
|
||||
if maxTokens <= 0 {
|
||||
return messages
|
||||
}
|
||||
fitted := append([]api.Message(nil), messages...)
|
||||
for range 16 {
|
||||
body, err := renderCompactionPrompt(previousSummary, fitted)
|
||||
if err != nil || estimateCompactionTokens(body) <= maxTokens {
|
||||
return fitted
|
||||
}
|
||||
|
||||
idx := largestCompactionContentMessage(fitted)
|
||||
if idx < 0 {
|
||||
return fitted
|
||||
}
|
||||
overageTokens := estimateCompactionTokens(body) - maxTokens
|
||||
currentRunes := len([]rune(fitted[idx].Content))
|
||||
nextRunes := currentRunes - overageTokens*4 - 256
|
||||
if nextRunes >= currentRunes {
|
||||
nextRunes = currentRunes / 2
|
||||
}
|
||||
fitted[idx].Content = truncateToolResultContentTo(fitted[idx].Content, nextRunes)
|
||||
}
|
||||
return fitted
|
||||
}
|
||||
|
||||
func largestCompactionContentMessage(messages []api.Message) int {
|
||||
idx := -1
|
||||
size := 0
|
||||
for i, msg := range messages {
|
||||
n := len([]rune(msg.Content))
|
||||
if n > size {
|
||||
idx = i
|
||||
size = n
|
||||
}
|
||||
}
|
||||
return idx
|
||||
}
|
||||
|
||||
func splitCompactionMessages(messages []api.Message, keepUserTurns int) (prefix []api.Message, previousSummary string, archive []api.Message, suffix []api.Message, keptUserTurns int, ok bool) {
|
||||
if keepUserTurns < 0 {
|
||||
keepUserTurns = defaultCompactionKeepUserTurns
|
||||
}
|
||||
|
||||
start := 0
|
||||
for start < len(messages) && messages[start].Role == "system" && !isCompactionSummary(messages[start]) {
|
||||
prefix = append(prefix, messages[start])
|
||||
start++
|
||||
}
|
||||
|
||||
candidates := make([]api.Message, 0, len(messages)-start)
|
||||
for i := start; i < len(messages); i++ {
|
||||
msg := messages[i]
|
||||
if isCompactionSummary(msg) {
|
||||
previousSummary = CompactionSummaryText(msg.Content)
|
||||
continue
|
||||
}
|
||||
if isCompactionToolCall(msg) {
|
||||
if i+1 < len(messages) && isCompactionSummary(messages[i+1]) {
|
||||
previousSummary = CompactionSummaryText(messages[i+1].Content)
|
||||
i++
|
||||
}
|
||||
continue
|
||||
}
|
||||
candidates = append(candidates, msg)
|
||||
}
|
||||
|
||||
userTurnIndexes := make([]int, 0, keepUserTurns)
|
||||
for i := len(candidates) - 1; i >= 0; i-- {
|
||||
if candidates[i].Role == "user" {
|
||||
userTurnIndexes = append(userTurnIndexes, i)
|
||||
}
|
||||
}
|
||||
keptUserTurns = keepUserTurns
|
||||
if len(userTurnIndexes) <= keptUserTurns {
|
||||
keptUserTurns = len(userTurnIndexes) - 1
|
||||
}
|
||||
if keptUserTurns < 0 {
|
||||
keptUserTurns = 0
|
||||
}
|
||||
|
||||
suffixStart := len(candidates)
|
||||
if keptUserTurns > 0 {
|
||||
suffixStart = userTurnIndexes[keptUserTurns-1]
|
||||
}
|
||||
if suffixStart <= 0 || len(candidates[:suffixStart]) == 0 {
|
||||
return prefix, previousSummary, nil, nil, keptUserTurns, false
|
||||
}
|
||||
|
||||
return prefix, previousSummary, candidates[:suffixStart], candidates[suffixStart:], keptUserTurns, true
|
||||
}
|
||||
|
||||
func isCompactionToolName(name string) bool {
|
||||
return name == CompactionToolName
|
||||
}
|
||||
|
||||
func isCompactionSummary(msg api.Message) bool {
|
||||
return (msg.Role == "user" || msg.Role == "system" || (msg.Role == "tool" && isCompactionToolName(msg.ToolName))) &&
|
||||
strings.HasPrefix(msg.Content, CompactionSummaryMessagePrefix)
|
||||
}
|
||||
|
||||
// IsCompactionSummary reports whether msg uses the canonical compaction
|
||||
// summary message representation.
|
||||
func IsCompactionSummary(msg api.Message) bool {
|
||||
return isCompactionSummary(msg)
|
||||
}
|
||||
|
||||
// CompactionSummaryContent returns the user-visible summary from msg when it
|
||||
// is a canonical compaction summary.
|
||||
func CompactionSummaryContent(msg api.Message) (string, bool) {
|
||||
if !isCompactionSummary(msg) {
|
||||
return "", false
|
||||
}
|
||||
return CompactionSummaryText(msg.Content), true
|
||||
}
|
||||
|
||||
// IsCompactionToolResult reports whether msg is the synthetic tool result used
|
||||
// to represent compaction in message history.
|
||||
func IsCompactionToolResult(msg api.Message) bool {
|
||||
return msg.Role == "tool" && (isCompactionToolName(msg.ToolName) || msg.ToolCallID == CompactionToolCallID)
|
||||
}
|
||||
|
||||
// IsCompactionToolCall reports whether msg is the synthetic assistant tool
|
||||
// call paired with a compaction summary result.
|
||||
func IsCompactionToolCall(msg api.Message) bool {
|
||||
return isCompactionToolCall(msg)
|
||||
}
|
||||
|
||||
func isCompactionToolCall(msg api.Message) bool {
|
||||
if msg.Role != "assistant" {
|
||||
return false
|
||||
}
|
||||
for _, call := range msg.ToolCalls {
|
||||
if isCompactionToolName(call.Function.Name) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// CompactionSummaryText reverses CompactionSummaryMessages, returning the
|
||||
// user-visible summary text with the prefix and any continuation instruction
|
||||
// removed.
|
||||
func CompactionSummaryText(content string) string {
|
||||
return strings.TrimSpace(strings.TrimSuffix(
|
||||
strings.TrimSpace(strings.TrimPrefix(content, CompactionSummaryMessagePrefix)),
|
||||
CompactionContinueInstruction,
|
||||
))
|
||||
}
|
||||
|
||||
func intOption(options map[string]any, key string) int {
|
||||
if options == nil {
|
||||
return 0
|
||||
}
|
||||
switch v := options[key].(type) {
|
||||
case int:
|
||||
return v
|
||||
case int64:
|
||||
return int(v)
|
||||
case float64:
|
||||
return int(v)
|
||||
case float32:
|
||||
return int(v)
|
||||
case json.Number:
|
||||
n, _ := v.Int64()
|
||||
return int(n)
|
||||
default:
|
||||
return 0
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,773 @@
|
||||
package agent
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/ollama/ollama/api"
|
||||
)
|
||||
|
||||
type scriptedCompactionClient struct {
|
||||
responses [][]api.ChatResponse
|
||||
errs []error
|
||||
requests []*api.ChatRequest
|
||||
}
|
||||
|
||||
func (c *scriptedCompactionClient) Chat(_ context.Context, req *api.ChatRequest, fn api.ChatResponseFunc) error {
|
||||
c.requests = append(c.requests, req)
|
||||
i := len(c.requests) - 1
|
||||
if i < len(c.responses) {
|
||||
for _, response := range c.responses[i] {
|
||||
if err := fn(response); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
}
|
||||
if i < len(c.errs) {
|
||||
return c.errs[i]
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func assertCompactionSummaryPair(t *testing.T, messages []api.Message) {
|
||||
t.Helper()
|
||||
if len(messages) != 2 {
|
||||
t.Fatalf("compaction summary pair len = %d, want 2: %#v", len(messages), messages)
|
||||
}
|
||||
if messages[0].Role != "assistant" || len(messages[0].ToolCalls) != 1 || messages[0].ToolCalls[0].Function.Name != CompactionToolName {
|
||||
t.Fatalf("compaction assistant message = %#v", messages[0])
|
||||
}
|
||||
if messages[0].ToolCalls[0].Function.Arguments.Len() != 0 {
|
||||
t.Fatalf("compaction summary tool call should not have arguments: %#v", messages[0].ToolCalls[0].Function.Arguments.ToMap())
|
||||
}
|
||||
if messages[1].Role != "tool" || messages[1].ToolName != CompactionToolName || messages[1].ToolCallID != messages[0].ToolCalls[0].ID {
|
||||
t.Fatalf("compaction tool result = %#v", messages[1])
|
||||
}
|
||||
if !strings.HasPrefix(messages[1].Content, CompactionSummaryMessagePrefix) {
|
||||
t.Fatalf("compaction tool result missing summary prefix: %#v", messages[1])
|
||||
}
|
||||
}
|
||||
|
||||
func TestSimpleCompactorSummarizesOldMessages(t *testing.T) {
|
||||
client := &fakeClient{
|
||||
responses: [][]api.ChatResponse{{
|
||||
{Message: api.Message{Role: "assistant", Content: "summary"}},
|
||||
}},
|
||||
}
|
||||
compactor := &SimpleCompactor{Client: client, Options: CompactionOptions{
|
||||
ContextWindowTokens: 16000,
|
||||
KeepUserTurns: 2,
|
||||
Threshold: 0.5,
|
||||
}}
|
||||
|
||||
messages := []api.Message{
|
||||
{Role: "system", Content: "stay pinned"},
|
||||
{Role: "user", Content: "old request"},
|
||||
{Role: "assistant", Content: "old answer", Thinking: "hidden"},
|
||||
{Role: "user", Content: "recent one"},
|
||||
{Role: "assistant", Content: "recent answer"},
|
||||
{Role: "user", Content: "recent two"},
|
||||
}
|
||||
|
||||
result, err := compactor.MaybeCompact(context.Background(), CompactionRequest{
|
||||
ChatID: "chat-1",
|
||||
Model: "model",
|
||||
Messages: messages,
|
||||
Latest: api.ChatResponse{Metrics: api.Metrics{PromptEvalCount: 12000}},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !result.Compacted {
|
||||
t.Fatal("expected compaction")
|
||||
}
|
||||
compacted := result.Messages
|
||||
if len(compacted) != 6 {
|
||||
t.Fatalf("compacted messages = %d, want 6", len(compacted))
|
||||
}
|
||||
if compacted[0].Content != "stay pinned" {
|
||||
t.Fatalf("first message = %#v", compacted[0])
|
||||
}
|
||||
if result.Summary != "summary" {
|
||||
t.Fatalf("result summary = %q", result.Summary)
|
||||
}
|
||||
assertCompactionSummaryPair(t, compacted[1:3])
|
||||
if compacted[3].Content != "recent one" || compacted[5].Content != "recent two" {
|
||||
t.Fatalf("recent turns were not kept: %#v", compacted)
|
||||
}
|
||||
if len(client.requests) != 1 {
|
||||
t.Fatalf("summary requests = %d, want 1", len(client.requests))
|
||||
}
|
||||
if strings.Contains(client.requests[0].Messages[1].Content, "hidden") {
|
||||
t.Fatal("compaction prompt should omit thinking")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSimpleCompactorKeepsOnlySummaryForSmallContext(t *testing.T) {
|
||||
client := &fakeClient{
|
||||
responses: [][]api.ChatResponse{{
|
||||
{Message: api.Message{Role: "assistant", Content: "small context summary"}},
|
||||
}},
|
||||
}
|
||||
compactor := &SimpleCompactor{Client: client, Options: CompactionOptions{
|
||||
ContextWindowTokens: compactOnlySummaryContextTokens - 1,
|
||||
KeepUserTurns: 3,
|
||||
Threshold: 0.5,
|
||||
}}
|
||||
|
||||
result, err := compactor.MaybeCompact(context.Background(), CompactionRequest{
|
||||
ChatID: "chat-1",
|
||||
Model: "model",
|
||||
ContinueTask: true,
|
||||
Messages: []api.Message{
|
||||
{Role: "system", Content: "pinned"},
|
||||
{Role: "user", Content: "old request"},
|
||||
{Role: "assistant", Content: "old answer"},
|
||||
{Role: "user", Content: "latest request"},
|
||||
},
|
||||
Latest: api.ChatResponse{Metrics: api.Metrics{PromptEvalCount: 12000}},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !result.Compacted {
|
||||
t.Fatal("expected compaction")
|
||||
}
|
||||
if len(result.Messages) != 3 {
|
||||
t.Fatalf("messages = %#v, want system plus compaction summary pair", result.Messages)
|
||||
}
|
||||
if result.Messages[0].Content != "pinned" {
|
||||
t.Fatalf("leading system message not kept: %#v", result.Messages)
|
||||
}
|
||||
assertCompactionSummaryPair(t, result.Messages[1:])
|
||||
if !strings.Contains(result.Messages[2].Content, CompactionContinueInstruction) {
|
||||
t.Fatalf("tool result missing continue instruction: %q", result.Messages[2].Content)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSimpleCompactorAddsContinueTaskInstructionOnlyToToolResult(t *testing.T) {
|
||||
client := &fakeClient{
|
||||
responses: [][]api.ChatResponse{{
|
||||
{Message: api.Message{Role: "assistant", Content: "summary"}},
|
||||
}},
|
||||
}
|
||||
compactor := &SimpleCompactor{Client: client, Options: CompactionOptions{
|
||||
ContextWindowTokens: 100,
|
||||
KeepUserTurns: 1,
|
||||
Threshold: 0.5,
|
||||
}}
|
||||
|
||||
result, err := compactor.MaybeCompact(context.Background(), CompactionRequest{
|
||||
ChatID: "chat-1",
|
||||
Model: "model",
|
||||
ContinueTask: true,
|
||||
Messages: []api.Message{
|
||||
{Role: "user", Content: "old request"},
|
||||
{Role: "assistant", Content: "old answer"},
|
||||
{Role: "user", Content: "recent request"},
|
||||
},
|
||||
Latest: api.ChatResponse{Metrics: api.Metrics{PromptEvalCount: 75}},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if result.Summary != "summary" {
|
||||
t.Fatalf("result summary = %q", result.Summary)
|
||||
}
|
||||
content := result.Messages[1].Content
|
||||
if !strings.Contains(content, CompactionContinueInstruction) {
|
||||
t.Fatalf("tool result missing continue instruction: %q", content)
|
||||
}
|
||||
if got := CompactionSummaryText(content); got != "summary" {
|
||||
t.Fatalf("visible summary text = %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSimpleCompactorTruncatesOversizedSummary(t *testing.T) {
|
||||
longSummary := strings.Repeat("x", maxCompactionSummaryRunes+1024)
|
||||
client := &fakeClient{
|
||||
responses: [][]api.ChatResponse{{
|
||||
{Message: api.Message{Role: "assistant", Content: longSummary}},
|
||||
}},
|
||||
}
|
||||
compactor := &SimpleCompactor{Client: client, Options: CompactionOptions{
|
||||
ContextWindowTokens: 100,
|
||||
KeepUserTurns: 1,
|
||||
Threshold: 0.5,
|
||||
}}
|
||||
|
||||
result, err := compactor.MaybeCompact(context.Background(), CompactionRequest{
|
||||
ChatID: "chat-1",
|
||||
Model: "model",
|
||||
Messages: []api.Message{
|
||||
{Role: "user", Content: "old one"},
|
||||
{Role: "assistant", Content: "old answer"},
|
||||
{Role: "user", Content: "recent one"},
|
||||
},
|
||||
Latest: api.ChatResponse{Metrics: api.Metrics{PromptEvalCount: 75}},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !result.Compacted {
|
||||
t.Fatal("expected compaction")
|
||||
}
|
||||
if runeCount := len([]rune(result.Summary)); runeCount > maxCompactionSummaryRunes+200 {
|
||||
t.Fatalf("summary runes = %d, want <= %d (plus marker)", runeCount, maxCompactionSummaryRunes)
|
||||
}
|
||||
if !strings.Contains(result.Summary, "[summary truncated:") {
|
||||
t.Fatalf("summary missing truncation marker: %q", result.Summary)
|
||||
}
|
||||
if !strings.Contains(result.Messages[1].Content, "[summary truncated:") {
|
||||
t.Fatalf("compacted message missing truncation marker: %#v", result.Messages)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSimpleCompactorRetriesEmptySummaryWithThinkFalse(t *testing.T) {
|
||||
client := &scriptedCompactionClient{
|
||||
responses: [][]api.ChatResponse{
|
||||
{{Message: api.Message{Role: "assistant", Thinking: "internal summary plan"}}},
|
||||
{{Message: api.Message{Role: "assistant", Content: "fallback summary"}}},
|
||||
},
|
||||
}
|
||||
compactor := &SimpleCompactor{Client: client, Options: CompactionOptions{
|
||||
ContextWindowTokens: 100,
|
||||
KeepUserTurns: 1,
|
||||
Threshold: 0.5,
|
||||
}}
|
||||
|
||||
result, err := compactor.MaybeCompact(context.Background(), CompactionRequest{
|
||||
Model: "model",
|
||||
Messages: []api.Message{
|
||||
{Role: "user", Content: "old request"},
|
||||
{Role: "assistant", Content: "old answer"},
|
||||
{Role: "user", Content: "recent request"},
|
||||
},
|
||||
Force: true,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !result.Compacted || result.Summary != "fallback summary" {
|
||||
t.Fatalf("compaction result = %#v", result)
|
||||
}
|
||||
if len(client.requests) != 2 {
|
||||
t.Fatalf("summary requests = %d, want 2", len(client.requests))
|
||||
}
|
||||
if client.requests[0].Think != nil {
|
||||
t.Fatalf("first summary request think = %#v, want nil", client.requests[0].Think)
|
||||
}
|
||||
if client.requests[1].Think == nil || client.requests[1].Think.Value != false {
|
||||
t.Fatalf("fallback summary request think = %#v, want false", client.requests[1].Think)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSimpleCompactorIgnoresUnsupportedThinkFalseFallback(t *testing.T) {
|
||||
client := &scriptedCompactionClient{
|
||||
responses: [][]api.ChatResponse{
|
||||
{{Message: api.Message{Role: "assistant", Thinking: "internal summary plan"}}},
|
||||
nil,
|
||||
},
|
||||
errs: []error{
|
||||
nil,
|
||||
api.StatusError{StatusCode: http.StatusBadRequest, ErrorMessage: "model does not support thinking"},
|
||||
},
|
||||
}
|
||||
compactor := &SimpleCompactor{Client: client, Options: CompactionOptions{
|
||||
ContextWindowTokens: 100,
|
||||
KeepUserTurns: 1,
|
||||
Threshold: 0.5,
|
||||
}}
|
||||
|
||||
result, err := compactor.MaybeCompact(context.Background(), CompactionRequest{
|
||||
Model: "model",
|
||||
Messages: []api.Message{
|
||||
{Role: "user", Content: "old request"},
|
||||
{Role: "assistant", Content: "old answer"},
|
||||
{Role: "user", Content: "recent request"},
|
||||
},
|
||||
Force: true,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if result.Compacted || result.Reason != "summary was empty" {
|
||||
t.Fatalf("compaction result = %#v", result)
|
||||
}
|
||||
if len(client.requests) != 2 {
|
||||
t.Fatalf("summary requests = %d, want 2", len(client.requests))
|
||||
}
|
||||
if client.requests[1].Think == nil || client.requests[1].Think.Value != false {
|
||||
t.Fatalf("fallback summary request think = %#v, want false", client.requests[1].Think)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSimpleCompactorFallsBackToUnsetThinkWhenThinkFalseUnsupported(t *testing.T) {
|
||||
client := &scriptedCompactionClient{
|
||||
responses: [][]api.ChatResponse{
|
||||
{{Message: api.Message{Role: "assistant", Thinking: "internal summary plan"}}},
|
||||
nil,
|
||||
{{Message: api.Message{Role: "assistant", Content: "unset think summary"}}},
|
||||
},
|
||||
errs: []error{
|
||||
nil,
|
||||
api.StatusError{StatusCode: http.StatusBadRequest, ErrorMessage: "think level is not supported"},
|
||||
nil,
|
||||
},
|
||||
}
|
||||
compactor := &SimpleCompactor{Client: client, Options: CompactionOptions{
|
||||
ContextWindowTokens: 100,
|
||||
KeepUserTurns: 1,
|
||||
Threshold: 0.5,
|
||||
}}
|
||||
thinkHigh := &api.ThinkValue{Value: "high"}
|
||||
|
||||
result, err := compactor.MaybeCompact(context.Background(), CompactionRequest{
|
||||
Model: "model",
|
||||
Messages: []api.Message{
|
||||
{Role: "user", Content: "old request"},
|
||||
{Role: "assistant", Content: "old answer"},
|
||||
{Role: "user", Content: "recent request"},
|
||||
},
|
||||
Think: thinkHigh,
|
||||
Force: true,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !result.Compacted || result.Summary != "unset think summary" {
|
||||
t.Fatalf("compaction result = %#v", result)
|
||||
}
|
||||
if len(client.requests) != 3 {
|
||||
t.Fatalf("summary requests = %d, want 3", len(client.requests))
|
||||
}
|
||||
if client.requests[0].Think != thinkHigh {
|
||||
t.Fatalf("first summary request think = %#v, want original", client.requests[0].Think)
|
||||
}
|
||||
if client.requests[1].Think == nil || client.requests[1].Think.Value != false {
|
||||
t.Fatalf("fallback summary request think = %#v, want false", client.requests[1].Think)
|
||||
}
|
||||
if client.requests[2].Think != nil {
|
||||
t.Fatalf("unsupported fallback retry think = %#v, want nil", client.requests[2].Think)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSimpleCompactorKeepsFewerTurnsForShortChats(t *testing.T) {
|
||||
client := &fakeClient{
|
||||
responses: [][]api.ChatResponse{{
|
||||
{Message: api.Message{Role: "assistant", Content: "short summary"}},
|
||||
}},
|
||||
}
|
||||
compactor := &SimpleCompactor{Client: client, Options: CompactionOptions{
|
||||
ContextWindowTokens: 16000,
|
||||
KeepUserTurns: 3,
|
||||
Threshold: 0.5,
|
||||
}}
|
||||
|
||||
result, err := compactor.MaybeCompact(context.Background(), CompactionRequest{
|
||||
ChatID: "chat-1",
|
||||
Model: "model",
|
||||
Messages: []api.Message{
|
||||
{Role: "user", Content: "old request"},
|
||||
{Role: "assistant", Content: "old answer"},
|
||||
{Role: "user", Content: "latest request"},
|
||||
},
|
||||
Latest: api.ChatResponse{Metrics: api.Metrics{PromptEvalCount: 12000}},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !result.Compacted {
|
||||
t.Fatal("expected compaction")
|
||||
}
|
||||
if len(result.Messages) != 3 {
|
||||
t.Fatalf("messages = %#v, want compaction tool pair plus latest request", result.Messages)
|
||||
}
|
||||
assertCompactionSummaryPair(t, result.Messages[:2])
|
||||
if result.Messages[2].Content != "latest request" {
|
||||
t.Fatalf("latest turn was not kept: %#v", result.Messages)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSimpleCompactorCanArchiveWholeShortChat(t *testing.T) {
|
||||
client := &fakeClient{
|
||||
responses: [][]api.ChatResponse{{
|
||||
{Message: api.Message{Role: "assistant", Content: "whole summary"}},
|
||||
}},
|
||||
}
|
||||
compactor := &SimpleCompactor{Client: client, Options: CompactionOptions{
|
||||
ContextWindowTokens: 100,
|
||||
KeepUserTurns: 3,
|
||||
Threshold: 0.5,
|
||||
}}
|
||||
|
||||
result, err := compactor.MaybeCompact(context.Background(), CompactionRequest{
|
||||
ChatID: "chat-1",
|
||||
Model: "model",
|
||||
Messages: []api.Message{
|
||||
{Role: "user", Content: "only request"},
|
||||
{Role: "assistant", Content: "only answer"},
|
||||
},
|
||||
Latest: api.ChatResponse{Metrics: api.Metrics{PromptEvalCount: 75}},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !result.Compacted {
|
||||
t.Fatal("expected compaction")
|
||||
}
|
||||
if len(result.Messages) != 2 {
|
||||
t.Fatalf("messages = %#v, want only compaction tool pair", result.Messages)
|
||||
}
|
||||
assertCompactionSummaryPair(t, result.Messages)
|
||||
}
|
||||
|
||||
func TestSimpleCompactorSkipsBelowThreshold(t *testing.T) {
|
||||
client := &fakeClient{}
|
||||
compactor := &SimpleCompactor{Client: client, Options: CompactionOptions{
|
||||
ContextWindowTokens: 100,
|
||||
Threshold: 0.8,
|
||||
}}
|
||||
|
||||
messages := []api.Message{
|
||||
{Role: "user", Content: "one"},
|
||||
{Role: "user", Content: "two"},
|
||||
{Role: "user", Content: "three"},
|
||||
{Role: "user", Content: "four"},
|
||||
{Role: "user", Content: "five"},
|
||||
}
|
||||
result, err := compactor.MaybeCompact(context.Background(), CompactionRequest{
|
||||
Model: "model",
|
||||
Messages: messages,
|
||||
Latest: api.ChatResponse{Metrics: api.Metrics{PromptEvalCount: 50}},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if result.Compacted {
|
||||
t.Fatal("did not expect compaction")
|
||||
}
|
||||
if result.Due {
|
||||
t.Fatal("below-threshold compaction should not be due")
|
||||
}
|
||||
if len(result.Messages) != len(messages) {
|
||||
t.Fatalf("messages changed below threshold: %#v", result.Messages)
|
||||
}
|
||||
if len(client.requests) != 0 {
|
||||
t.Fatalf("summary requests = %d, want 0", len(client.requests))
|
||||
}
|
||||
}
|
||||
|
||||
func TestSimpleCompactorUsesEstimatedMessagesWhenPromptEvalMissing(t *testing.T) {
|
||||
client := &fakeClient{
|
||||
responses: [][]api.ChatResponse{{
|
||||
{Message: api.Message{Role: "assistant", Content: "estimated summary"}},
|
||||
}},
|
||||
}
|
||||
compactor := &SimpleCompactor{Client: client, Options: CompactionOptions{
|
||||
ContextWindowTokens: 100,
|
||||
KeepUserTurns: 1,
|
||||
Threshold: 0.8,
|
||||
}}
|
||||
|
||||
result, err := compactor.MaybeCompact(context.Background(), CompactionRequest{
|
||||
Model: "model",
|
||||
Messages: []api.Message{
|
||||
{Role: "user", Content: "old request"},
|
||||
{Role: "assistant", Content: "old answer"},
|
||||
{Role: "user", Content: "read large output"},
|
||||
{Role: "assistant", ToolCalls: []api.ToolCall{{
|
||||
ID: "call-1",
|
||||
Function: api.ToolCallFunction{
|
||||
Name: "read",
|
||||
},
|
||||
}}},
|
||||
{Role: "tool", ToolName: "read", ToolCallID: "call-1", Content: strings.Repeat("x", 360)},
|
||||
},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !result.Due || !result.Compacted {
|
||||
t.Fatalf("expected estimate-driven compaction, got %#v", result)
|
||||
}
|
||||
if result.Summary != "estimated summary" {
|
||||
t.Fatalf("summary = %q", result.Summary)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSimpleCompactorEstimateIncludesRequestPreamble(t *testing.T) {
|
||||
compactor := &SimpleCompactor{Client: nil, Options: CompactionOptions{
|
||||
ContextWindowTokens: 100,
|
||||
Threshold: 0.8,
|
||||
}}
|
||||
|
||||
if !compactor.shouldCompact(CompactionRequest{
|
||||
SystemPrompt: strings.Repeat("system ", 360),
|
||||
Messages: []api.Message{{Role: "user", Content: "tiny"}},
|
||||
}) {
|
||||
t.Fatal("system prompt should count toward compaction estimate")
|
||||
}
|
||||
|
||||
if !compactor.shouldCompact(CompactionRequest{
|
||||
Messages: []api.Message{{Role: "user", Content: "tiny"}},
|
||||
Tools: api.Tools{{
|
||||
Type: "function",
|
||||
Function: api.ToolFunction{
|
||||
Name: "verbose_tool",
|
||||
Description: strings.Repeat("description ", 360),
|
||||
},
|
||||
}},
|
||||
}) {
|
||||
t.Fatal("tool definitions should count toward compaction estimate")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCompactionPromptFitsBudgetByTruncatingLargeToolOutput(t *testing.T) {
|
||||
largeToolOutput := strings.Repeat("x", 10_000)
|
||||
body, err := compactionPrompt("", []api.Message{
|
||||
{Role: "user", Content: "what changed?"},
|
||||
{Role: "assistant", ToolCalls: []api.ToolCall{{
|
||||
ID: "call-1",
|
||||
Function: api.ToolCallFunction{
|
||||
Name: "bash",
|
||||
},
|
||||
}}},
|
||||
{Role: "tool", ToolName: "bash", ToolCallID: "call-1", Content: largeToolOutput},
|
||||
}, 300)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if estimateCompactionTokens(body) > 300 {
|
||||
t.Fatalf("compaction prompt tokens = %d, want <= 300", estimateCompactionTokens(body))
|
||||
}
|
||||
if strings.Count(body, "x") >= len(largeToolOutput) {
|
||||
t.Fatal("large tool output was not truncated")
|
||||
}
|
||||
if !strings.Contains(body, "[tool output truncated: showing first ~") {
|
||||
t.Fatalf("truncation marker missing from compaction prompt: %q", body)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCompactionPromptRetruncatesAlreadyTruncatedToolOutput(t *testing.T) {
|
||||
alreadyTruncated := strings.Repeat("x", 7000) + "\n\n[tool output truncated: showing first ~100 tokens and last ~100 tokens; omitted ~99999 tokens. Use a narrower command, line range, or search query if more detail is needed.]\n\n" + strings.Repeat("y", 7000)
|
||||
body, err := compactionPrompt("", []api.Message{
|
||||
{Role: "user", Content: "what changed?"},
|
||||
{Role: "assistant", ToolCalls: []api.ToolCall{{
|
||||
ID: "call-1",
|
||||
Function: api.ToolCallFunction{
|
||||
Name: "bash",
|
||||
},
|
||||
}}},
|
||||
{Role: "tool", ToolName: "bash", ToolCallID: "call-1", Content: alreadyTruncated},
|
||||
}, 300)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if estimateCompactionTokens(body) > 300 {
|
||||
t.Fatalf("compaction prompt tokens = %d, want <= 300", estimateCompactionTokens(body))
|
||||
}
|
||||
if strings.Count(body, "x")+strings.Count(body, "y") >= 14_000 {
|
||||
t.Fatal("already-truncated tool output was not truncated again")
|
||||
}
|
||||
if !strings.Contains(body, "[tool output truncated: showing first ~") {
|
||||
t.Fatalf("truncation marker missing from compaction prompt: %q", body)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCompactionSummaryTextStripsPrefix(t *testing.T) {
|
||||
content := compactionSummaryMessageForTask("worked on branch changes", false)
|
||||
if got := CompactionSummaryText(content); got != "worked on branch changes" {
|
||||
t.Fatalf("summary text = %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCompactionSummaryCanTellModelToContinueTask(t *testing.T) {
|
||||
content := compactionSummaryMessageForTask("worked on branch changes", true)
|
||||
if !strings.Contains(content, CompactionContinueInstruction) {
|
||||
t.Fatalf("summary message missing continue instruction: %q", content)
|
||||
}
|
||||
if got := CompactionSummaryText(content); got != "worked on branch changes" {
|
||||
t.Fatalf("summary text = %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveContextWindowTokensPrefersExplicitNumCtx(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
options map[string]any
|
||||
configured int
|
||||
want int
|
||||
}{
|
||||
{
|
||||
name: "explicit smaller num ctx",
|
||||
options: map[string]any{"num_ctx": 4096},
|
||||
configured: 8192,
|
||||
want: 4096,
|
||||
},
|
||||
{
|
||||
name: "explicit num ctx can exceed configured metadata",
|
||||
options: map[string]any{"num_ctx": 131072},
|
||||
configured: 8192,
|
||||
want: 131072,
|
||||
},
|
||||
{
|
||||
name: "metadata without explicit num ctx",
|
||||
configured: 32768,
|
||||
want: 32768,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if got := ResolveContextWindowTokens(tt.options, tt.configured); got != tt.want {
|
||||
t.Fatalf("ResolveContextWindowTokens() = %d, want %d", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestSimpleCompactorForceCompactsWithoutPromptEvalCount(t *testing.T) {
|
||||
client := &fakeClient{
|
||||
responses: [][]api.ChatResponse{{
|
||||
{Message: api.Message{Role: "assistant", Content: "forced summary"}},
|
||||
}},
|
||||
}
|
||||
compactor := &SimpleCompactor{Client: client, Options: CompactionOptions{
|
||||
ContextWindowTokens: 100,
|
||||
KeepUserTurns: 1,
|
||||
Threshold: 0.8,
|
||||
}}
|
||||
|
||||
result, err := compactor.MaybeCompact(context.Background(), CompactionRequest{
|
||||
Model: "model",
|
||||
Messages: []api.Message{
|
||||
{Role: "user", Content: "old"},
|
||||
{Role: "assistant", Content: "old answer"},
|
||||
{Role: "user", Content: "recent"},
|
||||
},
|
||||
Force: true,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !result.Due || !result.Compacted {
|
||||
t.Fatalf("forced compaction result = %#v", result)
|
||||
}
|
||||
if result.Summary != "forced summary" {
|
||||
t.Fatalf("summary = %q", result.Summary)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSimpleCompactorDefaultsToKeepingThreeUserTurns(t *testing.T) {
|
||||
client := &fakeClient{
|
||||
responses: [][]api.ChatResponse{{
|
||||
{Message: api.Message{Role: "assistant", Content: "summary"}},
|
||||
}},
|
||||
}
|
||||
compactor := &SimpleCompactor{Client: client, Options: CompactionOptions{
|
||||
ContextWindowTokens: 16000,
|
||||
Threshold: 0.5,
|
||||
}}
|
||||
|
||||
result, err := compactor.MaybeCompact(context.Background(), CompactionRequest{
|
||||
ChatID: "chat-1",
|
||||
Model: "model",
|
||||
Messages: []api.Message{
|
||||
{Role: "user", Content: "old"},
|
||||
{Role: "assistant", Content: "old answer"},
|
||||
{Role: "user", Content: "one"},
|
||||
{Role: "assistant", Content: "one answer"},
|
||||
{Role: "user", Content: "two"},
|
||||
{Role: "assistant", Content: "two answer"},
|
||||
{Role: "user", Content: "three"},
|
||||
},
|
||||
Latest: api.ChatResponse{Metrics: api.Metrics{PromptEvalCount: 12000}},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !result.Compacted {
|
||||
t.Fatal("expected compaction")
|
||||
}
|
||||
assertCompactionSummaryPair(t, result.Messages[:2])
|
||||
if got := result.Messages[2].Content; got != "one" {
|
||||
t.Fatalf("first kept turn = %q, want one", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSimpleCompactorCarriesPreviousSummary(t *testing.T) {
|
||||
client := &fakeClient{
|
||||
responses: [][]api.ChatResponse{{
|
||||
{Message: api.Message{Role: "assistant", Content: "new summary"}},
|
||||
}},
|
||||
}
|
||||
compactor := &SimpleCompactor{Client: client, Options: CompactionOptions{
|
||||
ContextWindowTokens: 16000,
|
||||
KeepUserTurns: 1,
|
||||
Threshold: 0.5,
|
||||
}}
|
||||
|
||||
result, err := compactor.MaybeCompact(context.Background(), CompactionRequest{
|
||||
Model: "model",
|
||||
Messages: []api.Message{
|
||||
{Role: "system", Content: CompactionSummaryMessagePrefix + "old summary"},
|
||||
{Role: "user", Content: "old"},
|
||||
{Role: "assistant", Content: "old answer"},
|
||||
{Role: "user", Content: "recent"},
|
||||
},
|
||||
Latest: api.ChatResponse{Metrics: api.Metrics{PromptEvalCount: 12000}},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !result.Compacted {
|
||||
t.Fatal("expected compaction")
|
||||
}
|
||||
if !strings.Contains(client.requests[0].Messages[1].Content, "Previous summary:\nold summary") {
|
||||
t.Fatalf("previous summary missing from request: %q", client.requests[0].Messages[1].Content)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSimpleCompactorCarriesPreviousToolSummaryAndPlacesNewSummaryBeforeKeptSuffix(t *testing.T) {
|
||||
client := &fakeClient{
|
||||
responses: [][]api.ChatResponse{{
|
||||
{Message: api.Message{Role: "assistant", Content: "new summary"}},
|
||||
}},
|
||||
}
|
||||
compactor := &SimpleCompactor{Client: client, Options: CompactionOptions{
|
||||
ContextWindowTokens: 16000,
|
||||
KeepUserTurns: 1,
|
||||
Threshold: 0.5,
|
||||
}}
|
||||
|
||||
messages := []api.Message{
|
||||
{Role: "user", Content: "kept before old summary"},
|
||||
CompactionSummaryMessages("old summary", false)[0],
|
||||
CompactionSummaryMessages("old summary", false)[1],
|
||||
{Role: "user", Content: "latest request"},
|
||||
}
|
||||
result, err := compactor.MaybeCompact(context.Background(), CompactionRequest{
|
||||
Model: "model",
|
||||
Messages: messages,
|
||||
Latest: api.ChatResponse{Metrics: api.Metrics{PromptEvalCount: 12000}},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !result.Compacted {
|
||||
t.Fatal("expected compaction")
|
||||
}
|
||||
if !strings.Contains(client.requests[0].Messages[1].Content, "Previous summary:\nold summary") {
|
||||
t.Fatalf("previous summary missing from request: %q", client.requests[0].Messages[1].Content)
|
||||
}
|
||||
if len(result.Messages) != 3 {
|
||||
t.Fatalf("messages = %#v, want compaction pair plus latest request", result.Messages)
|
||||
}
|
||||
assertCompactionSummaryPair(t, result.Messages[:2])
|
||||
if result.Messages[2].Content != "latest request" {
|
||||
t.Fatalf("kept suffix = %#v", result.Messages)
|
||||
}
|
||||
}
|
||||
+177
@@ -0,0 +1,177 @@
|
||||
package agent
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
|
||||
"github.com/ollama/ollama/api"
|
||||
)
|
||||
|
||||
type EventType string
|
||||
|
||||
const (
|
||||
EventMessageDelta EventType = "message_delta"
|
||||
EventThinkingDelta EventType = "thinking_delta"
|
||||
EventToolCallDetected EventType = "tool_call_detected"
|
||||
EventToolStarted EventType = "tool_started"
|
||||
EventToolFinished EventType = "tool_finished"
|
||||
EventCompactionStarted EventType = "compaction_started"
|
||||
EventCompactionProgress EventType = "compaction_progress"
|
||||
EventCompacted EventType = "compacted"
|
||||
EventCompactionSkipped EventType = "compaction_skipped"
|
||||
EventRunFinished EventType = "run_finished"
|
||||
EventError EventType = "error"
|
||||
)
|
||||
|
||||
// ToolStatus is the typed lifecycle state for a tool call, carried on
|
||||
// Event.ToolStatus for tool events.
|
||||
type ToolStatus string
|
||||
|
||||
const (
|
||||
ToolStatusRunning ToolStatus = "running"
|
||||
ToolStatusDone ToolStatus = "done"
|
||||
ToolStatusFailed ToolStatus = "failed"
|
||||
ToolStatusDenied ToolStatus = "denied"
|
||||
ToolStatusDisabled ToolStatus = "disabled"
|
||||
ToolStatusSkipped ToolStatus = "skipped"
|
||||
)
|
||||
|
||||
// RunStatus is the typed terminal outcome of a run, carried on Event.Status for
|
||||
// run_finished events.
|
||||
type RunStatus string
|
||||
|
||||
const (
|
||||
RunStatusDone RunStatus = "done"
|
||||
RunStatusDenied RunStatus = "denied"
|
||||
RunStatusCanceled RunStatus = "canceled"
|
||||
)
|
||||
|
||||
// CompactionTrigger is the typed reason a compaction ran or was attempted,
|
||||
// carried on Event.CompactionTrigger for compaction events.
|
||||
type CompactionTrigger string
|
||||
|
||||
const (
|
||||
CompactionTriggerForce CompactionTrigger = "force"
|
||||
CompactionTriggerPromptEval CompactionTrigger = "prompt_eval"
|
||||
CompactionTriggerEstimate CompactionTrigger = "estimate"
|
||||
CompactionTriggerToolOutput CompactionTrigger = "tool_output"
|
||||
CompactionTriggerError CompactionTrigger = "error"
|
||||
CompactionTriggerDue CompactionTrigger = "due"
|
||||
)
|
||||
|
||||
type Event struct {
|
||||
Type EventType `json:"type"`
|
||||
RunID string `json:"runId,omitempty"`
|
||||
ChatID string `json:"chatId,omitempty"`
|
||||
Model string `json:"model,omitempty"`
|
||||
Status RunStatus `json:"status,omitempty"`
|
||||
ToolStatus ToolStatus `json:"toolStatus,omitempty"`
|
||||
CompactionTrigger CompactionTrigger `json:"compactionTrigger,omitempty"`
|
||||
ToolCallID string `json:"toolCallId,omitempty"`
|
||||
ToolName string `json:"toolName,omitempty"`
|
||||
WorkingDir string `json:"workingDir,omitempty"`
|
||||
Content string `json:"content,omitempty"`
|
||||
Thinking string `json:"thinking,omitempty"`
|
||||
ToolCalls []api.ToolCall `json:"toolCalls,omitempty"`
|
||||
Messages []api.Message `json:"messages,omitempty"`
|
||||
Args map[string]any `json:"args,omitempty"`
|
||||
Tokens int `json:"tokens,omitempty"`
|
||||
Error string `json:"error,omitempty"`
|
||||
}
|
||||
|
||||
type EventSink interface {
|
||||
Emit(Event) error
|
||||
}
|
||||
|
||||
type EventSinkFunc func(Event) error
|
||||
|
||||
func (fn EventSinkFunc) Emit(event Event) error {
|
||||
if fn == nil {
|
||||
return nil
|
||||
}
|
||||
return fn(event)
|
||||
}
|
||||
|
||||
// eventMetadata carries the run identification fields shared by all events.
|
||||
type eventMetadata struct {
|
||||
runID string
|
||||
chatID string
|
||||
model string
|
||||
}
|
||||
|
||||
func newEventMetadata(runID string, opts RunOptions) eventMetadata {
|
||||
return eventMetadata{runID: runID, chatID: opts.ChatID, model: opts.Model}
|
||||
}
|
||||
|
||||
func newMessageDelta(m eventMetadata, content string) Event {
|
||||
return Event{Type: EventMessageDelta, RunID: m.runID, ChatID: m.chatID, Model: m.model, Content: content}
|
||||
}
|
||||
|
||||
func newThinkingDelta(m eventMetadata, thinking string) Event {
|
||||
return Event{Type: EventThinkingDelta, RunID: m.runID, ChatID: m.chatID, Model: m.model, Thinking: thinking}
|
||||
}
|
||||
|
||||
func newToolCallDetected(m eventMetadata, calls []api.ToolCall) Event {
|
||||
return Event{Type: EventToolCallDetected, RunID: m.runID, ChatID: m.chatID, Model: m.model, ToolCalls: calls}
|
||||
}
|
||||
|
||||
func newToolStarted(m eventMetadata, callID, toolName, workingDir string, args map[string]any) Event {
|
||||
return Event{Type: EventToolStarted, RunID: m.runID, ChatID: m.chatID, Model: m.model, ToolStatus: ToolStatusRunning, ToolCallID: callID, ToolName: toolName, WorkingDir: workingDir, Args: args}
|
||||
}
|
||||
|
||||
func newToolFinished(m eventMetadata, status ToolStatus, callID, toolName, workingDir string, args map[string]any, content, errMsg string) Event {
|
||||
ev := Event{Type: EventToolFinished, RunID: m.runID, ChatID: m.chatID, Model: m.model, ToolStatus: status, ToolCallID: callID, ToolName: toolName, WorkingDir: workingDir, Args: args, Content: content}
|
||||
if errMsg != "" {
|
||||
ev.Error = errMsg
|
||||
}
|
||||
return ev
|
||||
}
|
||||
|
||||
func newRunFinished(m eventMetadata, status RunStatus) Event {
|
||||
return Event{Type: EventRunFinished, RunID: m.runID, ChatID: m.chatID, Model: m.model, Status: status}
|
||||
}
|
||||
|
||||
func newErrorEvent(m eventMetadata, errMsg string) Event {
|
||||
return Event{Type: EventError, RunID: m.runID, ChatID: m.chatID, Model: m.model, Error: errMsg}
|
||||
}
|
||||
|
||||
func newCompactionProgress(m eventMetadata, tokens int) Event {
|
||||
return Event{Type: EventCompactionProgress, RunID: m.runID, ChatID: m.chatID, Model: m.model, Tokens: tokens}
|
||||
}
|
||||
|
||||
func newCompactionStarted(m eventMetadata, trigger CompactionTrigger) Event {
|
||||
return Event{Type: EventCompactionStarted, RunID: m.runID, ChatID: m.chatID, Model: m.model, CompactionTrigger: trigger}
|
||||
}
|
||||
|
||||
func newCompactionSkipped(m eventMetadata, trigger CompactionTrigger, content string) Event {
|
||||
return Event{Type: EventCompactionSkipped, RunID: m.runID, ChatID: m.chatID, Model: m.model, CompactionTrigger: trigger, Content: content}
|
||||
}
|
||||
|
||||
func newCompacted(m eventMetadata, messages []api.Message, trigger CompactionTrigger, content string) Event {
|
||||
return Event{Type: EventCompacted, RunID: m.runID, ChatID: m.chatID, Model: m.model, CompactionTrigger: trigger, Content: content, Messages: messages}
|
||||
}
|
||||
|
||||
func (s *Session) emit(event Event) error {
|
||||
if s == nil {
|
||||
return nil
|
||||
}
|
||||
var errs []error
|
||||
for _, sink := range s.EventSinks {
|
||||
if sink == nil {
|
||||
continue
|
||||
}
|
||||
if err := sink.Emit(event); err != nil {
|
||||
errs = append(errs, err)
|
||||
}
|
||||
}
|
||||
return errors.Join(errs...)
|
||||
}
|
||||
|
||||
func (s *Session) emitIgnoringCanceled(ctx context.Context, event Event) error {
|
||||
err := s.emit(event)
|
||||
if err != nil && ctx != nil && ctx.Err() != nil {
|
||||
//nolint:nilerr // Event sinks may close during cancellation; cancellation is not a user-facing emit failure.
|
||||
return nil
|
||||
}
|
||||
return err
|
||||
}
|
||||
@@ -0,0 +1,104 @@
|
||||
package agent
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"sort"
|
||||
|
||||
"github.com/ollama/ollama/api"
|
||||
)
|
||||
|
||||
type ToolContext struct {
|
||||
WorkingDir string
|
||||
}
|
||||
|
||||
type ToolResult struct {
|
||||
Content string
|
||||
WorkingDir string
|
||||
}
|
||||
|
||||
type Tool interface {
|
||||
Name() string
|
||||
Description() string
|
||||
Schema() api.ToolFunction
|
||||
Execute(context.Context, ToolContext, map[string]any) (ToolResult, error)
|
||||
}
|
||||
|
||||
type ApprovalRequired interface {
|
||||
RequiresApproval(map[string]any) bool
|
||||
}
|
||||
|
||||
// ScopedTool is implemented by tools that need per-invocation approval
|
||||
// scoping beyond the tool name (e.g. shell commands scoped to the exact
|
||||
// command string). Tools that don't implement this are scoped by name only.
|
||||
type ScopedTool interface {
|
||||
ApprovalScope(args map[string]any) string
|
||||
}
|
||||
|
||||
type Registry struct {
|
||||
tools map[string]Tool
|
||||
}
|
||||
|
||||
func (r *Registry) Register(tool Tool) {
|
||||
if r == nil || tool == nil {
|
||||
return
|
||||
}
|
||||
if r.tools == nil {
|
||||
r.tools = make(map[string]Tool)
|
||||
}
|
||||
r.tools[tool.Name()] = tool
|
||||
}
|
||||
|
||||
func (r *Registry) Get(name string) (Tool, bool) {
|
||||
if r == nil {
|
||||
return nil, false
|
||||
}
|
||||
tool, ok := r.tools[name]
|
||||
return tool, ok
|
||||
}
|
||||
|
||||
func (r *Registry) Names() []string {
|
||||
if r == nil {
|
||||
return nil
|
||||
}
|
||||
names := make([]string, 0, len(r.tools))
|
||||
for name := range r.tools {
|
||||
names = append(names, name)
|
||||
}
|
||||
sort.Strings(names)
|
||||
return names
|
||||
}
|
||||
|
||||
func (r *Registry) Tools() api.Tools {
|
||||
if r == nil {
|
||||
return nil
|
||||
}
|
||||
names := r.Names()
|
||||
apiTools := make(api.Tools, 0, len(names))
|
||||
for _, name := range names {
|
||||
tool := r.tools[name]
|
||||
apiTools = append(apiTools, api.Tool{
|
||||
Type: "function",
|
||||
Function: tool.Schema(),
|
||||
})
|
||||
}
|
||||
return apiTools
|
||||
}
|
||||
|
||||
func (r *Registry) Execute(ctx context.Context, toolCtx ToolContext, call api.ToolCall) (ToolResult, error) {
|
||||
tool, ok := r.Get(call.Function.Name)
|
||||
if !ok {
|
||||
return ToolResult{}, fmt.Errorf("unknown tool: %s", call.Function.Name)
|
||||
}
|
||||
return tool.Execute(ctx, toolCtx, call.Function.Arguments.ToMap())
|
||||
}
|
||||
|
||||
func ToolRequiresApproval(tool Tool, args map[string]any) bool {
|
||||
if tool == nil {
|
||||
return false
|
||||
}
|
||||
if t, ok := tool.(ApprovalRequired); ok {
|
||||
return t.RequiresApproval(args)
|
||||
}
|
||||
return false
|
||||
}
|
||||
+1099
File diff suppressed because it is too large.
Load diff
File diff suppressed because it is too large.
Load diff
@@ -0,0 +1,57 @@
|
||||
package agent
|
||||
|
||||
import (
|
||||
"context"
|
||||
"strings"
|
||||
|
||||
"github.com/google/uuid"
|
||||
|
||||
"github.com/ollama/ollama/api"
|
||||
)
|
||||
|
||||
// activateSkill loads opts.SkillName from the catalog and injects a synthetic
|
||||
// assistant tool call plus tool result before the first model request, so the
|
||||
// transcript looks like a real skill tool invocation. It emits the same
|
||||
// tool_call_detected -> tool_started -> tool_finished lifecycle the model path
|
||||
// uses, and returns the messages to prepend. A blank SkillName is a no-op.
|
||||
func (s *Session) activateSkill(ctx context.Context, runID string, opts RunOptions) ([]api.Message, error) {
|
||||
name := strings.TrimSpace(opts.SkillName)
|
||||
if name == "" {
|
||||
return nil, nil
|
||||
}
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return nil, ctx.Err()
|
||||
default:
|
||||
}
|
||||
skill, err := s.Skills.Load(name)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
args := api.NewToolCallFunctionArguments()
|
||||
args.Set("name", skill.Name)
|
||||
call := api.ToolCall{
|
||||
ID: "call_skill_" + uuid.NewString(),
|
||||
Function: api.ToolCallFunction{Name: "skill", Arguments: args},
|
||||
}
|
||||
result := api.Message{
|
||||
Role: "tool",
|
||||
ToolName: "skill",
|
||||
ToolCallID: call.ID,
|
||||
Content: skill.Content(),
|
||||
}
|
||||
meta := newEventMetadata(runID, opts)
|
||||
if err := s.emit(newToolCallDetected(meta, []api.ToolCall{call})); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := s.emit(newToolStarted(meta, call.ID, "skill", s.currentWorkingDir(), args.ToMap())); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := s.emitIgnoringCanceled(ctx, newToolFinished(meta, ToolStatusDone, call.ID, "skill", s.currentWorkingDir(), args.ToMap(), result.Content, "")); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return []api.Message{
|
||||
{Role: "assistant", ToolCalls: []api.ToolCall{call}},
|
||||
result,
|
||||
}, nil
|
||||
}
|
||||
@@ -0,0 +1,74 @@
|
||||
package agent
|
||||
|
||||
import (
|
||||
"context"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/ollama/ollama/api"
|
||||
)
|
||||
|
||||
type skillTestClient struct{ requests []*api.ChatRequest }
|
||||
|
||||
func (c *skillTestClient) Chat(_ context.Context, req *api.ChatRequest, fn api.ChatResponseFunc) error {
|
||||
c.requests = append(c.requests, req)
|
||||
return fn(api.ChatResponse{Message: api.Message{Role: "assistant", Content: "Done."}})
|
||||
}
|
||||
|
||||
func testSkillCatalog(t *testing.T) *SkillCatalog {
|
||||
t.Helper()
|
||||
dir := t.TempDir()
|
||||
path := filepath.Join(dir, "release-notes")
|
||||
if err := os.Mkdir(path, 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := os.WriteFile(filepath.Join(path, "SKILL.md"), []byte("---\nname: release-notes\ndescription: Draft release notes.\n---\nUse concise bullets."), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
catalog, err := DiscoverSkills(dir)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return catalog
|
||||
}
|
||||
|
||||
func TestSessionSkillActivationPreservesCallAndResultOrder(t *testing.T) {
|
||||
catalog := testSkillCatalog(t)
|
||||
client := &skillTestClient{}
|
||||
events := &recordingEventSink{}
|
||||
result, err := (&Session{Client: client, Skills: catalog, EventSinks: []EventSink{events}}).Run(context.Background(), RunOptions{
|
||||
Model: "test",
|
||||
NewMessages: []api.Message{{Role: "user", Content: "draft release notes"}},
|
||||
SkillName: "release-notes",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(result.Messages) != 4 {
|
||||
t.Fatalf("transcript = %#v", result.Messages)
|
||||
}
|
||||
call, toolTranscript := result.Messages[1], result.Messages[2]
|
||||
if call.Role != "assistant" || len(call.ToolCalls) != 1 || call.ToolCalls[0].Function.Name != "skill" || !strings.HasPrefix(call.ToolCalls[0].ID, "call_skill_") {
|
||||
t.Fatalf("call message = %#v", call)
|
||||
}
|
||||
if toolTranscript.Role != "tool" || toolTranscript.ToolName != "skill" || toolTranscript.ToolCallID != call.ToolCalls[0].ID || !strings.Contains(toolTranscript.Content, "Use concise bullets.") {
|
||||
t.Fatalf("tool result = %#v", toolTranscript)
|
||||
}
|
||||
if len(client.requests) != 1 || len(client.requests[0].Messages) != 3 || client.requests[0].Messages[2].ToolCallID != call.ToolCalls[0].ID {
|
||||
t.Fatalf("model request did not preserve transcript: %#v", client.requests)
|
||||
}
|
||||
var skillEvents []EventType
|
||||
for _, event := range events.events {
|
||||
if event.ToolName == "skill" || event.Type == EventToolCallDetected {
|
||||
skillEvents = append(skillEvents, event.Type)
|
||||
}
|
||||
}
|
||||
if len(skillEvents) < 3 {
|
||||
t.Fatalf("skill event order = %#v, want tool_call_detected,tool_started,tool_finished", skillEvents)
|
||||
}
|
||||
if got, want := strings.Join([]string{string(skillEvents[0]), string(skillEvents[1]), string(skillEvents[2])}, ","), "tool_call_detected,tool_started,tool_finished"; got != want {
|
||||
t.Fatalf("skill event order = %#v, want %s", skillEvents, want)
|
||||
}
|
||||
}
|
||||
+438
@@ -0,0 +1,438 @@
|
||||
package agent
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"io/fs"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"regexp"
|
||||
"sort"
|
||||
"strings"
|
||||
|
||||
"gopkg.in/yaml.v3"
|
||||
)
|
||||
|
||||
const (
|
||||
// SkillsDirEnv overrides the user-level Ollama-owned skills directory. The
|
||||
// cross-client .agents/skills/ convention and project-level .ollama/skills/
|
||||
// are also scanned (see LoadDefaultSkills); on a name collision, Ollama-owned
|
||||
// directories take precedence over .agents/skills/, and project-level takes
|
||||
// precedence over user-level.
|
||||
SkillsDirEnv = "OLLAMA_SKILLS"
|
||||
skillFilename = "SKILL.md"
|
||||
maxSkillBytes = 1 << 20
|
||||
|
||||
bundledSkillCreatorName = "skill-creator"
|
||||
bundledSkillCreatorContent = `---
|
||||
name: skill-creator
|
||||
description: Create or improve reusable skills. Use when the user wants a reusable skill, asks how to author SKILL.md, or needs help installing a skill.
|
||||
---
|
||||
|
||||
# Create a skill
|
||||
|
||||
Create a focused, reusable instruction package. Treat a skill as guidance for the model, not as a way to gain new permissions or bypass safety controls.
|
||||
|
||||
## Choose the location
|
||||
|
||||
Create user skills beside this one. The skill directory shown in the loaded skill context is this skill's location; its parent is the user skill root. This bundled skill normally lives at ~/.ollama/skills/skill-creator, so new user skills normally go at ~/.ollama/skills/<skill-name>/SKILL.md.
|
||||
|
||||
Use a project-local skill directory only when the user asks to keep the skill with that project. Do not overwrite an existing skill without the user's approval. New and changed skills are discovered when the agent starts, so tell the user to begin a new agent session afterward.
|
||||
|
||||
## Follow the required shape
|
||||
|
||||
Use the directory name as the skill name. Use lowercase letters, numbers, and single hyphens only. Keep the name short and no longer than 64 characters.
|
||||
|
||||
Every skill needs a SKILL.md with YAML frontmatter followed by Markdown instructions:
|
||||
|
||||
~~~md
|
||||
---
|
||||
name: release-notes
|
||||
description: Draft concise release notes from completed changes. Use when the user asks for a changelog, release notes, or GitHub release copy.
|
||||
---
|
||||
|
||||
# Draft release notes
|
||||
|
||||
Write the workflow here.
|
||||
~~~
|
||||
|
||||
Require a non-empty description that says both what the skill does and when to use it. Keep the body procedural and concise. Put detailed schemas, long examples, and variant-specific guidance in references/ only when the skill needs them.
|
||||
|
||||
Use scripts/ for repeatable or fragile operations that benefit from deterministic execution. Use assets/ for files that belong in generated output. Do not add README files, changelogs, or setup notes that do not help the model perform the task.
|
||||
|
||||
## Create safely
|
||||
|
||||
1. Identify the repeated task, expected inputs, and useful output.
|
||||
2. Choose the smallest name and description that reliably trigger the skill.
|
||||
3. Create the folder and SKILL.md; add resources only when they remove real repeated work.
|
||||
4. Re-read the completed file and verify its frontmatter, directory-name match, and relative resource paths.
|
||||
5. Tell the user where it was created and that a new agent session will discover it.
|
||||
|
||||
Skills provide instructions only. They do not grant filesystem, network, shell, or approval privileges, and they do not make a tool available. Use only the tools that are actually available, follow their normal approval rules, and ask before actions that need user authorization.
|
||||
`
|
||||
)
|
||||
|
||||
var skillName = regexp.MustCompile(`^[a-z0-9]+(?:-[a-z0-9]+)*$`)
|
||||
|
||||
// SkillsDir returns the canonical runtime-owned skill directory.
|
||||
func SkillsDir() (string, error) {
|
||||
if path := strings.TrimSpace(os.Getenv(SkillsDirEnv)); path != "" {
|
||||
return filepath.Abs(path)
|
||||
}
|
||||
if xdg := strings.TrimSpace(os.Getenv("XDG_CONFIG_HOME")); xdg != "" {
|
||||
return filepath.Join(xdg, "ollama", "skills"), nil
|
||||
}
|
||||
home, err := os.UserHomeDir()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return filepath.Join(home, ".ollama", "skills"), nil
|
||||
}
|
||||
|
||||
// Skill is a validated, loadable instruction set. It never grants tool
|
||||
// permissions; it is supplied to the model as ordinary tool-result content.
|
||||
type Skill struct {
|
||||
Name string
|
||||
Description string
|
||||
Instructions string
|
||||
Path string
|
||||
}
|
||||
|
||||
func (s Skill) Content() string {
|
||||
var b strings.Builder
|
||||
fmt.Fprintf(&b, "<skill name=%q>\n%s\n", s.Name, strings.TrimSpace(s.Instructions))
|
||||
if s.Path != "" {
|
||||
dir := filepath.Dir(s.Path)
|
||||
fmt.Fprintf(&b, "Skill directory: %s\n", dir)
|
||||
b.WriteString("Relative paths in this skill are relative to the skill directory.\n")
|
||||
}
|
||||
if resources := s.resources(); len(resources) > 0 {
|
||||
b.WriteString("<skill_resources>\n")
|
||||
for _, r := range resources {
|
||||
fmt.Fprintf(&b, " <file>%s</file>\n", r)
|
||||
}
|
||||
b.WriteString("</skill_resources>\n")
|
||||
}
|
||||
b.WriteString("</skill>")
|
||||
return b.String()
|
||||
}
|
||||
|
||||
// resources lists bundled files one level deep under scripts/, references/,
|
||||
// and assets/ without reading them, so the model can load them on demand.
|
||||
func (s Skill) resources() []string {
|
||||
if s.Path == "" {
|
||||
return nil
|
||||
}
|
||||
dir := filepath.Dir(s.Path)
|
||||
var resources []string
|
||||
for _, sub := range []string{"scripts", "references", "assets"} {
|
||||
entries, err := os.ReadDir(filepath.Join(dir, sub))
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
for _, e := range entries {
|
||||
if e.IsDir() {
|
||||
continue
|
||||
}
|
||||
resources = append(resources, sub+"/"+e.Name())
|
||||
}
|
||||
}
|
||||
sort.Strings(resources)
|
||||
return resources
|
||||
}
|
||||
|
||||
// SkillCatalog contains valid skills and diagnostics for ignored invalid
|
||||
// entries, so one malformed skill cannot hide the rest.
|
||||
type SkillCatalog struct {
|
||||
dir string
|
||||
skills map[string]Skill
|
||||
diagnostics []error
|
||||
}
|
||||
|
||||
func DiscoverSkills(dir string) (*SkillCatalog, error) {
|
||||
dir, err := filepath.Abs(strings.TrimSpace(dir))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
catalog := &SkillCatalog{dir: dir, skills: make(map[string]Skill)}
|
||||
entries, err := os.ReadDir(dir)
|
||||
if errors.Is(err, fs.ErrNotExist) {
|
||||
return catalog, nil
|
||||
}
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("read skills directory: %w", err)
|
||||
}
|
||||
for _, entry := range entries {
|
||||
name := entry.Name()
|
||||
// Follow symlinks so users can point at shared skill repositories.
|
||||
// The link name (not the target) is the canonical skill name.
|
||||
info, err := os.Stat(filepath.Join(dir, name))
|
||||
if err != nil {
|
||||
if errors.Is(err, fs.ErrNotExist) {
|
||||
continue
|
||||
}
|
||||
catalog.diagnostics = append(catalog.diagnostics, fmt.Errorf("skill %q: %w", name, err))
|
||||
continue
|
||||
}
|
||||
if !info.IsDir() {
|
||||
continue
|
||||
}
|
||||
if !skillName.MatchString(name) {
|
||||
catalog.diagnostics = append(catalog.diagnostics, fmt.Errorf("invalid skill directory %q", name))
|
||||
continue
|
||||
}
|
||||
skill, err := parseSkill(filepath.Join(dir, name, skillFilename), name)
|
||||
if errors.Is(err, fs.ErrNotExist) {
|
||||
continue
|
||||
}
|
||||
if err != nil {
|
||||
catalog.diagnostics = append(catalog.diagnostics, err)
|
||||
continue
|
||||
}
|
||||
catalog.skills[skill.Name] = skill
|
||||
}
|
||||
return catalog, nil
|
||||
}
|
||||
|
||||
// LoadDefaultSkills discovers skills from the spec's scopes, merged with
|
||||
// deterministic precedence. Roots are scanned lowest-precedence first so later
|
||||
// roots override earlier ones on name collisions (recording a diagnostic):
|
||||
//
|
||||
// 1. ~/.agents/skills/ (user, cross-client)
|
||||
// 2. user Ollama skills dir (user, Ollama-owned; SkillsDir)
|
||||
// 3. <project>/.agents/skills/ (project, cross-client)
|
||||
// 4. <project>/.ollama/skills/ (project, Ollama-owned)
|
||||
//
|
||||
// Project-level overrides user-level, and within a scope Ollama-owned
|
||||
// directories override .agents/skills/. projectDir is the agent's working
|
||||
// directory at startup (discovery is a session-start snapshot per the spec).
|
||||
func LoadDefaultSkills(projectDir string) (*SkillCatalog, error) {
|
||||
roots, err := defaultSkillRoots(projectDir)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
catalog := &SkillCatalog{skills: make(map[string]Skill)}
|
||||
bundled, err := bundledSkillCreator()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
catalog.skills[bundled.Name] = bundled
|
||||
if err := installBundledSkillCreator(); err != nil {
|
||||
catalog.diagnostics = append(catalog.diagnostics, err)
|
||||
}
|
||||
for _, root := range roots {
|
||||
sub, err := DiscoverSkills(root.path)
|
||||
if err != nil {
|
||||
catalog.diagnostics = append(catalog.diagnostics, fmt.Errorf("discover skills in %s: %w", root.path, err))
|
||||
continue
|
||||
}
|
||||
catalog.diagnostics = append(catalog.diagnostics, sub.diagnostics...)
|
||||
for _, skill := range sub.skills {
|
||||
// Name collisions across roots are expected precedence resolution,
|
||||
// not errors: later (higher-precedence) roots legitimately override
|
||||
// earlier ones. The skill is still loaded; no diagnostic needed.
|
||||
catalog.skills[skill.Name] = skill
|
||||
}
|
||||
}
|
||||
return catalog, nil
|
||||
}
|
||||
|
||||
func bundledSkillCreator() (Skill, error) {
|
||||
skill, err := parseSkillContent("", bundledSkillCreatorName, bundledSkillCreatorContent)
|
||||
if err != nil {
|
||||
return Skill{}, fmt.Errorf("load bundled %s skill: %w", bundledSkillCreatorName, err)
|
||||
}
|
||||
return skill, nil
|
||||
}
|
||||
|
||||
func installBundledSkillCreator() error {
|
||||
dir, err := SkillsDir()
|
||||
if err != nil {
|
||||
return fmt.Errorf("resolve bundled skill directory: %w", err)
|
||||
}
|
||||
path := filepath.Join(dir, bundledSkillCreatorName, skillFilename)
|
||||
if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil {
|
||||
return fmt.Errorf("create bundled skill directory: %w", err)
|
||||
}
|
||||
contents, err := os.ReadFile(path)
|
||||
if err == nil && string(contents) == bundledSkillCreatorContent {
|
||||
return nil
|
||||
}
|
||||
if err != nil && !errors.Is(err, fs.ErrNotExist) {
|
||||
return fmt.Errorf("read bundled skill: %w", err)
|
||||
}
|
||||
if err := os.WriteFile(path, []byte(bundledSkillCreatorContent), 0o644); err != nil {
|
||||
return fmt.Errorf("write bundled skill: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
type skillRoot struct {
|
||||
path string
|
||||
}
|
||||
|
||||
// defaultSkillRoots returns skill directories ordered lowest- to
|
||||
// highest-precedence. Non-existent directories are scanned harmlessly
|
||||
// (DiscoverSkills skips them).
|
||||
func defaultSkillRoots(projectDir string) ([]skillRoot, error) {
|
||||
var roots []skillRoot
|
||||
|
||||
if home, err := os.UserHomeDir(); err == nil && home != "" {
|
||||
roots = append(roots, skillRoot{path: filepath.Join(home, ".agents", "skills")})
|
||||
}
|
||||
|
||||
userOllama, err := SkillsDir()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
roots = append(roots, skillRoot{path: userOllama})
|
||||
|
||||
projectDir = strings.TrimSpace(projectDir)
|
||||
if projectDir != "" {
|
||||
if abs, err := filepath.Abs(projectDir); err == nil {
|
||||
roots = append(roots,
|
||||
skillRoot{path: filepath.Join(abs, ".agents", "skills")},
|
||||
skillRoot{path: filepath.Join(abs, ".ollama", "skills")},
|
||||
)
|
||||
}
|
||||
}
|
||||
return roots, nil
|
||||
}
|
||||
|
||||
func (c *SkillCatalog) Dir() string {
|
||||
if c == nil {
|
||||
return ""
|
||||
}
|
||||
return c.dir
|
||||
}
|
||||
|
||||
func (c *SkillCatalog) List() []Skill {
|
||||
if c == nil {
|
||||
return nil
|
||||
}
|
||||
list := make([]Skill, 0, len(c.skills))
|
||||
for _, skill := range c.skills {
|
||||
list = append(list, skill)
|
||||
}
|
||||
sort.Slice(list, func(i, j int) bool { return list[i].Name < list[j].Name })
|
||||
return list
|
||||
}
|
||||
|
||||
func (c *SkillCatalog) Diagnostics() []error {
|
||||
if c == nil {
|
||||
return nil
|
||||
}
|
||||
return append([]error(nil), c.diagnostics...)
|
||||
}
|
||||
|
||||
func (c *SkillCatalog) Load(name string) (Skill, error) {
|
||||
name = strings.TrimSpace(name)
|
||||
if !skillName.MatchString(name) {
|
||||
return Skill{}, fmt.Errorf("invalid skill name %q", name)
|
||||
}
|
||||
if c == nil {
|
||||
return Skill{}, errors.New("skills are unavailable")
|
||||
}
|
||||
skill, ok := c.skills[name]
|
||||
if !ok {
|
||||
return Skill{}, fmt.Errorf("skill %q not found in %s", name, c.dir)
|
||||
}
|
||||
return skill, nil
|
||||
}
|
||||
|
||||
// SystemContext advertises the catalog without expanding full instructions in
|
||||
// every request. The skill call is the explicit loading boundary.
|
||||
func (c *SkillCatalog) SystemContext() string {
|
||||
list := c.List()
|
||||
if len(list) == 0 {
|
||||
return ""
|
||||
}
|
||||
lines := []string{"<available_skills>"}
|
||||
for _, skill := range list {
|
||||
description := skill.Description
|
||||
if description == "" {
|
||||
description = "No description provided."
|
||||
}
|
||||
lines = append(lines, fmt.Sprintf("- %s: %s", skill.Name, description))
|
||||
}
|
||||
lines = append(lines, "</available_skills>", "Load a matching skill with the skill tool before following its instructions. Skills only provide instructions; use ordinary tools for filesystem or network access, with their normal approval rules.")
|
||||
return strings.Join(lines, "\n")
|
||||
}
|
||||
|
||||
func parseSkill(path, directoryName string) (Skill, error) {
|
||||
// Stat (not Lstat) so a symlinked SKILL.md resolves to its target file.
|
||||
info, err := os.Stat(path)
|
||||
if err != nil {
|
||||
return Skill{}, err
|
||||
}
|
||||
if !info.Mode().IsRegular() {
|
||||
return Skill{}, fmt.Errorf("skill %q: %s is not a regular file", directoryName, skillFilename)
|
||||
}
|
||||
if info.Size() > maxSkillBytes {
|
||||
return Skill{}, fmt.Errorf("skill %q: %s exceeds %d bytes", directoryName, skillFilename, maxSkillBytes)
|
||||
}
|
||||
data, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
return Skill{}, fmt.Errorf("read skill %q: %w", directoryName, err)
|
||||
}
|
||||
return parseSkillContent(path, directoryName, string(data))
|
||||
}
|
||||
|
||||
func parseSkillContent(path, directoryName, input string) (Skill, error) {
|
||||
instructions := strings.TrimSpace(input)
|
||||
if instructions == "" {
|
||||
return Skill{}, fmt.Errorf("skill %q: %s is empty", directoryName, skillFilename)
|
||||
}
|
||||
if !strings.HasPrefix(instructions, "---\n") && !strings.HasPrefix(instructions, "---\r\n") {
|
||||
return Skill{}, fmt.Errorf("skill %q: missing YAML front matter", directoryName)
|
||||
}
|
||||
metadata, body, err := skillFrontMatter(instructions)
|
||||
if err != nil {
|
||||
return Skill{}, fmt.Errorf("skill %q: %w", directoryName, err)
|
||||
}
|
||||
if metadata.Name == "" {
|
||||
return Skill{}, fmt.Errorf("skill %q: front matter requires name", directoryName)
|
||||
}
|
||||
if metadata.Description == "" {
|
||||
return Skill{}, fmt.Errorf("skill %q: front matter requires description", directoryName)
|
||||
}
|
||||
if !skillName.MatchString(metadata.Name) {
|
||||
return Skill{}, fmt.Errorf("skill %q: invalid front matter name %q", directoryName, metadata.Name)
|
||||
}
|
||||
if metadata.Name != directoryName {
|
||||
return Skill{}, fmt.Errorf("skill %q: front matter name %q must match directory name", directoryName, metadata.Name)
|
||||
}
|
||||
skill := Skill{Name: metadata.Name, Description: metadata.Description, Path: path}
|
||||
instructions = body
|
||||
if strings.TrimSpace(instructions) == "" {
|
||||
return Skill{}, fmt.Errorf("skill %q: instructions are empty", directoryName)
|
||||
}
|
||||
skill.Instructions = strings.TrimSpace(instructions)
|
||||
return skill, nil
|
||||
}
|
||||
|
||||
type skillFrontMatterMetadata struct {
|
||||
Name string `yaml:"name"`
|
||||
Description string `yaml:"description"`
|
||||
Metadata map[string]any `yaml:"metadata"`
|
||||
}
|
||||
|
||||
func skillFrontMatter(input string) (skillFrontMatterMetadata, string, error) {
|
||||
input = strings.ReplaceAll(input, "\r\n", "\n")
|
||||
lines := strings.Split(input, "\n")
|
||||
if len(lines) < 3 || lines[0] != "---" {
|
||||
return skillFrontMatterMetadata{}, "", errors.New("invalid front matter")
|
||||
}
|
||||
for i := 1; i < len(lines); i++ {
|
||||
if lines[i] == "---" {
|
||||
var metadata skillFrontMatterMetadata
|
||||
if err := yaml.Unmarshal([]byte(strings.Join(lines[1:i], "\n")), &metadata); err != nil {
|
||||
return skillFrontMatterMetadata{}, "", fmt.Errorf("parse YAML front matter: %w", err)
|
||||
}
|
||||
metadata.Name = strings.TrimSpace(metadata.Name)
|
||||
metadata.Description = strings.TrimSpace(metadata.Description)
|
||||
return metadata, strings.Join(lines[i+1:], "\n"), nil
|
||||
}
|
||||
}
|
||||
return skillFrontMatterMetadata{}, "", errors.New("front matter is not closed")
|
||||
}
|
||||
@@ -0,0 +1,293 @@
|
||||
package agent
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func writeCatalogSkill(t *testing.T, dir, name, content string) {
|
||||
t.Helper()
|
||||
path := filepath.Join(dir, name)
|
||||
if err := os.MkdirAll(path, 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !strings.HasPrefix(content, "---") {
|
||||
content = "---\nname: " + name + "\ndescription: Test skill.\n---\n" + content
|
||||
}
|
||||
if err := os.WriteFile(filepath.Join(path, skillFilename), []byte(content), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDiscoverAndLoadSkills(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
writeCatalogSkill(t, dir, "release-notes", "---\nname: release-notes\ndescription: Draft concise release notes.\nmetadata:\n author: Ollama\n labels:\n - release\n - docs\n---\n# Release notes\n\nUse short bullets.")
|
||||
catalog, err := DiscoverSkills(dir)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
list := catalog.List()
|
||||
if len(list) != 1 || list[0].Name != "release-notes" || list[0].Description != "Draft concise release notes." {
|
||||
t.Fatalf("skills = %#v", list)
|
||||
}
|
||||
skill, err := catalog.Load("release-notes")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !strings.Contains(skill.Content(), `<skill name="release-notes">`) || !strings.Contains(skill.Content(), "Use short bullets.") {
|
||||
t.Fatalf("skill content = %q", skill.Content())
|
||||
}
|
||||
if context := catalog.SystemContext(); !strings.Contains(context, "release-notes: Draft concise release notes.") || !strings.Contains(context, "normal approval rules") {
|
||||
t.Fatalf("system context = %q", context)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDiscoverSkillsSkipsMalformedEntries(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
writeCatalogSkill(t, dir, "valid", "do the useful thing")
|
||||
writeCatalogSkill(t, dir, "mismatched", "---\nname: whatever\ndescription: wrong name\n---\nbody")
|
||||
// Genuinely malformed front matter (a line without a key:value pair) is still rejected.
|
||||
writeCatalogSkill(t, dir, "broken", "---\nname: broken\ndescription\n---\nnope")
|
||||
writeCatalogSkill(t, dir, "missing-name", "---\ndescription: missing name\n---\nbody")
|
||||
writeCatalogSkill(t, dir, "missing-description", "---\nname: missing-description\n---\nbody")
|
||||
writeCatalogSkill(t, dir, "bad-name", "---\nname: bad_name\ndescription: invalid name\n---\nbody")
|
||||
writeCatalogSkill(t, dir, "under_score", "---\nname: under_score\ndescription: invalid directory\n---\nbody")
|
||||
if err := os.MkdirAll(filepath.Join(dir, "no-front-matter"), 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := os.WriteFile(filepath.Join(dir, "no-front-matter", skillFilename), []byte("body"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
catalog, err := DiscoverSkills(dir)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if got, want := len(catalog.List()), 1; got != want {
|
||||
t.Fatalf("valid skills = %d, want %d", got, want)
|
||||
}
|
||||
if got, want := len(catalog.Diagnostics()), 7; got != want {
|
||||
t.Fatalf("diagnostics = %d, want %d: %#v", got, want, catalog.Diagnostics())
|
||||
}
|
||||
if _, err := catalog.Load("broken"); err == nil || !strings.Contains(err.Error(), "not found") {
|
||||
t.Fatalf("load broken error = %v", err)
|
||||
}
|
||||
if _, err := catalog.Load("../valid"); err == nil || !strings.Contains(err.Error(), "invalid skill name") {
|
||||
t.Fatalf("unsafe name error = %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDiscoverSkillsFollowsSymlinks(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
target := t.TempDir()
|
||||
writeCatalogSkill(t, target, "shared", "---\nname: shared\ndescription: From a linked repo.\n---\nshared instructions")
|
||||
if err := os.Symlink(filepath.Join(target, "shared"), filepath.Join(dir, "shared")); err != nil {
|
||||
t.Skipf("symlink not supported: %v", err)
|
||||
}
|
||||
catalog, err := DiscoverSkills(dir)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
list := catalog.List()
|
||||
if len(list) != 1 || list[0].Name != "shared" || list[0].Description != "From a linked repo." {
|
||||
t.Fatalf("symlinked skills = %#v", list)
|
||||
}
|
||||
if !strings.Contains(list[0].Content(), "shared instructions") {
|
||||
t.Fatalf("symlinked skill content = %q", list[0].Content())
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadDefaultSkillsContinuesAfterBadRoot(t *testing.T) {
|
||||
home := t.TempDir()
|
||||
t.Setenv("HOME", home)
|
||||
t.Setenv("USERPROFILE", home)
|
||||
project := t.TempDir()
|
||||
writeCatalogSkill(t, filepath.Join(project, ".ollama", "skills"), "release-notes", "project instructions")
|
||||
|
||||
badRoot := filepath.Join(t.TempDir(), "not-a-directory")
|
||||
if err := os.WriteFile(badRoot, []byte("not a directory"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Setenv(SkillsDirEnv, badRoot)
|
||||
|
||||
catalog, err := LoadDefaultSkills(project)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := catalog.Load("release-notes"); err != nil {
|
||||
t.Fatalf("valid skill was hidden by bad root: %v", err)
|
||||
}
|
||||
if _, err := catalog.Load(bundledSkillCreatorName); err != nil {
|
||||
t.Fatalf("bundled skill was hidden by bad root: %v", err)
|
||||
}
|
||||
var foundDiagnostic bool
|
||||
for _, diagnostic := range catalog.Diagnostics() {
|
||||
if strings.Contains(diagnostic.Error(), badRoot) {
|
||||
foundDiagnostic = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if !foundDiagnostic {
|
||||
t.Fatalf("diagnostics = %#v, want bad root %q", catalog.Diagnostics(), badRoot)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadDefaultSkillsInstallsBundledSkillCreator(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
t.Setenv(SkillsDirEnv, dir)
|
||||
|
||||
catalog, err := LoadDefaultSkills("")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
skill, err := catalog.Load(bundledSkillCreatorName)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
path := filepath.Join(dir, bundledSkillCreatorName, skillFilename)
|
||||
contents, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if string(contents) != bundledSkillCreatorContent {
|
||||
t.Fatalf("installed skill = %q, want bundled contents", contents)
|
||||
}
|
||||
if skill.Path != path {
|
||||
t.Fatalf("skill path = %q, want %q", skill.Path, path)
|
||||
}
|
||||
if !strings.Contains(skill.Content(), "Skill directory: "+filepath.Dir(path)) {
|
||||
t.Fatalf("skill content does not identify its directory: %q", skill.Content())
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadDefaultSkillsUpdatesExistingSkillCreator(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
t.Setenv(SkillsDirEnv, dir)
|
||||
writeCatalogSkill(t, dir, bundledSkillCreatorName, "custom instructions")
|
||||
|
||||
if _, err := LoadDefaultSkills(""); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
contents, err := os.ReadFile(filepath.Join(dir, bundledSkillCreatorName, skillFilename))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if string(contents) != bundledSkillCreatorContent {
|
||||
t.Fatalf("installed skill = %q, want bundled contents", contents)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSkillsDirUsesOverrideAndXDG(t *testing.T) {
|
||||
base := t.TempDir()
|
||||
|
||||
override := filepath.Join(base, "skills-override")
|
||||
t.Setenv(SkillsDirEnv, override)
|
||||
got, err := SkillsDir()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
want, err := filepath.Abs(override)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if got != want {
|
||||
t.Fatalf("SkillsDir override = %q, want %q", got, want)
|
||||
}
|
||||
|
||||
t.Setenv(SkillsDirEnv, "")
|
||||
xdg := filepath.Join(base, "xdg")
|
||||
t.Setenv("XDG_CONFIG_HOME", xdg)
|
||||
if got, err := SkillsDir(); err != nil || got != filepath.Join(xdg, "ollama", "skills") {
|
||||
t.Fatalf("SkillsDir xdg = %q, want %q, %v", got, filepath.Join(xdg, "ollama", "skills"), err)
|
||||
}
|
||||
|
||||
t.Setenv("XDG_CONFIG_HOME", "")
|
||||
home := filepath.Join(base, "home")
|
||||
t.Setenv("HOME", home)
|
||||
t.Setenv("USERPROFILE", home)
|
||||
if got, err := SkillsDir(); err != nil || got != filepath.Join(home, ".ollama", "skills") {
|
||||
t.Fatalf("SkillsDir default = %q, want %q, %v", got, filepath.Join(home, ".ollama", "skills"), err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadDefaultSkillsPrecedenceAndCollisions(t *testing.T) {
|
||||
home := t.TempDir()
|
||||
t.Setenv("HOME", home)
|
||||
t.Setenv("USERPROFILE", home) // Windows: os.UserHomeDir uses %USERPROFILE%
|
||||
|
||||
userOllama := t.TempDir()
|
||||
t.Setenv(SkillsDirEnv, userOllama)
|
||||
|
||||
userAgents := filepath.Join(home, ".agents", "skills")
|
||||
project := t.TempDir()
|
||||
projectAgents := filepath.Join(project, ".agents", "skills")
|
||||
projectOllama := filepath.Join(project, ".ollama", "skills")
|
||||
|
||||
// release-notes exists in all four roots; project ollama must win.
|
||||
writeCatalogSkill(t, userAgents, "release-notes", "from user agents")
|
||||
writeCatalogSkill(t, userOllama, "release-notes", "from user ollama")
|
||||
writeCatalogSkill(t, projectOllama, "release-notes", "from project ollama")
|
||||
// code-review exists in both project roots; project ollama beats project agents.
|
||||
writeCatalogSkill(t, projectAgents, "code-review", "from project agents")
|
||||
writeCatalogSkill(t, projectOllama, "code-review", "from project ollama")
|
||||
// unique appears only in user ollama (via env override).
|
||||
writeCatalogSkill(t, userOllama, "unique", "only here")
|
||||
|
||||
catalog, err := LoadDefaultSkills(project)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
rn, err := catalog.Load("release-notes")
|
||||
if err != nil || !strings.Contains(rn.Instructions, "from project ollama") || !strings.Contains(rn.Path, ".ollama") {
|
||||
t.Fatalf("release-notes = %#v, want project ollama to win", rn)
|
||||
}
|
||||
cr, err := catalog.Load("code-review")
|
||||
if err != nil || !strings.Contains(cr.Instructions, "from project ollama") {
|
||||
t.Fatalf("code-review = %#v, want project ollama to win over project agents", cr)
|
||||
}
|
||||
if _, err := catalog.Load("unique"); err != nil {
|
||||
t.Fatalf("unique should load from user ollama: %v", err)
|
||||
}
|
||||
// Collisions are resolved silently by precedence — no diagnostics.
|
||||
for _, d := range catalog.Diagnostics() {
|
||||
if strings.Contains(d.Error(), "shadows") {
|
||||
t.Fatalf("unexpected shadow diagnostic: %v", d)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestSkillContentListsDirectoryAndResources(t *testing.T) {
|
||||
root := t.TempDir()
|
||||
skillDir := filepath.Join(root, "pdf-processing")
|
||||
if err := os.MkdirAll(filepath.Join(skillDir, "scripts"), 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := os.MkdirAll(filepath.Join(skillDir, "references"), 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := os.WriteFile(filepath.Join(skillDir, "SKILL.md"), []byte("---\nname: pdf-processing\ndescription: Handle PDFs.\n---\nHandle PDFs."), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := os.WriteFile(filepath.Join(skillDir, "scripts", "extract.py"), []byte("#!/usr/bin/env python3"), 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := os.WriteFile(filepath.Join(skillDir, "references", "ref.md"), []byte("ref"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
catalog, err := DiscoverSkills(root)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
skill, err := catalog.Load("pdf-processing")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
content := skill.Content()
|
||||
if !strings.Contains(content, "Skill directory:") || !strings.Contains(content, skillDir) {
|
||||
t.Fatalf("content missing skill directory: %q", content)
|
||||
}
|
||||
if !strings.Contains(content, "<file>scripts/extract.py</file>") || !strings.Contains(content, "<file>references/ref.md</file>") {
|
||||
t.Fatalf("content missing resource listing: %q", content)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,450 @@
|
||||
package tools
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"os"
|
||||
"os/exec"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"strings"
|
||||
"time"
|
||||
"unicode/utf8"
|
||||
|
||||
"github.com/ollama/ollama/agent"
|
||||
"github.com/ollama/ollama/api"
|
||||
)
|
||||
|
||||
const (
|
||||
bashTimeout = 3 * time.Minute
|
||||
bashWaitDelay = 1 * time.Second
|
||||
maxBashOutputBytes = 60_000
|
||||
)
|
||||
|
||||
type Bash struct{}
|
||||
|
||||
func (b *Bash) Name() string {
|
||||
return shellToolName()
|
||||
}
|
||||
|
||||
func (b *Bash) Description() string {
|
||||
return shellToolDescription()
|
||||
}
|
||||
|
||||
func (b *Bash) Schema() api.ToolFunction {
|
||||
props := api.NewToolPropertiesMap()
|
||||
props.Set("command", api.ToolProperty{
|
||||
Type: api.PropertyType{"string"},
|
||||
Description: shellCommandDescription(),
|
||||
})
|
||||
return api.ToolFunction{
|
||||
Name: b.Name(),
|
||||
Description: b.Description(),
|
||||
Parameters: api.ToolFunctionParameters{
|
||||
Type: "object",
|
||||
Properties: props,
|
||||
Required: []string{"command"},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func (b *Bash) RequiresApproval(map[string]any) bool {
|
||||
return true
|
||||
}
|
||||
|
||||
// ApprovalScope scopes shell approval to the exact, trimmed command string
|
||||
// using a NUL separator: "<tool>\x00<command>". "Always allow this command"
|
||||
// matches ONLY that precise string — any whitespace, quoting, or casing
|
||||
// variant re-prompts. The NUL separator is safe because a shell command
|
||||
// string cannot contain a literal NUL.
|
||||
func (b *Bash) ApprovalScope(args map[string]any) string {
|
||||
name := b.Name()
|
||||
if command, ok := args["command"].(string); ok {
|
||||
command = strings.TrimSpace(command)
|
||||
if command != "" {
|
||||
return name + "\x00" + command
|
||||
}
|
||||
}
|
||||
return name
|
||||
}
|
||||
|
||||
func (b *Bash) Execute(ctx context.Context, toolCtx agent.ToolContext, args map[string]any) (agent.ToolResult, error) {
|
||||
// TODO: use shared agent.RequiredStringArg for the "command" parameter (see agent package cleanup plan).
|
||||
command, ok := args["command"].(string)
|
||||
if !ok || strings.TrimSpace(command) == "" {
|
||||
return agent.ToolResult{}, fmt.Errorf("command parameter is required")
|
||||
}
|
||||
if err := rejectUnsafeShellCommand(command); err != nil {
|
||||
return agent.ToolResult{}, err
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithTimeout(ctx, bashTimeout)
|
||||
defer cancel()
|
||||
|
||||
cwdFile, err := os.CreateTemp("", "ollama-agent-cwd-*")
|
||||
if err != nil {
|
||||
return agent.ToolResult{}, err
|
||||
}
|
||||
cwdPath := cwdFile.Name()
|
||||
_ = cwdFile.Close()
|
||||
defer os.Remove(cwdPath)
|
||||
|
||||
cmd := newBashCommand(ctx, command, cwdPath)
|
||||
cmd.WaitDelay = bashWaitDelay
|
||||
cmd.Cancel = func() error {
|
||||
return killBashCommand(cmd)
|
||||
}
|
||||
if toolCtx.WorkingDir != "" {
|
||||
cmd.Dir = toolCtx.WorkingDir
|
||||
}
|
||||
|
||||
var stdout, stderr boundedOutput
|
||||
stdout.Limit = maxBashOutputBytes
|
||||
stderr.Limit = maxBashOutputBytes
|
||||
cmd.Stdout = &stdout
|
||||
cmd.Stderr = &stderr
|
||||
|
||||
err = runBashCommand(cmd)
|
||||
finalWorkingDir := readFinalWorkingDir(cwdPath)
|
||||
|
||||
var sb strings.Builder
|
||||
if stdout.Len() > 0 {
|
||||
sb.WriteString(stdout.String("stdout"))
|
||||
}
|
||||
if stderr.Len() > 0 {
|
||||
if sb.Len() > 0 {
|
||||
sb.WriteString("\n")
|
||||
}
|
||||
sb.WriteString("stderr:\n")
|
||||
sb.WriteString(stderr.String("stderr"))
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
if ctx.Err() == context.DeadlineExceeded {
|
||||
return agent.ToolResult{Content: bashContentWithError(sb.String(), "Error: command timed out after "+bashTimeout.String()), WorkingDir: finalWorkingDir}, nil
|
||||
}
|
||||
if ctx.Err() == context.Canceled {
|
||||
return agent.ToolResult{Content: bashContentWithError(sb.String(), "Error: command was canceled"), WorkingDir: finalWorkingDir}, nil
|
||||
}
|
||||
if errors.Is(err, exec.ErrWaitDelay) {
|
||||
_ = killBashCommand(cmd)
|
||||
return agent.ToolResult{Content: bashContentWithError(sb.String(), "Error: command output pipes did not close after "+bashWaitDelay.String()), WorkingDir: finalWorkingDir}, nil
|
||||
}
|
||||
if exitErr, ok := err.(*exec.ExitError); ok {
|
||||
return agent.ToolResult{Content: bashContentWithError(sb.String(), fmt.Sprintf("Exit code: %d", exitErr.ExitCode())), WorkingDir: finalWorkingDir}, nil
|
||||
}
|
||||
return agent.ToolResult{Content: sb.String(), WorkingDir: finalWorkingDir}, fmt.Errorf("executing command: %w", err)
|
||||
}
|
||||
|
||||
if sb.Len() == 0 {
|
||||
return agent.ToolResult{Content: "(no output)", WorkingDir: finalWorkingDir}, nil
|
||||
}
|
||||
return agent.ToolResult{Content: sb.String(), WorkingDir: finalWorkingDir}, nil
|
||||
}
|
||||
|
||||
func bashContentWithError(content, msg string) string {
|
||||
if content == "" {
|
||||
return msg
|
||||
}
|
||||
return content + "\n\n" + msg
|
||||
}
|
||||
|
||||
// rejectUnsafeShellCommand applies a best-effort blocklist for obviously
|
||||
// destructive or credential-exfiltrating commands. It is defense-in-depth
|
||||
// ONLY: the interactive approval prompt is the real security control, and
|
||||
// this check must not be relied upon as a sandbox. Sophisticated or novel
|
||||
// dangerous commands (e.g. find / -delete, dd, fork bombs, custom binaries)
|
||||
// are NOT caught here and will simply be routed through approval like any
|
||||
// other command. Keep the approval prompt as the gate.
|
||||
func rejectUnsafeShellCommand(command string) error {
|
||||
switch {
|
||||
case hasUnsafeRecursiveDelete(command):
|
||||
return fmt.Errorf("refusing to run unsafe command: recursive delete target is too broad")
|
||||
case readsCredentialPath(command):
|
||||
return fmt.Errorf("refusing to run unsafe command: credential file reads are not allowed")
|
||||
default:
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
func hasUnsafeRecursiveDelete(command string) bool {
|
||||
// Check each command segment independently. shellSafetyText flattens
|
||||
// separators (; & | newlines) to spaces, which would otherwise let the
|
||||
// rm target scan bleed across command boundaries — e.g.
|
||||
// "rm -rf build && echo ~/.ssh/config" flattened to one token stream
|
||||
// would treat the unrelated ~/.ssh/config (a ~/-prefixed "unsafe
|
||||
// target") as an rm argument. Splitting on separators first restores
|
||||
// command boundaries while still catching multi-target single commands
|
||||
// like "rm -rf build /etc".
|
||||
for _, segment := range shellSegments(command) {
|
||||
fields := shellSafetyFields(segment)
|
||||
for i, field := range fields {
|
||||
if isRMCommand(field) && rmCommandDeletesUnsafeTarget(fields[i+1:]) {
|
||||
return true
|
||||
}
|
||||
if isPowerShellDeleteCommand(field) && powerShellDeleteCommandDeletesUnsafeTarget(fields[i+1:]) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// shellSegments splits a command on shell control operators (;, &, |, &&,
|
||||
// ||) and newlines, returning the individual command segments. It operates on
|
||||
// the lowercased raw command before quote/separator normalization so that
|
||||
// command boundaries are preserved for per-segment checks. Subshell parens are
|
||||
// intentionally NOT treated as separators: splitting on them would fragment
|
||||
// command substitutions like "rm -rf $(echo /)" into "rm -rf $" and "echo /",
|
||||
// hiding the destructive "/" target from the per-segment scan. Empty segments
|
||||
// are dropped.
|
||||
func shellSegments(command string) []string {
|
||||
command = strings.ToLower(command)
|
||||
var segments []string
|
||||
for _, segment := range strings.FieldsFunc(command, func(r rune) bool {
|
||||
switch r {
|
||||
case ';', '&', '|', '\n', '\r':
|
||||
return true
|
||||
}
|
||||
return false
|
||||
}) {
|
||||
if segment = strings.TrimSpace(segment); segment != "" {
|
||||
segments = append(segments, segment)
|
||||
}
|
||||
}
|
||||
return segments
|
||||
}
|
||||
|
||||
func rmCommandDeletesUnsafeTarget(fields []string) bool {
|
||||
var flags string
|
||||
for _, field := range fields {
|
||||
if field == "--" {
|
||||
continue
|
||||
}
|
||||
if strings.HasPrefix(field, "-") {
|
||||
flags += field
|
||||
continue
|
||||
}
|
||||
if strings.Contains(flags, "r") && strings.Contains(flags, "f") && isUnsafeDeleteTarget(field) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func powerShellDeleteCommandDeletesUnsafeTarget(fields []string) bool {
|
||||
var recurse, force bool
|
||||
var targets []string
|
||||
for _, field := range fields {
|
||||
switch field {
|
||||
case "-r", "-recurse", "-recursive":
|
||||
recurse = true
|
||||
case "-f", "-force":
|
||||
force = true
|
||||
default:
|
||||
if !strings.HasPrefix(field, "-") {
|
||||
targets = append(targets, field)
|
||||
}
|
||||
}
|
||||
}
|
||||
if !recurse || !force {
|
||||
return false
|
||||
}
|
||||
for _, target := range targets {
|
||||
if isUnsafeDeleteTarget(target) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func readsCredentialPath(command string) bool {
|
||||
fields := shellSafetyFields(command)
|
||||
if !hasCredentialReadVerb(fields) {
|
||||
return false
|
||||
}
|
||||
normalized := shellSafetyText(command)
|
||||
for _, fragment := range []string{
|
||||
"/.ssh/id_rsa",
|
||||
"/.ssh/id_dsa",
|
||||
"/.ssh/id_ecdsa",
|
||||
"/.ssh/id_ed25519",
|
||||
"/.ssh/config",
|
||||
"/.ssh/known_hosts",
|
||||
"/.aws/credentials",
|
||||
"/.aws/config",
|
||||
"/.config/gcloud/application_default_credentials.json",
|
||||
"/.kube/config",
|
||||
"/.netrc",
|
||||
"/.npmrc",
|
||||
"/.docker/config.json",
|
||||
"/.config/gh/hosts.yml",
|
||||
"/.gnupg/",
|
||||
"/etc/shadow",
|
||||
} {
|
||||
if strings.Contains(normalized, fragment) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func hasCredentialReadVerb(fields []string) bool {
|
||||
for _, field := range fields {
|
||||
switch field {
|
||||
case "cat", "less", "more", "head", "tail", "type", "get-content", "gc", "select-string", "grep", "rg", "sed", "awk":
|
||||
return true
|
||||
case "env", "printenv":
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func isRMCommand(field string) bool {
|
||||
return field == "rm" || strings.HasSuffix(field, "/rm")
|
||||
}
|
||||
|
||||
func isPowerShellDeleteCommand(field string) bool {
|
||||
switch field {
|
||||
case "remove-item", "del", "erase", "rd", "rmdir":
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func isUnsafeDeleteTarget(target string) bool {
|
||||
if target == "." || target == "./" || target == "*" {
|
||||
return true
|
||||
}
|
||||
if target == "/*" {
|
||||
return true
|
||||
}
|
||||
target = strings.TrimSuffix(target, "/*")
|
||||
for _, prefix := range []string{"~/", "$home/", "${home}/", "$env:home/", "$env:userprofile/", "%userprofile%/"} {
|
||||
if strings.HasPrefix(target, prefix) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
for _, prefix := range []string{"/etc/", "/bin/", "/sbin/", "/usr/", "/var/", "/lib/", "/library/", "/system/", "/applications/", "c:/windows/", "c:/program files/"} {
|
||||
if strings.HasPrefix(target, prefix) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
for _, exact := range []string{"/", "~", "$home", "${home}", "$env:home", "$env:userprofile", "%userprofile%", "c:", "c:/", "/etc", "/bin", "/sbin", "/usr", "/var", "/lib", "/library", "/system", "/applications", "c:/windows", "c:/program files"} {
|
||||
if target == exact {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func shellSafetyFields(command string) []string {
|
||||
return strings.Fields(shellSafetyText(command))
|
||||
}
|
||||
|
||||
func shellSafetyText(command string) string {
|
||||
command = strings.ToLower(command)
|
||||
return strings.NewReplacer(
|
||||
"\\", "/",
|
||||
"\n", " ",
|
||||
"\t", " ",
|
||||
";", " ",
|
||||
"&", " ",
|
||||
"|", " ",
|
||||
"(", " ",
|
||||
")", " ",
|
||||
"\"", "",
|
||||
"'", "",
|
||||
"`", "",
|
||||
).Replace(command)
|
||||
}
|
||||
|
||||
func readFinalWorkingDir(path string) string {
|
||||
content, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
workingDir := strings.TrimPrefix(string(content), "\ufeff")
|
||||
workingDir = strings.TrimSpace(workingDir)
|
||||
if workingDir == "" {
|
||||
return ""
|
||||
}
|
||||
workingDir = normalizeBashWorkingDir(workingDir)
|
||||
info, err := os.Stat(workingDir)
|
||||
if err != nil || !info.IsDir() {
|
||||
return ""
|
||||
}
|
||||
return workingDir
|
||||
}
|
||||
|
||||
func normalizeBashWorkingDir(workingDir string) string {
|
||||
if runtime.GOOS == "windows" && len(workingDir) >= 3 && workingDir[0] == '/' && workingDir[2] == '/' && isASCIIAlpha(workingDir[1]) {
|
||||
workingDir = strings.ToUpper(string(workingDir[1])) + ":" + workingDir[2:]
|
||||
}
|
||||
workingDir = filepath.Clean(filepath.FromSlash(workingDir))
|
||||
if runtime.GOOS == "windows" && len(workingDir) >= 2 && workingDir[1] == ':' && isASCIIAlpha(workingDir[0]) {
|
||||
workingDir = strings.ToUpper(string(workingDir[0])) + workingDir[1:]
|
||||
}
|
||||
return workingDir
|
||||
}
|
||||
|
||||
func isASCIIAlpha(b byte) bool {
|
||||
return (b >= 'a' && b <= 'z') || (b >= 'A' && b <= 'Z')
|
||||
}
|
||||
|
||||
type boundedOutput struct {
|
||||
Limit int
|
||||
buf []byte
|
||||
omitted int
|
||||
}
|
||||
|
||||
func (b *boundedOutput) Write(p []byte) (int, error) {
|
||||
if b.Limit <= 0 {
|
||||
b.omitted += len(p)
|
||||
return len(p), nil
|
||||
}
|
||||
remaining := b.Limit - len(b.buf)
|
||||
if remaining <= 0 {
|
||||
b.omitted += len(p)
|
||||
return len(p), nil
|
||||
}
|
||||
if len(p) <= remaining {
|
||||
b.buf = append(b.buf, p...)
|
||||
return len(p), nil
|
||||
}
|
||||
writeLen := utf8SafePrefixLen(p[:remaining])
|
||||
b.buf = append(b.buf, p[:writeLen]...)
|
||||
b.omitted += len(p) - writeLen
|
||||
return len(p), nil
|
||||
}
|
||||
|
||||
func (b *boundedOutput) Len() int {
|
||||
return len(b.buf) + b.omitted
|
||||
}
|
||||
|
||||
func (b *boundedOutput) String(label string) string {
|
||||
safeLen := utf8SafePrefixLen(b.buf)
|
||||
content := string(b.buf[:safeLen])
|
||||
omitted := b.omitted + len(b.buf) - safeLen
|
||||
if omitted == 0 {
|
||||
return content
|
||||
}
|
||||
return content + agent.TruncMarker(label, safeLen, 0, omitted, false, "")
|
||||
}
|
||||
|
||||
func utf8SafePrefixLen(p []byte) int {
|
||||
if len(p) == 0 {
|
||||
return 0
|
||||
}
|
||||
for i := 0; i < len(p); {
|
||||
r, size := utf8.DecodeRune(p[i:])
|
||||
if r == utf8.RuneError && size == 1 {
|
||||
return i
|
||||
}
|
||||
i += size
|
||||
}
|
||||
return len(p)
|
||||
}
|
||||
@@ -0,0 +1,258 @@
|
||||
package tools
|
||||
|
||||
import (
|
||||
"context"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"strings"
|
||||
"testing"
|
||||
"unicode/utf8"
|
||||
|
||||
"github.com/ollama/ollama/agent"
|
||||
)
|
||||
|
||||
func TestBashReportsFinalWorkingDir(t *testing.T) {
|
||||
root := t.TempDir()
|
||||
subdir := filepath.Join(root, "sub")
|
||||
if err := os.Mkdir(subdir, 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
result, err := (&Bash{}).Execute(context.Background(), agent.ToolContext{WorkingDir: root}, map[string]any{
|
||||
"command": shellTestCommand("cd sub && pwd", "Set-Location sub; Get-Location"),
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
wantDir, err := filepath.EvalSymlinks(subdir)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if result.WorkingDir != wantDir {
|
||||
t.Fatalf("working dir = %q, want %q", result.WorkingDir, wantDir)
|
||||
}
|
||||
if !strings.Contains(result.Content, "sub") {
|
||||
t.Fatalf("content = %q, want pwd output", result.Content)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBashBoundsOutputWhileRunning(t *testing.T) {
|
||||
result, err := (&Bash{}).Execute(context.Background(), agent.ToolContext{WorkingDir: t.TempDir()}, map[string]any{
|
||||
"command": shellTestCommand("yes x | head -c 70000", "[Console]::Out.Write(('x' * 70000))"),
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !strings.Contains(result.Content, "[stdout truncated: showing first ~") || !strings.Contains(result.Content, "omitted ~") || !strings.Contains(result.Content, " tokens.]") {
|
||||
t.Fatalf("content = %q, want stdout truncation marker", result.Content)
|
||||
}
|
||||
if count, want := strings.Count(result.Content, "x"), shellTestCapturedXCount(); count != want {
|
||||
t.Fatalf("captured x count = %d, want %d", count, want)
|
||||
}
|
||||
if len(result.Content) > maxBashOutputBytes+200 {
|
||||
t.Fatalf("content length = %d, want bounded output", len(result.Content))
|
||||
}
|
||||
}
|
||||
|
||||
func TestBoundedOutputTruncatesAtUTF8Boundary(t *testing.T) {
|
||||
var out boundedOutput
|
||||
out.Limit = len([]byte("abc")) + 1
|
||||
|
||||
if _, err := out.Write([]byte("abcédef")); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
content := out.String("stdout")
|
||||
if !utf8.ValidString(content) {
|
||||
t.Fatalf("content is not valid UTF-8: %q", content)
|
||||
}
|
||||
if strings.ContainsRune(content, utf8.RuneError) {
|
||||
t.Fatalf("content contains replacement rune: %q", content)
|
||||
}
|
||||
if !strings.HasPrefix(content, "abc\n\n[stdout truncated:") {
|
||||
t.Fatalf("content = %q, want complete ASCII prefix and truncation marker", content)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBoundedOutputKeepsCompleteUTF8AtBoundary(t *testing.T) {
|
||||
var out boundedOutput
|
||||
out.Limit = len([]byte("abcé"))
|
||||
|
||||
if _, err := out.Write([]byte("abcédef")); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if content := out.String("stdout"); !strings.HasPrefix(content, "abcé\n\n[stdout truncated:") {
|
||||
t.Fatalf("content = %q, want complete UTF-8 prefix", content)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBoundedOutputTrimsTrailingPartialUTF8(t *testing.T) {
|
||||
var out boundedOutput
|
||||
out.Limit = 4
|
||||
|
||||
if _, err := out.Write([]byte{'a', 'b', 'c', 0xc3}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := out.Write([]byte{0xa9}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if content := out.String("stdout"); !utf8.ValidString(content) || !strings.HasPrefix(content, "abc\n\n[stdout truncated:") {
|
||||
t.Fatalf("content = %q, want valid UTF-8 with partial suffix trimmed", content)
|
||||
}
|
||||
}
|
||||
|
||||
func TestUTF8SafePrefixRejectsMalformedLeadByte(t *testing.T) {
|
||||
input := []byte{'a', 0xc0, 0x80, 'b'}
|
||||
if got := utf8SafePrefixLen(input); got != 1 {
|
||||
t.Fatalf("safe prefix length = %d, want 1", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBoundedOutputDropsMalformedUTF8(t *testing.T) {
|
||||
var out boundedOutput
|
||||
out.Limit = 4
|
||||
|
||||
if _, err := out.Write([]byte{'a', 0xc0, 0x80, 'b'}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
content := out.String("stdout")
|
||||
if !utf8.ValidString(content) {
|
||||
t.Fatalf("content is not valid UTF-8: %q", content)
|
||||
}
|
||||
if strings.ContainsRune(content, utf8.RuneError) {
|
||||
t.Fatalf("content contains replacement rune: %q", content)
|
||||
}
|
||||
if !strings.HasPrefix(content, "a\n\n[stdout truncated:") {
|
||||
t.Fatalf("content = %q, want valid prefix and truncation marker", content)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBashReportsCanceledCommand(t *testing.T) {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
|
||||
result, err := (&Bash{}).Execute(ctx, agent.ToolContext{WorkingDir: t.TempDir()}, map[string]any{
|
||||
"command": shellTestCommand("sleep 10", "Start-Sleep -Seconds 10"),
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !strings.Contains(result.Content, "Error: command was canceled") {
|
||||
t.Fatalf("content = %q, want canceled message", result.Content)
|
||||
}
|
||||
if strings.Contains(result.Content, "Exit code: -1") {
|
||||
t.Fatalf("content = %q, should not mask cancellation as exit code", result.Content)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRejectUnsafeShellCommand(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
command string
|
||||
wantErr bool
|
||||
}{
|
||||
{name: "rm root", command: "rm -rf /", wantErr: true},
|
||||
{name: "sudo rm root", command: "sudo rm -rf -- /", wantErr: true},
|
||||
{name: "rm home", command: "rm -fr $HOME", wantErr: true},
|
||||
{name: "rm root wildcard", command: "rm -rf /*", wantErr: true},
|
||||
{name: "rm system subdir", command: "rm -rf /etc/ssh", wantErr: true},
|
||||
{name: "rm cwd", command: "rm -rf .", wantErr: true},
|
||||
{name: "powershell remove root", command: `Remove-Item -Recurse -Force C:\`, wantErr: true},
|
||||
{name: "powershell remove system subdir", command: `Remove-Item -Recurse -Force C:\Windows\Temp`, wantErr: true},
|
||||
{name: "ssh private key", command: "cat ~/.ssh/id_rsa", wantErr: true},
|
||||
{name: "aws credentials", command: "Get-Content $HOME/.aws/credentials", wantErr: true},
|
||||
{name: "shadow", command: "head /etc/shadow", wantErr: true},
|
||||
{name: "netrc", command: "cat ~/.netrc", wantErr: true},
|
||||
{name: "docker config", command: "cat ~/.docker/config.json", wantErr: true},
|
||||
{name: "gnupg dir", command: "cat ~/.gnupg/private-keys-v1.d/key", wantErr: true},
|
||||
{name: "gh hosts", command: "cat ~/.config/gh/hosts.yml", wantErr: true},
|
||||
{name: "ssh config", command: "cat ~/.ssh/config", wantErr: true},
|
||||
{name: "printenv dump", command: "printenv", wantErr: false},
|
||||
{name: "delete build dir", command: "rm -rf build", wantErr: false},
|
||||
{name: "read project file", command: "cat README.md", wantErr: false},
|
||||
{name: "mention key text", command: "rg id_rsa docs", wantErr: false},
|
||||
{name: "env example", command: "cat .env.example", wantErr: false},
|
||||
{name: "rm build then unrelated tilde path", command: "rm -rf build && echo ~/.ssh/config", wantErr: false},
|
||||
{name: "rm build then unrelated slash path", command: "rm -rf build; cat /etc/passwd", wantErr: false},
|
||||
{name: "rm build then unrelated star glob", command: "rm -rf build && ls *.go", wantErr: false},
|
||||
{name: "rm multiple targets one unsafe", command: "rm -rf build /etc", wantErr: true},
|
||||
{name: "rm unsafe then safe piped", command: "rm -rf / | tee log", wantErr: true},
|
||||
{name: "rm unsafe via command substitution", command: "rm -rf $(echo /)", wantErr: true},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
err := rejectUnsafeShellCommand(tt.command)
|
||||
if tt.wantErr && err == nil {
|
||||
t.Fatal("expected unsafe command to be rejected")
|
||||
}
|
||||
if !tt.wantErr && err != nil {
|
||||
t.Fatalf("command rejected: %v", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestBashRejectsUnsafeCommandBeforeExecution(t *testing.T) {
|
||||
_, err := (&Bash{}).Execute(context.Background(), agent.ToolContext{WorkingDir: t.TempDir()}, map[string]any{
|
||||
"command": "rm -rf /",
|
||||
})
|
||||
if err == nil || !strings.Contains(err.Error(), "refusing to run unsafe command") {
|
||||
t.Fatalf("err = %v, want unsafe command rejection", err)
|
||||
}
|
||||
}
|
||||
|
||||
func shellTestCommand(unix, windows string) string {
|
||||
if runtime.GOOS == "windows" {
|
||||
return windows
|
||||
}
|
||||
return unix
|
||||
}
|
||||
|
||||
func shellTestCapturedXCount() int {
|
||||
if runtime.GOOS == "windows" {
|
||||
return maxBashOutputBytes
|
||||
}
|
||||
return maxBashOutputBytes / 2
|
||||
}
|
||||
|
||||
func TestReadFinalWorkingDirRejectsInvalidPaths(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
cwdFile := filepath.Join(dir, "cwd")
|
||||
notDir := filepath.Join(dir, "file.txt")
|
||||
if err := os.WriteFile(notDir, []byte("not a dir"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
if err := os.WriteFile(cwdFile, []byte(notDir+"\n"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if got := readFinalWorkingDir(cwdFile); got != "" {
|
||||
t.Fatalf("regular file cwd = %q, want empty", got)
|
||||
}
|
||||
|
||||
if err := os.WriteFile(cwdFile, []byte(filepath.Join(dir, "missing")+"\n"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if got := readFinalWorkingDir(cwdFile); got != "" {
|
||||
t.Fatalf("missing cwd = %q, want empty", got)
|
||||
}
|
||||
|
||||
if err := os.WriteFile(cwdFile, []byte(dir+"\n"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if got := readFinalWorkingDir(cwdFile); got != dir {
|
||||
t.Fatalf("directory cwd = %q, want %q", got, dir)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNormalizeBashWorkingDirWindowsDriveLetter(t *testing.T) {
|
||||
if runtime.GOOS != "windows" {
|
||||
t.Skip("windows path normalization")
|
||||
}
|
||||
got := normalizeBashWorkingDir("/c/Users/jdoe/project")
|
||||
want := filepath.Clean(`C:\Users\jdoe\project`)
|
||||
if got != want {
|
||||
t.Fatalf("working dir = %q, want %q", got, want)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,49 @@
|
||||
//go:build !windows
|
||||
|
||||
package tools
|
||||
|
||||
import (
|
||||
"context"
|
||||
"os/exec"
|
||||
"strings"
|
||||
"syscall"
|
||||
)
|
||||
|
||||
func shellToolName() string {
|
||||
return "bash"
|
||||
}
|
||||
|
||||
func shellToolDescription() string {
|
||||
return "Execute a bash command on the system. Use this to inspect files, run tests, and perform development tasks."
|
||||
}
|
||||
|
||||
func shellCommandDescription() string {
|
||||
return "The bash command to execute."
|
||||
}
|
||||
|
||||
func newBashCommand(ctx context.Context, command, cwdPath string) *exec.Cmd {
|
||||
script := command + "\n__ollama_status=$?\npwd -P > " + shellQuote(cwdPath) + "\nexit $__ollama_status"
|
||||
cmd := exec.CommandContext(ctx, "bash", "-c", script)
|
||||
configureBashCommand(cmd)
|
||||
return cmd
|
||||
}
|
||||
|
||||
func shellQuote(value string) string {
|
||||
return "'" + strings.ReplaceAll(value, "'", "'\\''") + "'"
|
||||
}
|
||||
|
||||
func configureBashCommand(cmd *exec.Cmd) {
|
||||
cmd.SysProcAttr = &syscall.SysProcAttr{Setpgid: true}
|
||||
}
|
||||
|
||||
func runBashCommand(cmd *exec.Cmd) error {
|
||||
return cmd.Run()
|
||||
}
|
||||
|
||||
func killBashCommand(cmd *exec.Cmd) error {
|
||||
if cmd == nil || cmd.Process == nil {
|
||||
return nil
|
||||
}
|
||||
_ = syscall.Kill(-cmd.Process.Pid, syscall.SIGKILL)
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,40 @@
|
||||
//go:build !windows
|
||||
|
||||
package tools
|
||||
|
||||
import (
|
||||
"context"
|
||||
"os/exec"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/ollama/ollama/agent"
|
||||
)
|
||||
|
||||
func TestConfigureBashCommandSetsProcessGroup(t *testing.T) {
|
||||
cmd := exec.Command("bash", "-c", "true")
|
||||
configureBashCommand(cmd)
|
||||
if cmd.SysProcAttr == nil || !cmd.SysProcAttr.Setpgid {
|
||||
t.Fatalf("configureBashCommand should start bash in a new process group")
|
||||
}
|
||||
}
|
||||
|
||||
func TestBashWaitDelayBoundsBackgroundOutputPipe(t *testing.T) {
|
||||
start := time.Now()
|
||||
result, err := (&Bash{}).Execute(context.Background(), agent.ToolContext{WorkingDir: t.TempDir()}, map[string]any{
|
||||
"command": "sleep 5 & echo done",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if elapsed := time.Since(start); elapsed > bashWaitDelay+2*time.Second {
|
||||
t.Fatalf("command elapsed = %s, want bounded near %s", elapsed, bashWaitDelay)
|
||||
}
|
||||
if !strings.Contains(result.Content, "done") {
|
||||
t.Fatalf("content = %q, want command output", result.Content)
|
||||
}
|
||||
if !strings.Contains(result.Content, "output pipes did not close") {
|
||||
t.Fatalf("content = %q, want wait delay message", result.Content)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,134 @@
|
||||
//go:build windows
|
||||
|
||||
package tools
|
||||
|
||||
import (
|
||||
"context"
|
||||
"os/exec"
|
||||
"strings"
|
||||
"sync"
|
||||
"unsafe"
|
||||
|
||||
"golang.org/x/sys/windows"
|
||||
)
|
||||
|
||||
var bashJobHandles sync.Map
|
||||
|
||||
func shellToolName() string {
|
||||
return "powershell"
|
||||
}
|
||||
|
||||
func shellToolDescription() string {
|
||||
return "Execute a PowerShell command on the system. Use this to inspect files, run tests, and perform development tasks."
|
||||
}
|
||||
|
||||
func shellCommandDescription() string {
|
||||
return "The PowerShell command to execute."
|
||||
}
|
||||
|
||||
func newBashCommand(ctx context.Context, command, cwdPath string) *exec.Cmd {
|
||||
return exec.CommandContext(
|
||||
ctx,
|
||||
"powershell.exe",
|
||||
"-NoLogo",
|
||||
"-NoProfile",
|
||||
"-NonInteractive",
|
||||
"-ExecutionPolicy",
|
||||
"Bypass",
|
||||
"-Command",
|
||||
powerShellCommandScript(command, cwdPath),
|
||||
)
|
||||
}
|
||||
|
||||
func powerShellCommandScript(command, cwdPath string) string {
|
||||
cwdPath = powerShellSingleQuote(cwdPath)
|
||||
return strings.Join([]string{
|
||||
"$__ollama_status = 0",
|
||||
". {",
|
||||
"try {",
|
||||
command,
|
||||
" $__ollama_success = $?",
|
||||
" $__ollama_last_exit = $global:LASTEXITCODE",
|
||||
" if ($__ollama_success) {",
|
||||
" $__ollama_status = 0",
|
||||
" } elseif ($__ollama_last_exit -is [int] -and $__ollama_last_exit -ne 0) {",
|
||||
" $__ollama_status = $__ollama_last_exit",
|
||||
" } else {",
|
||||
" $__ollama_status = 1",
|
||||
" }",
|
||||
"} catch {",
|
||||
" Write-Error $_",
|
||||
" $__ollama_status = 1",
|
||||
"} finally {",
|
||||
" try { [System.IO.File]::WriteAllText(" + cwdPath + ", (Get-Location).ProviderPath, [System.Text.Encoding]::UTF8) } catch {}",
|
||||
"}",
|
||||
"} | Out-String -Stream -Width 4096",
|
||||
"exit $__ollama_status",
|
||||
}, "\n")
|
||||
}
|
||||
|
||||
func powerShellSingleQuote(value string) string {
|
||||
return "'" + strings.ReplaceAll(value, "'", "''") + "'"
|
||||
}
|
||||
|
||||
func runBashCommand(cmd *exec.Cmd) error {
|
||||
if err := cmd.Start(); err != nil {
|
||||
return err
|
||||
}
|
||||
if job, err := createBashJob(cmd.Process.Pid); err == nil {
|
||||
bashJobHandles.Store(cmd.Process.Pid, job)
|
||||
defer releaseBashJob(cmd.Process.Pid)
|
||||
}
|
||||
return cmd.Wait()
|
||||
}
|
||||
|
||||
func killBashCommand(cmd *exec.Cmd) error {
|
||||
if cmd == nil || cmd.Process == nil {
|
||||
return nil
|
||||
}
|
||||
releaseBashJob(cmd.Process.Pid)
|
||||
_ = cmd.Process.Kill()
|
||||
return nil
|
||||
}
|
||||
|
||||
func createBashJob(pid int) (windows.Handle, error) {
|
||||
job, err := windows.CreateJobObject(nil, nil)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
|
||||
info := windows.JOBOBJECT_EXTENDED_LIMIT_INFORMATION{}
|
||||
info.BasicLimitInformation.LimitFlags = windows.JOB_OBJECT_LIMIT_KILL_ON_JOB_CLOSE
|
||||
if _, err := windows.SetInformationJobObject(
|
||||
job,
|
||||
windows.JobObjectExtendedLimitInformation,
|
||||
uintptr(unsafe.Pointer(&info)),
|
||||
uint32(unsafe.Sizeof(info)),
|
||||
); err != nil {
|
||||
_ = windows.CloseHandle(job)
|
||||
return 0, err
|
||||
}
|
||||
|
||||
process, err := windows.OpenProcess(windows.PROCESS_SET_QUOTA|windows.PROCESS_TERMINATE, false, uint32(pid))
|
||||
if err != nil {
|
||||
_ = windows.CloseHandle(job)
|
||||
return 0, err
|
||||
}
|
||||
defer windows.CloseHandle(process)
|
||||
|
||||
if err := windows.AssignProcessToJobObject(job, process); err != nil {
|
||||
_ = windows.CloseHandle(job)
|
||||
return 0, err
|
||||
}
|
||||
return job, nil
|
||||
}
|
||||
|
||||
func releaseBashJob(pid int) {
|
||||
value, ok := bashJobHandles.LoadAndDelete(pid)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
if job, ok := value.(windows.Handle); ok {
|
||||
_ = windows.CloseHandle(job)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,15 @@
|
||||
//go:build windows
|
||||
|
||||
package tools
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestPowerShellCommandScriptUsesWideOutString(t *testing.T) {
|
||||
script := powerShellCommandScript("Get-ChildItem", `C:\cwd.txt`)
|
||||
if !strings.Contains(script, "Out-String -Stream -Width 4096") {
|
||||
t.Fatalf("script = %q, want explicit Out-String width", script)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,558 @@
|
||||
package tools
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"context"
|
||||
"fmt"
|
||||
"io"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"github.com/ollama/ollama/agent"
|
||||
"github.com/ollama/ollama/api"
|
||||
)
|
||||
|
||||
const (
|
||||
maxReadBytes = 200000
|
||||
)
|
||||
|
||||
type Read struct{}
|
||||
|
||||
func (r *Read) Name() string {
|
||||
return "read"
|
||||
}
|
||||
|
||||
func (r *Read) Description() string {
|
||||
return "Read a text file from the current working directory."
|
||||
}
|
||||
|
||||
func (r *Read) Schema() api.ToolFunction {
|
||||
props := api.NewToolPropertiesMap()
|
||||
props.Set("path", api.ToolProperty{
|
||||
Type: api.PropertyType{"string"},
|
||||
Description: "Path to the file to read, relative to the working directory.",
|
||||
})
|
||||
props.Set("start", api.ToolProperty{
|
||||
Type: api.PropertyType{"integer"},
|
||||
Description: "Optional 1-based line to start reading from.",
|
||||
})
|
||||
props.Set("end", api.ToolProperty{
|
||||
Type: api.PropertyType{"integer"},
|
||||
Description: "Optional 1-based inclusive line to stop reading at.",
|
||||
})
|
||||
return api.ToolFunction{
|
||||
Name: r.Name(),
|
||||
Description: r.Description(),
|
||||
Parameters: api.ToolFunctionParameters{
|
||||
Type: "object",
|
||||
Properties: props,
|
||||
Required: []string{"path"},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func (r *Read) RequiresApproval(map[string]any) bool {
|
||||
return true
|
||||
}
|
||||
|
||||
func (r *Read) Execute(ctx context.Context, toolCtx agent.ToolContext, args map[string]any) (agent.ToolResult, error) {
|
||||
// TODO: use shared agent.RequiredStringArg / agent.OptionalIntArg for args (see agent package cleanup plan).
|
||||
path, ok := args["path"].(string)
|
||||
if !ok || strings.TrimSpace(path) == "" {
|
||||
return agent.ToolResult{}, fmt.Errorf("path parameter is required")
|
||||
}
|
||||
|
||||
file, info, err := openRegularFile(toolCtx.WorkingDir, path, true)
|
||||
if err != nil {
|
||||
return agent.ToolResult{}, err
|
||||
}
|
||||
defer file.Close()
|
||||
|
||||
selection, err := readSelectionFromArgs(args)
|
||||
if err != nil {
|
||||
return agent.ToolResult{}, err
|
||||
}
|
||||
if !selection.enabled && info.Size() > maxReadBytes {
|
||||
return agent.ToolResult{}, fmt.Errorf("%s is too large to read (%d bytes)", path, info.Size())
|
||||
}
|
||||
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return agent.ToolResult{}, ctx.Err()
|
||||
default:
|
||||
}
|
||||
|
||||
var content string
|
||||
if selection.enabled {
|
||||
content, err = readLineSelection(file, selection)
|
||||
} else {
|
||||
var contentBytes []byte
|
||||
contentBytes, err = readAllWithinLimit(file, maxReadBytes)
|
||||
content = string(contentBytes)
|
||||
}
|
||||
if err != nil {
|
||||
return agent.ToolResult{}, err
|
||||
}
|
||||
return agent.ToolResult{Content: content}, nil
|
||||
}
|
||||
|
||||
type Edit struct{}
|
||||
|
||||
func (e *Edit) Name() string {
|
||||
return "edit"
|
||||
}
|
||||
|
||||
func (e *Edit) Description() string {
|
||||
return "Edit a text file in the current working directory by replacing exact text."
|
||||
}
|
||||
|
||||
func (e *Edit) Schema() api.ToolFunction {
|
||||
props := api.NewToolPropertiesMap()
|
||||
props.Set("path", api.ToolProperty{
|
||||
Type: api.PropertyType{"string"},
|
||||
Description: "Path to the file to edit, relative to the working directory.",
|
||||
})
|
||||
props.Set("old_text", api.ToolProperty{
|
||||
Type: api.PropertyType{"string"},
|
||||
Description: "Exact text to replace.",
|
||||
})
|
||||
props.Set("new_text", api.ToolProperty{
|
||||
Type: api.PropertyType{"string"},
|
||||
Description: "Replacement text.",
|
||||
})
|
||||
props.Set("replace_all", api.ToolProperty{
|
||||
Type: api.PropertyType{"boolean"},
|
||||
Description: "Replace every occurrence. Defaults to false and requires old_text to match exactly once.",
|
||||
})
|
||||
return api.ToolFunction{
|
||||
Name: e.Name(),
|
||||
Description: e.Description(),
|
||||
Parameters: api.ToolFunctionParameters{
|
||||
Type: "object",
|
||||
Properties: props,
|
||||
Required: []string{"path", "old_text", "new_text"},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func (e *Edit) RequiresApproval(map[string]any) bool {
|
||||
return true
|
||||
}
|
||||
|
||||
func (e *Edit) Execute(ctx context.Context, toolCtx agent.ToolContext, args map[string]any) (agent.ToolResult, error) {
|
||||
// TODO: use shared agent.RequiredStringArg / agent.OptionalBoolArg for args (see agent package cleanup plan).
|
||||
path, ok := args["path"].(string)
|
||||
if !ok || strings.TrimSpace(path) == "" {
|
||||
return agent.ToolResult{}, fmt.Errorf("path parameter is required")
|
||||
}
|
||||
|
||||
oldText, ok := args["old_text"].(string)
|
||||
if !ok || oldText == "" {
|
||||
return agent.ToolResult{}, fmt.Errorf("old_text parameter is required")
|
||||
}
|
||||
|
||||
newText, ok := args["new_text"].(string)
|
||||
if !ok {
|
||||
return agent.ToolResult{}, fmt.Errorf("new_text parameter is required")
|
||||
}
|
||||
|
||||
replaceAll, _ := args["replace_all"].(bool)
|
||||
|
||||
if err := rejectFinalSymlink(toolCtx.WorkingDir, path); err != nil {
|
||||
return agent.ToolResult{}, err
|
||||
}
|
||||
|
||||
file, info, err := openRegularFile(toolCtx.WorkingDir, path, false)
|
||||
if err != nil {
|
||||
return agent.ToolResult{}, err
|
||||
}
|
||||
if info.Size() > maxReadBytes {
|
||||
file.Close()
|
||||
return agent.ToolResult{}, fmt.Errorf("%s is too large to edit (%d bytes)", path, info.Size())
|
||||
}
|
||||
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
file.Close()
|
||||
return agent.ToolResult{}, ctx.Err()
|
||||
default:
|
||||
}
|
||||
|
||||
contentBytes, err := readAllWithinLimit(file, maxReadBytes)
|
||||
if closeErr := file.Close(); err == nil && closeErr != nil {
|
||||
err = closeErr
|
||||
}
|
||||
if err != nil {
|
||||
return agent.ToolResult{}, err
|
||||
}
|
||||
content := string(contentBytes)
|
||||
matches := strings.Count(content, oldText)
|
||||
if matches == 0 {
|
||||
return agent.ToolResult{}, fmt.Errorf("old_text was not found in %s", path)
|
||||
}
|
||||
if matches > 1 && !replaceAll {
|
||||
return agent.ToolResult{}, fmt.Errorf("old_text matched %d times in %s; set replace_all to true to replace every match", matches, path)
|
||||
}
|
||||
|
||||
var updated string
|
||||
if replaceAll {
|
||||
updated = strings.ReplaceAll(content, oldText, newText)
|
||||
} else {
|
||||
updated = strings.Replace(content, oldText, newText, 1)
|
||||
}
|
||||
if len(updated) > maxReadBytes {
|
||||
return agent.ToolResult{}, fmt.Errorf("edited content is too large (%d bytes)", len(updated))
|
||||
}
|
||||
|
||||
if err := writeFileAtomic(toolCtx.WorkingDir, path, []byte(updated), info.Mode().Perm()); err != nil {
|
||||
return agent.ToolResult{}, err
|
||||
}
|
||||
|
||||
return agent.ToolResult{Content: fmt.Sprintf("Updated %s (%d replacement%s).", path, matches, plural(matches))}, nil
|
||||
}
|
||||
|
||||
func cleanRelativePath(path string) (string, error) {
|
||||
path = strings.TrimSpace(path)
|
||||
if path == "" {
|
||||
return "", fmt.Errorf("path parameter is required")
|
||||
}
|
||||
if filepath.IsAbs(path) {
|
||||
return "", fmt.Errorf("absolute paths are not allowed")
|
||||
}
|
||||
cleaned := filepath.Clean(path)
|
||||
if cleaned == "." || cleaned == ".." || strings.HasPrefix(cleaned, ".."+string(os.PathSeparator)) {
|
||||
return "", fmt.Errorf("path escapes working directory")
|
||||
}
|
||||
return cleaned, nil
|
||||
}
|
||||
|
||||
func openRegularFile(workingDir, path string, allowAbsolute bool) (*os.File, os.FileInfo, error) {
|
||||
path = strings.TrimSpace(path)
|
||||
if path == "" {
|
||||
return nil, nil, fmt.Errorf("path parameter is required")
|
||||
}
|
||||
if allowAbsolute && filepath.IsAbs(path) {
|
||||
cleaned := filepath.Clean(path)
|
||||
info, err := os.Lstat(cleaned)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
if info.Mode()&os.ModeSymlink != 0 {
|
||||
return nil, nil, fmt.Errorf("%s is a symlink; read the target file directly", path)
|
||||
}
|
||||
if err := rejectNonRegularFile(path, info); err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
file, err := os.Open(cleaned)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
info, err = file.Stat()
|
||||
if err != nil {
|
||||
file.Close()
|
||||
return nil, nil, err
|
||||
}
|
||||
if err := rejectNonRegularFile(path, info); err != nil {
|
||||
file.Close()
|
||||
return nil, nil, err
|
||||
}
|
||||
return file, info, nil
|
||||
}
|
||||
|
||||
rel, err := cleanRelativePath(path)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
root, err := openWorkingRoot(workingDir)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
defer root.Close()
|
||||
|
||||
if _, err := regularRootFileInfo(root, rel, path); err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
file, err := root.Open(rel)
|
||||
if err != nil {
|
||||
return nil, nil, rootPathError(err)
|
||||
}
|
||||
info, err := file.Stat()
|
||||
if err != nil {
|
||||
file.Close()
|
||||
return nil, nil, err
|
||||
}
|
||||
if err := rejectNonRegularFile(path, info); err != nil {
|
||||
file.Close()
|
||||
return nil, nil, err
|
||||
}
|
||||
return file, info, nil
|
||||
}
|
||||
|
||||
func regularRootFileInfo(root *os.Root, rel, path string) (os.FileInfo, error) {
|
||||
info, err := root.Lstat(rel)
|
||||
if err != nil {
|
||||
return nil, rootPathError(err)
|
||||
}
|
||||
// Reject symlinks outright. os.Root.Open follows symlinks via openat
|
||||
// without O_NOFOLLOW, so a symlink inside the working root that points
|
||||
// outside it (e.g. ./notes -> ~/.ssh/id_rsa) would otherwise be read
|
||||
// transparently, bypassing the working-directory confinement that the
|
||||
// bash denylist enforces for direct credential reads. The caller must
|
||||
// operate on the real target file instead.
|
||||
if info.Mode()&os.ModeSymlink != 0 {
|
||||
return nil, fmt.Errorf("%s is a symlink; read the target file directly", path)
|
||||
}
|
||||
if err := rejectNonRegularFile(path, info); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return info, nil
|
||||
}
|
||||
|
||||
func rejectNonRegularFile(path string, info os.FileInfo) error {
|
||||
if info.IsDir() {
|
||||
return fmt.Errorf("%s is a directory", path)
|
||||
}
|
||||
if !info.Mode().IsRegular() {
|
||||
return fmt.Errorf("%s is not a regular file", path)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func writeFileAtomic(workingDir, path string, data []byte, perm os.FileMode) error {
|
||||
rel, err := cleanRelativePath(path)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
root, err := openWorkingRoot(workingDir)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer root.Close()
|
||||
if err := rejectRootFinalSymlink(root, rel, path); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
parent, name := filepath.Split(rel)
|
||||
tmpBase := fmt.Sprintf(".%s.ollama-tmp-%d", name, os.Getpid())
|
||||
for i := 0; ; i++ {
|
||||
candidateName := tmpBase
|
||||
if i > 0 {
|
||||
candidateName = fmt.Sprintf("%s-%d", tmpBase, i)
|
||||
}
|
||||
candidate := filepath.Join(parent, candidateName)
|
||||
file, err := root.OpenFile(candidate, os.O_WRONLY|os.O_CREATE|os.O_EXCL, perm)
|
||||
if os.IsExist(err) {
|
||||
continue
|
||||
}
|
||||
if err != nil {
|
||||
return rootPathError(err)
|
||||
}
|
||||
if err := file.Chmod(perm); err != nil {
|
||||
closeErr := file.Close()
|
||||
_ = root.Remove(candidate)
|
||||
if closeErr != nil {
|
||||
return closeErr
|
||||
}
|
||||
return err
|
||||
}
|
||||
writeErr := writeAllAndSync(file, data)
|
||||
closeErr := file.Close()
|
||||
if writeErr != nil || closeErr != nil {
|
||||
_ = root.Remove(candidate)
|
||||
if writeErr != nil {
|
||||
return writeErr
|
||||
}
|
||||
return closeErr
|
||||
}
|
||||
if err := root.Rename(candidate, rel); err != nil {
|
||||
_ = root.Remove(candidate)
|
||||
return rootPathError(err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
func rejectFinalSymlink(workingDir, path string) error {
|
||||
rel, err := cleanRelativePath(path)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
root, err := openWorkingRoot(workingDir)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer root.Close()
|
||||
return rejectRootFinalSymlink(root, rel, path)
|
||||
}
|
||||
|
||||
func rejectRootFinalSymlink(root *os.Root, rel, path string) error {
|
||||
info, err := root.Lstat(rel)
|
||||
if err != nil {
|
||||
return rootPathError(err)
|
||||
}
|
||||
if info.Mode()&os.ModeSymlink != 0 {
|
||||
return fmt.Errorf("%s is a symlink; edit the target file directly", path)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func rootPathError(err error) error {
|
||||
if err != nil && strings.Contains(err.Error(), "path escapes") {
|
||||
return fmt.Errorf("path escapes working directory")
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
func openWorkingRoot(workingDir string) (*os.Root, error) {
|
||||
base, err := workingDirAbs(workingDir)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return os.OpenRoot(base)
|
||||
}
|
||||
|
||||
func writeAllAndSync(file *os.File, data []byte) error {
|
||||
if _, err := file.Write(data); err != nil {
|
||||
return err
|
||||
}
|
||||
return file.Sync()
|
||||
}
|
||||
|
||||
func readAllWithinLimit(reader io.Reader, limit int) ([]byte, error) {
|
||||
if limit < 0 {
|
||||
limit = 0
|
||||
}
|
||||
content, err := io.ReadAll(io.LimitReader(reader, int64(limit)+1))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if len(content) > limit {
|
||||
return nil, fmt.Errorf("content is too large (%d byte limit)", limit)
|
||||
}
|
||||
return content, nil
|
||||
}
|
||||
|
||||
func workingDirAbs(workingDir string) (string, error) {
|
||||
base := workingDir
|
||||
if base == "" {
|
||||
var err error
|
||||
base, err = os.Getwd()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
}
|
||||
return canonicalPath(base)
|
||||
}
|
||||
|
||||
func canonicalPath(path string) (string, error) {
|
||||
abs, err := filepath.Abs(path)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
resolved, err := filepath.EvalSymlinks(abs)
|
||||
if err == nil {
|
||||
return resolved, nil
|
||||
}
|
||||
return abs, nil
|
||||
}
|
||||
|
||||
type readSelection struct {
|
||||
enabled bool
|
||||
start int
|
||||
end int
|
||||
}
|
||||
|
||||
func readSelectionFromArgs(args map[string]any) (readSelection, error) {
|
||||
selection := readSelection{start: 1}
|
||||
|
||||
if start, ok, err := intReadArg(args, "start"); err != nil {
|
||||
return readSelection{}, err
|
||||
} else if ok {
|
||||
selection.enabled = true
|
||||
selection.start = start
|
||||
}
|
||||
if end, ok, err := intReadArg(args, "end"); err != nil {
|
||||
return readSelection{}, err
|
||||
} else if ok {
|
||||
selection.enabled = true
|
||||
selection.end = end
|
||||
}
|
||||
|
||||
if !selection.enabled {
|
||||
return selection, nil
|
||||
}
|
||||
if selection.start < 1 {
|
||||
return readSelection{}, fmt.Errorf("start must be greater than 0")
|
||||
}
|
||||
if selection.end > 0 && selection.end < selection.start {
|
||||
return readSelection{}, fmt.Errorf("end must be greater than or equal to start")
|
||||
}
|
||||
return selection, nil
|
||||
}
|
||||
|
||||
func readLineSelection(file *os.File, selection readSelection) (string, error) {
|
||||
reader := bufio.NewReader(file)
|
||||
var b strings.Builder
|
||||
for lineNo := 1; ; {
|
||||
line, err := reader.ReadSlice('\n')
|
||||
if lineNo >= selection.start && (selection.end == 0 || lineNo <= selection.end) {
|
||||
if b.Len()+len(line) > maxReadBytes {
|
||||
return "", fmt.Errorf("selected content is too large (%d byte limit)", maxReadBytes)
|
||||
}
|
||||
b.Write(line)
|
||||
}
|
||||
if err != nil {
|
||||
if err == bufio.ErrBufferFull {
|
||||
continue
|
||||
}
|
||||
if err == io.EOF {
|
||||
break
|
||||
}
|
||||
return "", err
|
||||
}
|
||||
if selection.end > 0 && lineNo >= selection.end {
|
||||
break
|
||||
}
|
||||
lineNo++
|
||||
}
|
||||
return b.String(), nil
|
||||
}
|
||||
|
||||
func intReadArg(args map[string]any, key string) (int, bool, error) {
|
||||
value, ok := args[key]
|
||||
if !ok {
|
||||
return 0, false, nil
|
||||
}
|
||||
switch v := value.(type) {
|
||||
case int:
|
||||
return v, true, nil
|
||||
case int64:
|
||||
return int(v), true, nil
|
||||
case float64:
|
||||
if v != float64(int(v)) {
|
||||
return 0, true, fmt.Errorf("%s must be a whole number", key)
|
||||
}
|
||||
return int(v), true, nil
|
||||
case string:
|
||||
v = strings.TrimSpace(v)
|
||||
if v == "" {
|
||||
return 0, false, nil
|
||||
}
|
||||
n, err := strconv.Atoi(v)
|
||||
if err != nil {
|
||||
return 0, true, fmt.Errorf("%s must be a whole number", key)
|
||||
}
|
||||
return n, true, nil
|
||||
default:
|
||||
return 0, true, fmt.Errorf("%s must be a whole number", key)
|
||||
}
|
||||
}
|
||||
|
||||
func plural(n int) string {
|
||||
if n == 1 {
|
||||
return ""
|
||||
}
|
||||
return "s"
|
||||
}
|
||||
@@ -0,0 +1,338 @@
|
||||
package tools
|
||||
|
||||
import (
|
||||
"context"
|
||||
"io"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/ollama/ollama/agent"
|
||||
)
|
||||
|
||||
func TestEditReplacesUniqueText(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
path := filepath.Join(dir, "note.txt")
|
||||
if err := os.WriteFile(path, []byte("hello world\n"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
result, err := (&Edit{}).Execute(context.Background(), agent.ToolContext{WorkingDir: dir}, map[string]any{
|
||||
"path": "note.txt",
|
||||
"old_text": "hello",
|
||||
"new_text": "hi",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !strings.Contains(result.Content, "Updated note.txt") {
|
||||
t.Fatalf("result = %q", result.Content)
|
||||
}
|
||||
|
||||
content, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if string(content) != "hi world\n" {
|
||||
t.Fatalf("content = %q", content)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEditRequiresUniqueMatchByDefault(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
path := filepath.Join(dir, "note.txt")
|
||||
if err := os.WriteFile(path, []byte("same same\n"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
_, err := (&Edit{}).Execute(context.Background(), agent.ToolContext{WorkingDir: dir}, map[string]any{
|
||||
"path": "note.txt",
|
||||
"old_text": "same",
|
||||
"new_text": "other",
|
||||
})
|
||||
if err == nil {
|
||||
t.Fatal("expected ambiguous edit to fail")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "matched 2 times") {
|
||||
t.Fatalf("err = %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEditRejectsEscapingPath(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
_, err := (&Edit{}).Execute(context.Background(), agent.ToolContext{WorkingDir: dir}, map[string]any{
|
||||
"path": "../outside.txt",
|
||||
"old_text": "old",
|
||||
"new_text": "new",
|
||||
})
|
||||
if err == nil {
|
||||
t.Fatal("expected escaping path to fail")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "path escapes working directory") {
|
||||
t.Fatalf("err = %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEditRejectsSymlinkEscape(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
outside := t.TempDir()
|
||||
if err := os.WriteFile(filepath.Join(outside, "note.txt"), []byte("old\n"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := os.Symlink(outside, filepath.Join(dir, "link")); err != nil {
|
||||
t.Skipf("symlinks unavailable: %v", err)
|
||||
}
|
||||
|
||||
_, err := (&Edit{}).Execute(context.Background(), agent.ToolContext{WorkingDir: dir}, map[string]any{
|
||||
"path": filepath.Join("link", "note.txt"),
|
||||
"old_text": "old",
|
||||
"new_text": "new",
|
||||
})
|
||||
if err == nil {
|
||||
t.Fatal("expected symlink escape to fail")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "path escapes working directory") {
|
||||
t.Fatalf("err = %v", err)
|
||||
}
|
||||
|
||||
content, err := os.ReadFile(filepath.Join(outside, "note.txt"))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if string(content) != "old\n" {
|
||||
t.Fatalf("outside content changed to %q", content)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEditRejectsFinalSymlink(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
target := filepath.Join(dir, "target.txt")
|
||||
if err := os.WriteFile(target, []byte("old\n"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
link := filepath.Join(dir, "link.txt")
|
||||
if err := os.Symlink("target.txt", link); err != nil {
|
||||
t.Skipf("symlinks unavailable: %v", err)
|
||||
}
|
||||
|
||||
_, err := (&Edit{}).Execute(context.Background(), agent.ToolContext{WorkingDir: dir}, map[string]any{
|
||||
"path": "link.txt",
|
||||
"old_text": "old",
|
||||
"new_text": "new",
|
||||
})
|
||||
if err == nil {
|
||||
t.Fatal("expected final symlink edit to fail")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "is a symlink") {
|
||||
t.Fatalf("err = %v", err)
|
||||
}
|
||||
content, err := os.ReadFile(target)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if string(content) != "old\n" {
|
||||
t.Fatalf("target content changed to %q", content)
|
||||
}
|
||||
info, err := os.Lstat(link)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if info.Mode()&os.ModeSymlink == 0 {
|
||||
t.Fatalf("link mode = %v, want symlink", info.Mode())
|
||||
}
|
||||
}
|
||||
|
||||
func TestReadRejectsParentOutsideCurrentWorkingDir(t *testing.T) {
|
||||
root := t.TempDir()
|
||||
subdir := filepath.Join(root, "sub")
|
||||
if err := os.Mkdir(subdir, 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := os.WriteFile(filepath.Join(root, "note.txt"), []byte("hello"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
_, err := (&Read{}).Execute(context.Background(), agent.ToolContext{WorkingDir: subdir}, map[string]any{
|
||||
"path": "../note.txt",
|
||||
})
|
||||
if err == nil {
|
||||
t.Fatal("expected parent path to fail")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "path escapes working directory") {
|
||||
t.Fatalf("err = %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestReadRequiresApproval(t *testing.T) {
|
||||
if !agent.ToolRequiresApproval((&Read{}), map[string]any{"path": "note.txt"}) {
|
||||
t.Fatal("read should require approval")
|
||||
}
|
||||
}
|
||||
|
||||
func TestReadDefaultsToEntireFile(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
content := "one\ntwo\nthree\n"
|
||||
if err := os.WriteFile(filepath.Join(dir, "note.txt"), []byte(content), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
result, err := (&Read{}).Execute(context.Background(), agent.ToolContext{WorkingDir: dir}, map[string]any{
|
||||
"path": "note.txt",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if result.Content != content {
|
||||
t.Fatalf("content = %q", result.Content)
|
||||
}
|
||||
}
|
||||
|
||||
func TestReadAllowsAbsolutePath(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
path := filepath.Join(dir, "note.txt")
|
||||
content := "one\ntwo\nthree\n"
|
||||
if err := os.WriteFile(path, []byte(content), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
result, err := (&Read{}).Execute(context.Background(), agent.ToolContext{WorkingDir: t.TempDir()}, map[string]any{
|
||||
"path": path,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if result.Content != content {
|
||||
t.Fatalf("content = %q", result.Content)
|
||||
}
|
||||
}
|
||||
|
||||
func TestReadRejectsAbsoluteSymlink(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
target := filepath.Join(dir, "target.txt")
|
||||
if err := os.WriteFile(target, []byte("hello\n"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
link := filepath.Join(dir, "alias")
|
||||
if err := os.Symlink(target, link); err != nil {
|
||||
t.Skipf("symlinks unavailable: %v", err)
|
||||
}
|
||||
|
||||
_, err := (&Read{}).Execute(context.Background(), agent.ToolContext{WorkingDir: t.TempDir()}, map[string]any{
|
||||
"path": link,
|
||||
})
|
||||
if err == nil {
|
||||
t.Fatal("expected absolute symlink to be rejected")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "symlink") {
|
||||
t.Fatalf("err = %v, want symlink rejection", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestReadStartEnd(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
if err := os.WriteFile(filepath.Join(dir, "note.txt"), []byte("one\ntwo\nthree\nfour\n"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
result, err := (&Read{}).Execute(context.Background(), agent.ToolContext{WorkingDir: dir}, map[string]any{
|
||||
"path": "note.txt",
|
||||
"start": 2,
|
||||
"end": 3,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if result.Content != "two\nthree\n" {
|
||||
t.Fatalf("content = %q", result.Content)
|
||||
}
|
||||
}
|
||||
|
||||
func TestReadStartOnly(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
if err := os.WriteFile(filepath.Join(dir, "note.txt"), []byte("one\ntwo\nthree\nfour\n"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
result, err := (&Read{}).Execute(context.Background(), agent.ToolContext{WorkingDir: dir}, map[string]any{
|
||||
"path": "note.txt",
|
||||
"start": 3,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if result.Content != "three\nfour\n" {
|
||||
t.Fatalf("content = %q", result.Content)
|
||||
}
|
||||
}
|
||||
|
||||
func TestReadEndOnly(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
if err := os.WriteFile(filepath.Join(dir, "note.txt"), []byte("one\ntwo\nthree\nfour\n"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
result, err := (&Read{}).Execute(context.Background(), agent.ToolContext{WorkingDir: dir}, map[string]any{
|
||||
"path": "note.txt",
|
||||
"end": 2,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if result.Content != "one\ntwo\n" {
|
||||
t.Fatalf("content = %q", result.Content)
|
||||
}
|
||||
}
|
||||
|
||||
func TestReadSelectionRejectsHugeSingleLine(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
if err := os.WriteFile(filepath.Join(dir, "note.txt"), []byte(strings.Repeat("x", maxReadBytes+1)), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
_, err := (&Read{}).Execute(context.Background(), agent.ToolContext{WorkingDir: dir}, map[string]any{
|
||||
"path": "note.txt",
|
||||
"start": 1,
|
||||
"end": 1,
|
||||
})
|
||||
if err == nil {
|
||||
t.Fatal("expected huge selected line to fail")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "selected content is too large") {
|
||||
t.Fatalf("err = %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestReadAllWithinLimitRejectsGrowingRead(t *testing.T) {
|
||||
reader := io.MultiReader(
|
||||
strings.NewReader(strings.Repeat("x", maxReadBytes)),
|
||||
strings.NewReader("x"),
|
||||
)
|
||||
|
||||
_, err := readAllWithinLimit(reader, maxReadBytes)
|
||||
if err == nil {
|
||||
t.Fatal("expected over-limit read to fail")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "content is too large") {
|
||||
t.Fatalf("err = %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestReadRejectsInvalidRange(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
if err := os.WriteFile(filepath.Join(dir, "note.txt"), []byte("one\ntwo\n"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
_, err := (&Read{}).Execute(context.Background(), agent.ToolContext{WorkingDir: dir}, map[string]any{
|
||||
"path": "note.txt",
|
||||
"start": 4,
|
||||
"end": 2,
|
||||
})
|
||||
if err == nil {
|
||||
t.Fatal("expected invalid range to fail")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "end must") {
|
||||
t.Fatalf("err = %v", err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,121 @@
|
||||
//go:build !windows
|
||||
|
||||
package tools
|
||||
|
||||
import (
|
||||
"context"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"syscall"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/ollama/ollama/agent"
|
||||
)
|
||||
|
||||
func TestOpenRegularFileRejectsFIFO(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
path := filepath.Join(dir, "pipe")
|
||||
if err := syscall.Mkfifo(path, 0o600); err != nil {
|
||||
t.Skipf("mkfifo unavailable: %v", err)
|
||||
}
|
||||
|
||||
done := make(chan error, 1)
|
||||
go func() {
|
||||
file, _, err := openRegularFile(dir, "pipe", false)
|
||||
if file != nil {
|
||||
file.Close()
|
||||
}
|
||||
done <- err
|
||||
}()
|
||||
|
||||
select {
|
||||
case err := <-done:
|
||||
if err == nil {
|
||||
t.Fatal("expected FIFO to be rejected")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "not a regular file") {
|
||||
t.Fatalf("err = %v", err)
|
||||
}
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("openRegularFile blocked on FIFO")
|
||||
}
|
||||
}
|
||||
|
||||
func TestEditPreservesModeDespiteUmask(t *testing.T) {
|
||||
oldUmask := syscall.Umask(0o077)
|
||||
defer syscall.Umask(oldUmask)
|
||||
|
||||
dir := t.TempDir()
|
||||
path := filepath.Join(dir, "note.txt")
|
||||
if err := os.WriteFile(path, []byte("hello\n"), 0o666); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := os.Chmod(path, 0o666); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
_, err := (&Edit{}).Execute(context.Background(), agent.ToolContext{WorkingDir: dir}, map[string]any{
|
||||
"path": "note.txt",
|
||||
"old_text": "hello",
|
||||
"new_text": "hi",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
info, err := os.Stat(path)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if got := info.Mode().Perm(); got != 0o666 {
|
||||
t.Fatalf("mode = %#o, want 0666", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestReadRejectsSymlinkEscapingWorkingDir(t *testing.T) {
|
||||
root := t.TempDir()
|
||||
secret := filepath.Join(t.TempDir(), "secret.txt")
|
||||
if err := os.WriteFile(secret, []byte("top secret\n"), 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
link := filepath.Join(root, "notes")
|
||||
if err := os.Symlink(secret, link); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
_, err := (&Read{}).Execute(context.Background(), agent.ToolContext{WorkingDir: root}, map[string]any{
|
||||
"path": "notes",
|
||||
})
|
||||
if err == nil {
|
||||
t.Fatal("expected symlink escaping working dir to be rejected")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "symlink") {
|
||||
t.Fatalf("err = %v, want symlink rejection", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestReadRejectsSymlinkInsideWorkingDirToOutside(t *testing.T) {
|
||||
root := t.TempDir()
|
||||
target := filepath.Join(root, "real.txt")
|
||||
if err := os.WriteFile(target, []byte("hello\n"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// A symlink to a sibling file still resolves inside the root; Read must
|
||||
// reject it regardless, consistent with Edit's rejectFinalSymlink.
|
||||
link := filepath.Join(root, "alias")
|
||||
if err := os.Symlink(target, link); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
_, err := (&Read{}).Execute(context.Background(), agent.ToolContext{WorkingDir: root}, map[string]any{
|
||||
"path": "alias",
|
||||
})
|
||||
if err == nil {
|
||||
t.Fatal("expected symlink to be rejected even when target is inside root")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "symlink") {
|
||||
t.Fatalf("err = %v, want symlink rejection", err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,38 @@
|
||||
package tools
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
|
||||
"github.com/ollama/ollama/agent"
|
||||
"github.com/ollama/ollama/api"
|
||||
)
|
||||
|
||||
// Skill is the model-facing adapter for the core agent skill catalog.
|
||||
// It only supplies instructions; regular tools retain their own approval
|
||||
// requirements for filesystem or network access.
|
||||
type Skill struct{ Catalog *agent.SkillCatalog }
|
||||
|
||||
func (t *Skill) Name() string { return "skill" }
|
||||
|
||||
func (t *Skill) Description() string {
|
||||
return "Load a named Ollama skill and return its instructions."
|
||||
}
|
||||
|
||||
func (t *Skill) Schema() api.ToolFunction {
|
||||
props := api.NewToolPropertiesMap()
|
||||
props.Set("name", api.ToolProperty{Type: api.PropertyType{"string"}, Description: "Name of the skill to load."})
|
||||
return api.ToolFunction{Name: t.Name(), Description: t.Description(), Parameters: api.ToolFunctionParameters{Type: "object", Properties: props, Required: []string{"name"}}}
|
||||
}
|
||||
|
||||
func (t *Skill) Execute(_ context.Context, _ agent.ToolContext, args map[string]any) (agent.ToolResult, error) {
|
||||
name, ok := args["name"].(string)
|
||||
if !ok {
|
||||
return agent.ToolResult{}, errors.New("name parameter is required")
|
||||
}
|
||||
skill, err := t.Catalog.Load(name)
|
||||
if err != nil {
|
||||
return agent.ToolResult{}, err
|
||||
}
|
||||
return agent.ToolResult{Content: skill.Content()}, nil
|
||||
}
|
||||
@@ -0,0 +1,34 @@
|
||||
package tools
|
||||
|
||||
import (
|
||||
"context"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/ollama/ollama/agent"
|
||||
)
|
||||
|
||||
func TestSkillLoadsCoreCatalogWithoutApproval(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
path := filepath.Join(dir, "release-notes")
|
||||
if err := os.Mkdir(path, 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := os.WriteFile(filepath.Join(path, "SKILL.md"), []byte("---\nname: release-notes\ndescription: Draft release notes.\n---\nUse concise bullets."), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
catalog, err := agent.DiscoverSkills(dir)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
tool := &Skill{Catalog: catalog}
|
||||
if agent.ToolRequiresApproval(tool, map[string]any{"name": "release-notes"}) {
|
||||
t.Fatal("loading a skill must not change ordinary tool approval semantics")
|
||||
}
|
||||
result, err := tool.Execute(context.Background(), agent.ToolContext{}, map[string]any{"name": "release-notes"})
|
||||
if err != nil || !strings.Contains(result.Content, "Use concise bullets.") {
|
||||
t.Fatalf("tool result = %#v, %v", result, err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,186 @@
|
||||
package tools
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/url"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/ollama/ollama/agent"
|
||||
"github.com/ollama/ollama/api"
|
||||
internalcloud "github.com/ollama/ollama/internal/cloud"
|
||||
)
|
||||
|
||||
const (
|
||||
maxWebFetchContentRunes = 60_000
|
||||
webSearchTimeout = 15 * time.Second
|
||||
webFetchTimeout = 30 * time.Second
|
||||
)
|
||||
|
||||
var ErrWebAuthRequired = errors.New("Not authenticated. Run `ollama signin` and try again.")
|
||||
|
||||
type WebSearch struct{}
|
||||
|
||||
func (w *WebSearch) Name() string {
|
||||
return "web_search"
|
||||
}
|
||||
|
||||
func (w *WebSearch) Description() string {
|
||||
return "Search the web for current information that may not be in the model's training data."
|
||||
}
|
||||
|
||||
func (w *WebSearch) Schema() api.ToolFunction {
|
||||
props := api.NewToolPropertiesMap()
|
||||
props.Set("query", api.ToolProperty{
|
||||
Type: api.PropertyType{"string"},
|
||||
Description: "The search query to look up on the web.",
|
||||
})
|
||||
return api.ToolFunction{
|
||||
Name: w.Name(),
|
||||
Description: w.Description(),
|
||||
Parameters: api.ToolFunctionParameters{
|
||||
Type: "object",
|
||||
Properties: props,
|
||||
Required: []string{"query"},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func (w *WebSearch) RequiresApproval(map[string]any) bool {
|
||||
return true
|
||||
}
|
||||
|
||||
func (w *WebSearch) Execute(ctx context.Context, _ agent.ToolContext, args map[string]any) (agent.ToolResult, error) {
|
||||
// TODO: use shared agent.RequiredStringArg for the "query" parameter (see agent package cleanup plan).
|
||||
if internalcloud.Disabled() {
|
||||
return agent.ToolResult{}, errors.New(internalcloud.DisabledError("web search is unavailable"))
|
||||
}
|
||||
query, ok := args["query"].(string)
|
||||
if !ok || strings.TrimSpace(query) == "" {
|
||||
return agent.ToolResult{}, fmt.Errorf("query parameter is required")
|
||||
}
|
||||
|
||||
client, err := api.ClientFromEnvironment()
|
||||
if err != nil {
|
||||
return agent.ToolResult{}, err
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithTimeout(ctx, webSearchTimeout)
|
||||
defer cancel()
|
||||
|
||||
searchResp, err := client.WebSearchExperimental(ctx, &api.WebSearchRequest{Query: query, MaxResults: 5})
|
||||
if err != nil {
|
||||
var authErr api.AuthorizationError
|
||||
if errors.As(err, &authErr) {
|
||||
return agent.ToolResult{}, ErrWebAuthRequired
|
||||
}
|
||||
return agent.ToolResult{}, err
|
||||
}
|
||||
if len(searchResp.Results) == 0 {
|
||||
return agent.ToolResult{Content: "No results found for query: " + query}, nil
|
||||
}
|
||||
|
||||
var sb strings.Builder
|
||||
sb.WriteString(fmt.Sprintf("Search results for: %s\n\n", query))
|
||||
for i, result := range searchResp.Results {
|
||||
sb.WriteString(fmt.Sprintf("%d. %s\n", i+1, result.Title))
|
||||
sb.WriteString(fmt.Sprintf(" URL: %s\n", result.URL))
|
||||
if result.Content != "" {
|
||||
content := []rune(result.Content)
|
||||
if len(content) > 300 {
|
||||
content = append(content[:300], []rune("...")...)
|
||||
}
|
||||
sb.WriteString(fmt.Sprintf(" %s\n", string(content)))
|
||||
}
|
||||
sb.WriteByte('\n')
|
||||
}
|
||||
return agent.ToolResult{Content: sb.String()}, nil
|
||||
}
|
||||
|
||||
type WebFetch struct{}
|
||||
|
||||
func (w *WebFetch) Name() string {
|
||||
return "web_fetch"
|
||||
}
|
||||
|
||||
func (w *WebFetch) Description() string {
|
||||
return "Fetch and extract text content from a web page."
|
||||
}
|
||||
|
||||
func (w *WebFetch) Schema() api.ToolFunction {
|
||||
props := api.NewToolPropertiesMap()
|
||||
props.Set("url", api.ToolProperty{
|
||||
Type: api.PropertyType{"string"},
|
||||
Description: "The URL to fetch and extract content from.",
|
||||
})
|
||||
return api.ToolFunction{
|
||||
Name: w.Name(),
|
||||
Description: w.Description(),
|
||||
Parameters: api.ToolFunctionParameters{
|
||||
Type: "object",
|
||||
Properties: props,
|
||||
Required: []string{"url"},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func (w *WebFetch) RequiresApproval(map[string]any) bool {
|
||||
return true
|
||||
}
|
||||
|
||||
func (w *WebFetch) Execute(ctx context.Context, _ agent.ToolContext, args map[string]any) (agent.ToolResult, error) {
|
||||
// TODO: use shared agent.RequiredStringArg for the "url" parameter (see agent package cleanup plan).
|
||||
if internalcloud.Disabled() {
|
||||
return agent.ToolResult{}, errors.New(internalcloud.DisabledError("web fetch is unavailable"))
|
||||
}
|
||||
urlStr, ok := args["url"].(string)
|
||||
if !ok || strings.TrimSpace(urlStr) == "" {
|
||||
return agent.ToolResult{}, fmt.Errorf("url parameter is required")
|
||||
}
|
||||
parsed, err := url.Parse(urlStr)
|
||||
if err != nil {
|
||||
return agent.ToolResult{}, fmt.Errorf("invalid URL: %w", err)
|
||||
}
|
||||
if scheme := strings.ToLower(parsed.Scheme); scheme != "http" && scheme != "https" {
|
||||
return agent.ToolResult{}, fmt.Errorf("unsupported URL scheme %q: only http and https are allowed", parsed.Scheme)
|
||||
}
|
||||
|
||||
client, err := api.ClientFromEnvironment()
|
||||
if err != nil {
|
||||
return agent.ToolResult{}, err
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithTimeout(ctx, webFetchTimeout)
|
||||
defer cancel()
|
||||
|
||||
fetchResp, err := client.WebFetchExperimental(ctx, &api.WebFetchRequest{URL: urlStr})
|
||||
if err != nil {
|
||||
var authErr api.AuthorizationError
|
||||
if errors.As(err, &authErr) {
|
||||
return agent.ToolResult{}, ErrWebAuthRequired
|
||||
}
|
||||
return agent.ToolResult{}, err
|
||||
}
|
||||
|
||||
var sb strings.Builder
|
||||
if fetchResp.Title != "" {
|
||||
sb.WriteString(fmt.Sprintf("Title: %s\n\n", fetchResp.Title))
|
||||
}
|
||||
if fetchResp.Content != "" {
|
||||
sb.WriteString("Content:\n")
|
||||
sb.WriteString(truncateWebFetchContent(fetchResp.Content))
|
||||
} else {
|
||||
sb.WriteString("No content could be extracted from the page.")
|
||||
}
|
||||
return agent.ToolResult{Content: sb.String()}, nil
|
||||
}
|
||||
|
||||
func truncateWebFetchContent(content string) string {
|
||||
return agent.Truncate(content, agent.TruncateConfig{
|
||||
MaxRunes: maxWebFetchContentRunes,
|
||||
Label: "tool output",
|
||||
Hint: "Use a narrower request or search query if more detail is needed.",
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,214 @@
|
||||
package tools
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
coreagent "github.com/ollama/ollama/agent"
|
||||
"github.com/ollama/ollama/api"
|
||||
"github.com/ollama/ollama/envconfig"
|
||||
internalcloud "github.com/ollama/ollama/internal/cloud"
|
||||
)
|
||||
|
||||
func TestWebToolsRequireApproval(t *testing.T) {
|
||||
if !coreagent.ToolRequiresApproval((&WebSearch{}), map[string]any{"query": "ollama"}) {
|
||||
t.Fatal("web search should require approval")
|
||||
}
|
||||
if !coreagent.ToolRequiresApproval((&WebFetch{}), map[string]any{"url": "https://ollama.com"}) {
|
||||
t.Fatal("web fetch should require approval")
|
||||
}
|
||||
}
|
||||
|
||||
var webToolCases = []struct {
|
||||
name string
|
||||
tool coreagent.Tool
|
||||
args map[string]any
|
||||
path string
|
||||
operation string
|
||||
}{
|
||||
{"search", &WebSearch{}, map[string]any{"query": "ollama"}, "/api/experimental/web_search", "web search is unavailable"},
|
||||
{"fetch", &WebFetch{}, map[string]any{"url": "https://ollama.com"}, "/api/experimental/web_fetch", "web fetch is unavailable"},
|
||||
}
|
||||
|
||||
// enableWebToolsForTest isolates web tool tests from the runner's cloud
|
||||
// policy. In particular, Windows can inherit both OLLAMA_NO_CLOUD and a
|
||||
// server.json from USERPROFILE.
|
||||
func enableWebToolsForTest(t *testing.T) {
|
||||
t.Helper()
|
||||
|
||||
// Register before t.Setenv so the cache is refreshed after t.Setenv has
|
||||
// restored the runner's environment during cleanup.
|
||||
t.Cleanup(envconfig.ReloadServerConfig)
|
||||
|
||||
home := t.TempDir()
|
||||
t.Setenv("HOME", home)
|
||||
t.Setenv("USERPROFILE", home)
|
||||
t.Setenv("OLLAMA_NO_CLOUD", "")
|
||||
envconfig.ReloadServerConfig()
|
||||
}
|
||||
|
||||
// runWebTool executes tool against a stub server that responds to every
|
||||
// request with status and body, returning the resulting error.
|
||||
func runWebTool(t *testing.T, tool coreagent.Tool, args map[string]any, path string, status int, body string) error {
|
||||
t.Helper()
|
||||
ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.URL.Path != path {
|
||||
t.Fatalf("path = %q, want %q", r.URL.Path, path)
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
w.WriteHeader(status)
|
||||
_, _ = w.Write([]byte(body))
|
||||
}))
|
||||
t.Cleanup(ts.Close)
|
||||
t.Setenv("OLLAMA_HOST", ts.URL)
|
||||
_, err := tool.Execute(t.Context(), coreagent.ToolContext{}, args)
|
||||
return err
|
||||
}
|
||||
|
||||
func TestWebToolsReportAuthenticationError(t *testing.T) {
|
||||
enableWebToolsForTest(t)
|
||||
|
||||
for _, tt := range webToolCases {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
err := runWebTool(t, tt.tool, tt.args, tt.path, http.StatusUnauthorized,
|
||||
`{"error":"unauthorized","signin_url":"https://ollama.com/signin"}`)
|
||||
if !errors.Is(err, ErrWebAuthRequired) {
|
||||
t.Fatalf("error = %v, want %v", err, ErrWebAuthRequired)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestWebToolsPreserveNonAuthenticationErrors(t *testing.T) {
|
||||
enableWebToolsForTest(t)
|
||||
|
||||
for _, tt := range webToolCases {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
err := runWebTool(t, tt.tool, tt.args, tt.path, http.StatusTooManyRequests,
|
||||
`{"error":"web search quota exceeded"}`)
|
||||
if err == nil {
|
||||
t.Fatal("expected error")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "web search quota exceeded") {
|
||||
t.Fatalf("error = %q, want original error message", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestWebToolsIgnoreInheritedCloudPolicy(t *testing.T) {
|
||||
// This cleanup is registered before the test environment, so it restores
|
||||
// the server config cache after t.Setenv restores the runner's values.
|
||||
t.Cleanup(envconfig.ReloadServerConfig)
|
||||
|
||||
home := t.TempDir()
|
||||
configPath := filepath.Join(home, ".ollama", "server.json")
|
||||
if err := os.MkdirAll(filepath.Dir(configPath), 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := os.WriteFile(configPath, []byte(`{"disable_ollama_cloud":true}`), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Setenv("HOME", home)
|
||||
t.Setenv("USERPROFILE", home)
|
||||
t.Setenv("OLLAMA_NO_CLOUD", "1")
|
||||
envconfig.ReloadServerConfig()
|
||||
|
||||
enableWebToolsForTest(t)
|
||||
err := runWebTool(t, &WebSearch{}, map[string]any{"query": "ollama"}, "/api/experimental/web_search", http.StatusUnauthorized,
|
||||
`{"error":"unauthorized","signin_url":"https://ollama.com/signin"}`)
|
||||
if !errors.Is(err, ErrWebAuthRequired) {
|
||||
t.Fatalf("error = %v, want %v", err, ErrWebAuthRequired)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWebFetchRejectsUnsupportedScheme(t *testing.T) {
|
||||
enableWebToolsForTest(t)
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
url string
|
||||
wantErr bool
|
||||
}{
|
||||
{name: "file scheme", url: "file:///etc/passwd", wantErr: true},
|
||||
{name: "data scheme", url: "data:text/plain,secret", wantErr: true},
|
||||
{name: "ftp scheme", url: "ftp://example.com/secret", wantErr: true},
|
||||
{name: "http allowed", url: "http://example.com", wantErr: false},
|
||||
{name: "https allowed", url: "https://example.com", wantErr: false},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
_, err := (&WebFetch{}).Execute(t.Context(), coreagent.ToolContext{}, map[string]any{"url": tt.url})
|
||||
if tt.wantErr && err == nil {
|
||||
t.Fatal("expected unsupported scheme to be rejected")
|
||||
}
|
||||
// For allowed schemes we expect an error only from the missing
|
||||
// server/auth path, not from scheme validation. The http/https
|
||||
// cases reach the client and may fail on connection/auth; we only
|
||||
// assert that the error is NOT a scheme error.
|
||||
if !tt.wantErr && err != nil && strings.Contains(err.Error(), "unsupported URL scheme") {
|
||||
t.Fatalf("http/https rejected as unsupported: %v", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestWebFetchBoundsContentBeforeReturning(t *testing.T) {
|
||||
enableWebToolsForTest(t)
|
||||
|
||||
ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.URL.Path != "/api/experimental/web_fetch" {
|
||||
t.Fatalf("path = %q, want /api/experimental/web_fetch", r.URL.Path)
|
||||
}
|
||||
var req api.WebFetchRequest
|
||||
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if req.URL != "https://ollama.com" {
|
||||
t.Fatalf("request URL = %q, want https://ollama.com", req.URL)
|
||||
}
|
||||
if err := json.NewEncoder(w).Encode(api.WebFetchResponse{
|
||||
Title: "Ollama",
|
||||
Content: strings.Repeat("x", maxWebFetchContentRunes+25),
|
||||
}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}))
|
||||
defer ts.Close()
|
||||
t.Setenv("OLLAMA_HOST", ts.URL)
|
||||
|
||||
result, err := (&WebFetch{}).Execute(t.Context(), coreagent.ToolContext{}, map[string]any{
|
||||
"url": "https://ollama.com",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !strings.Contains(result.Content, "[tool output truncated: showing first ~") ||
|
||||
!strings.Contains(result.Content, "omitted ~7 tokens") ||
|
||||
!strings.Contains(result.Content, "Use a narrower request or search query") {
|
||||
t.Fatalf("content missing truncation marker: %q", result.Content)
|
||||
}
|
||||
if count := strings.Count(result.Content, "x"); count != maxWebFetchContentRunes {
|
||||
t.Fatalf("captured content count = %d, want %d", count, maxWebFetchContentRunes)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWebToolsRejectWhenCloudDisabled(t *testing.T) {
|
||||
t.Setenv("OLLAMA_NO_CLOUD", "1")
|
||||
|
||||
for _, tt := range webToolCases {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
_, err := tt.tool.Execute(t.Context(), coreagent.ToolContext{}, tt.args)
|
||||
want := internalcloud.DisabledError(tt.operation)
|
||||
if err == nil || err.Error() != want {
|
||||
t.Fatalf("error = %v, want %q", err, want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -777,6 +777,18 @@ func (c *StreamConverter) Process(r api.ChatResponse) []StreamEvent {
|
||||
}
|
||||
|
||||
if r.Message.Thinking != "" && !c.thinkingDone {
|
||||
if c.textStarted {
|
||||
events = append(events, StreamEvent{
|
||||
Event: "content_block_stop",
|
||||
Data: ContentBlockStopEvent{
|
||||
Type: "content_block_stop",
|
||||
Index: c.contentIndex,
|
||||
},
|
||||
})
|
||||
c.contentIndex++
|
||||
c.textStarted = false
|
||||
}
|
||||
|
||||
if !c.thinkingStarted {
|
||||
c.thinkingStarted = true
|
||||
events = append(events, StreamEvent{
|
||||
|
||||
@@ -3,6 +3,7 @@ package anthropic
|
||||
import (
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
@@ -1140,6 +1141,56 @@ func TestStreamConverter_ThinkingDirectlyFollowedByToolCall(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestStreamConverter_TextBeforeThinking(t *testing.T) {
|
||||
conv := NewStreamConverter("msg_123", "test-model", 0)
|
||||
|
||||
responses := []api.ChatResponse{
|
||||
{Message: api.Message{Role: "assistant", Content: "---\n"}},
|
||||
{Message: api.Message{Role: "assistant", Thinking: "Let me think."}},
|
||||
{
|
||||
Message: api.Message{Role: "assistant", Content: "The answer."},
|
||||
Done: true,
|
||||
DoneReason: "stop",
|
||||
Metrics: api.Metrics{PromptEvalCount: 10, EvalCount: 5},
|
||||
},
|
||||
}
|
||||
|
||||
var got []string
|
||||
for _, response := range responses {
|
||||
for _, event := range conv.Process(response) {
|
||||
switch data := event.Data.(type) {
|
||||
case ContentBlockStartEvent:
|
||||
got = append(got, fmt.Sprintf("%s:%s:%d", event.Event, data.ContentBlock.Type, data.Index))
|
||||
case ContentBlockDeltaEvent:
|
||||
got = append(got, fmt.Sprintf("%s:%s:%d", event.Event, data.Delta.Type, data.Index))
|
||||
case ContentBlockStopEvent:
|
||||
got = append(got, fmt.Sprintf("%s:%d", event.Event, data.Index))
|
||||
default:
|
||||
got = append(got, event.Event)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
want := []string{
|
||||
"message_start",
|
||||
"content_block_start:text:0",
|
||||
"content_block_delta:text_delta:0",
|
||||
"content_block_stop:0",
|
||||
"content_block_start:thinking:1",
|
||||
"content_block_delta:thinking_delta:1",
|
||||
"content_block_stop:1",
|
||||
"content_block_start:text:2",
|
||||
"content_block_delta:text_delta:2",
|
||||
"content_block_stop:2",
|
||||
"message_delta",
|
||||
"message_stop",
|
||||
}
|
||||
|
||||
if diff := cmp.Diff(want, got); diff != "" {
|
||||
t.Fatalf("unexpected stream events (-want +got):\n%s", diff)
|
||||
}
|
||||
}
|
||||
|
||||
func TestStreamConverter_ToolCallWithUnmarshalableArgs(t *testing.T) {
|
||||
// Test that unmarshalable arguments (like channels) are handled gracefully
|
||||
// and don't cause a panic or corrupt stream
|
||||
|
||||
@@ -473,6 +473,26 @@ func (c *Client) CloudStatusExperimental(ctx context.Context) (*StatusResponse,
|
||||
return &status, nil
|
||||
}
|
||||
|
||||
// WebSearchExperimental searches the web through the local server's
|
||||
// experimental web search endpoint.
|
||||
func (c *Client) WebSearchExperimental(ctx context.Context, req *WebSearchRequest) (*WebSearchResponse, error) {
|
||||
var resp WebSearchResponse
|
||||
if err := c.do(ctx, http.MethodPost, "/api/experimental/web_search", req, &resp); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &resp, nil
|
||||
}
|
||||
|
||||
// WebFetchExperimental fetches web page content through the local server's
|
||||
// experimental web fetch endpoint.
|
||||
func (c *Client) WebFetchExperimental(ctx context.Context, req *WebFetchRequest) (*WebFetchResponse, error) {
|
||||
var resp WebFetchResponse
|
||||
if err := c.do(ctx, http.MethodPost, "/api/experimental/web_fetch", req, &resp); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &resp, nil
|
||||
}
|
||||
|
||||
// Signout will signout a client for a local ollama server.
|
||||
func (c *Client) Signout(ctx context.Context) error {
|
||||
return c.do(ctx, http.MethodPost, "/api/signout", nil, nil)
|
||||
@@ -490,3 +510,13 @@ func (c *Client) Whoami(ctx context.Context) (*UserResponse, error) {
|
||||
}
|
||||
return &resp, nil
|
||||
}
|
||||
|
||||
// Usage returns the authenticated user's recent activity and included-usage
|
||||
// limits.
|
||||
func (c *Client) Usage(ctx context.Context) (*UsageResponse, error) {
|
||||
var resp UsageResponse
|
||||
if err := c.do(ctx, http.MethodGet, "/api/usage", nil, &resp); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &resp, nil
|
||||
}
|
||||
@@ -51,6 +51,32 @@ func TestClientFromEnvironment(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestClientUsage(t *testing.T) {
|
||||
ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodGet || r.URL.Path != "/api/usage" {
|
||||
t.Fatalf("request = %s %s, want GET /api/usage", r.Method, r.URL.Path)
|
||||
}
|
||||
fmt.Fprint(w, `{"activity":{"cost":"0.00709","period":{"type":"last_4_weeks","starting_at":"2026-06-29T00:00:00Z","ending_at":"2026-07-27T00:00:00Z"},"models":[{"name":"qwen3-coder:480b","request_count":1,"cost":"0.00709"}]},"limits":{"session":{"usage":0.006,"models":[]},"weekly":{"usage":0,"models":[]}}}`)
|
||||
}))
|
||||
defer ts.Close()
|
||||
|
||||
base, err := url.Parse(ts.URL)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
got, err := NewClient(base, ts.Client()).Usage(t.Context())
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if got.Activity.Cost != "0.00709" {
|
||||
t.Errorf("activity cost = %q, want 0.00709", got.Activity.Cost)
|
||||
}
|
||||
if len(got.Activity.Models) != 1 || got.Activity.Models[0].Name != "qwen3-coder:480b" {
|
||||
t.Errorf("activity models = %#v, want qwen3-coder:480b", got.Activity.Models)
|
||||
}
|
||||
}
|
||||
|
||||
// testError represents an internal error type with status code and message
|
||||
// this is used since the error response from the server is not a standard error struct
|
||||
type testError struct {
|
||||
@@ -351,6 +377,82 @@ func TestClientDo(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestClientWebSearchExperimentalUsesLocalRoute(t *testing.T) {
|
||||
var gotPath string
|
||||
var gotMethod string
|
||||
var gotRequest WebSearchRequest
|
||||
ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
gotPath = r.URL.Path
|
||||
gotMethod = r.Method
|
||||
if err := json.NewDecoder(r.Body).Decode(&gotRequest); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := json.NewEncoder(w).Encode(WebSearchResponse{
|
||||
Results: []WebSearchResult{{Title: "Ollama", URL: "https://ollama.com", Content: "models"}},
|
||||
}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}))
|
||||
defer ts.Close()
|
||||
|
||||
client := NewClient(&url.URL{Scheme: "http", Host: ts.Listener.Addr().String()}, http.DefaultClient)
|
||||
resp, err := client.WebSearchExperimental(t.Context(), &WebSearchRequest{Query: "ollama", MaxResults: 3})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if gotMethod != http.MethodPost {
|
||||
t.Fatalf("method = %q, want POST", gotMethod)
|
||||
}
|
||||
if gotPath != "/api/experimental/web_search" {
|
||||
t.Fatalf("path = %q, want /api/experimental/web_search", gotPath)
|
||||
}
|
||||
if gotRequest.Query != "ollama" || gotRequest.MaxResults != 3 {
|
||||
t.Fatalf("request = %#v", gotRequest)
|
||||
}
|
||||
if len(resp.Results) != 1 || resp.Results[0].Title != "Ollama" {
|
||||
t.Fatalf("response = %#v", resp)
|
||||
}
|
||||
}
|
||||
|
||||
func TestClientWebFetchExperimentalUsesLocalRoute(t *testing.T) {
|
||||
var gotPath string
|
||||
var gotMethod string
|
||||
var gotRequest WebFetchRequest
|
||||
ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
gotPath = r.URL.Path
|
||||
gotMethod = r.Method
|
||||
if err := json.NewDecoder(r.Body).Decode(&gotRequest); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := json.NewEncoder(w).Encode(WebFetchResponse{
|
||||
Title: "Ollama",
|
||||
Content: "models",
|
||||
Links: []string{"https://ollama.com/library"},
|
||||
}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}))
|
||||
defer ts.Close()
|
||||
|
||||
client := NewClient(&url.URL{Scheme: "http", Host: ts.Listener.Addr().String()}, http.DefaultClient)
|
||||
resp, err := client.WebFetchExperimental(t.Context(), &WebFetchRequest{URL: "https://ollama.com"})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if gotMethod != http.MethodPost {
|
||||
t.Fatalf("method = %q, want POST", gotMethod)
|
||||
}
|
||||
if gotPath != "/api/experimental/web_fetch" {
|
||||
t.Fatalf("path = %q, want /api/experimental/web_fetch", gotPath)
|
||||
}
|
||||
if gotRequest.URL != "https://ollama.com" {
|
||||
t.Fatalf("request = %#v", gotRequest)
|
||||
}
|
||||
if resp.Title != "Ollama" || resp.Content != "models" {
|
||||
t.Fatalf("response = %#v", resp)
|
||||
}
|
||||
}
|
||||
|
||||
type roundTripFunc func(*http.Request) (*http.Response, error)
|
||||
|
||||
func (f roundTripFunc) RoundTrip(req *http.Request) (*http.Response, error) {
|
||||
|
||||
@@ -868,6 +868,36 @@ type StatusResponse struct {
|
||||
Cloud CloudStatus `json:"cloud"`
|
||||
}
|
||||
|
||||
// WebSearchRequest is the request for [Client.WebSearchExperimental].
|
||||
type WebSearchRequest struct {
|
||||
Query string `json:"query"`
|
||||
MaxResults int `json:"max_results,omitempty"`
|
||||
}
|
||||
|
||||
// WebSearchResult is a single result from [Client.WebSearchExperimental].
|
||||
type WebSearchResult struct {
|
||||
Title string `json:"title"`
|
||||
URL string `json:"url"`
|
||||
Content string `json:"content"`
|
||||
}
|
||||
|
||||
// WebSearchResponse is the response from [Client.WebSearchExperimental].
|
||||
type WebSearchResponse struct {
|
||||
Results []WebSearchResult `json:"results"`
|
||||
}
|
||||
|
||||
// WebFetchRequest is the request for [Client.WebFetchExperimental].
|
||||
type WebFetchRequest struct {
|
||||
URL string `json:"url"`
|
||||
}
|
||||
|
||||
// WebFetchResponse is the response from [Client.WebFetchExperimental].
|
||||
type WebFetchResponse struct {
|
||||
Title string `json:"title"`
|
||||
Content string `json:"content"`
|
||||
Links []string `json:"links,omitempty"`
|
||||
}
|
||||
|
||||
// GenerateResponse is the response passed into [GenerateResponseFunc].
|
||||
type GenerateResponse struct {
|
||||
// Model is the model name that generated the response.
|
||||
@@ -948,6 +978,45 @@ type UserResponse struct {
|
||||
Plan string `json:"plan,omitempty"`
|
||||
}
|
||||
|
||||
// UsageResponse reports recent activity and included-usage limits.
|
||||
type UsageResponse struct {
|
||||
Activity UsageActivity `json:"activity"`
|
||||
Limits UsageLimits `json:"limits"`
|
||||
}
|
||||
|
||||
// UsageActivity reports usage activity over a period.
|
||||
type UsageActivity struct {
|
||||
Cost string `json:"cost"`
|
||||
Period UsagePeriod `json:"period"`
|
||||
Models []UsageModel `json:"models"`
|
||||
}
|
||||
|
||||
// UsagePeriod describes the time window the usage covers.
|
||||
type UsagePeriod struct {
|
||||
Type string `json:"type"`
|
||||
StartingAt time.Time `json:"starting_at"`
|
||||
EndingAt time.Time `json:"ending_at"`
|
||||
}
|
||||
|
||||
// UsageLimits reports included usage for the current session and week.
|
||||
type UsageLimits struct {
|
||||
Session UsageLimit `json:"session"`
|
||||
Weekly UsageLimit `json:"weekly"`
|
||||
}
|
||||
|
||||
// UsageLimit reports the consumed fraction of an included-usage limit.
|
||||
type UsageLimit struct {
|
||||
Usage float64 `json:"usage"`
|
||||
Models []UsageModel `json:"models"`
|
||||
}
|
||||
|
||||
// UsageModel reports a model's activity.
|
||||
type UsageModel struct {
|
||||
Name string `json:"name"`
|
||||
RequestCount int `json:"request_count"`
|
||||
Cost string `json:"cost,omitempty"`
|
||||
}
|
||||
|
||||
// Tensor describes the metadata for a given tensor.
|
||||
type Tensor struct {
|
||||
Name string `json:"name"`
|
||||
|
||||
@@ -22,10 +22,10 @@ const LAUNCH_COMMANDS: LaunchCommand[] = [
|
||||
iconClassName: "h-7 w-7",
|
||||
},
|
||||
{
|
||||
id: "codex-app",
|
||||
name: "Codex App",
|
||||
command: "ollama launch codex-app",
|
||||
description: "An AI agent you can delegate real work to, by OpenAI",
|
||||
id: "chatgpt",
|
||||
name: "ChatGPT",
|
||||
command: "ollama launch chatgpt",
|
||||
description: "Complete work with ChatGPT",
|
||||
icon: "/launch-icons/codex-app.png",
|
||||
iconClassName: "h-full w-full",
|
||||
},
|
||||
|
||||
+96
-6
@@ -59,7 +59,7 @@ function(ollama_macos_major_version output)
|
||||
RESULT_VARIABLE _macos_result
|
||||
ERROR_QUIET)
|
||||
if(_macos_result EQUAL 0)
|
||||
string(REGEX MATCH "^[0-9]+" _macos_major "${_macos_version}")
|
||||
string(REGEX MATCH "^[0-9]+(\\.[0-9]+)?" _macos_major "${_macos_version}")
|
||||
endif()
|
||||
set(${output} "${_macos_major}" PARENT_SCOPE)
|
||||
endfunction()
|
||||
@@ -72,7 +72,7 @@ function(ollama_macos_sdk_major_version output)
|
||||
RESULT_VARIABLE _sdk_result
|
||||
ERROR_QUIET)
|
||||
if(_sdk_result EQUAL 0)
|
||||
string(REGEX MATCH "^[0-9]+" _sdk_major "${_sdk_version}")
|
||||
string(REGEX MATCH "^[0-9]+(\\.[0-9]+)?" _sdk_major "${_sdk_version}")
|
||||
endif()
|
||||
set(${output} "${_sdk_major}" PARENT_SCOPE)
|
||||
endfunction()
|
||||
@@ -83,7 +83,9 @@ function(ollama_default_mlx_backends output)
|
||||
ollama_check_metal_toolchain(_metal_version)
|
||||
ollama_macos_major_version(_macos_major)
|
||||
ollama_macos_sdk_major_version(_sdk_major)
|
||||
if(_macos_major AND _sdk_major AND _macos_major GREATER_EQUAL 26 AND _sdk_major GREATER_EQUAL 26)
|
||||
if(_macos_major AND _sdk_major
|
||||
AND _macos_major VERSION_GREATER_EQUAL 26.2
|
||||
AND _sdk_major VERSION_GREATER_EQUAL 26.2)
|
||||
set(_backends "metal_v4")
|
||||
else()
|
||||
set(_backends "metal_v3")
|
||||
@@ -192,8 +194,16 @@ if(OLLAMA_MLX_BACKENDS)
|
||||
add_custom_target(ollama-mlx-sources DEPENDS ${_mlx_source_targets})
|
||||
endif()
|
||||
|
||||
set(OLLAMA_BUILD_PARALLEL "" CACHE STRING
|
||||
"Number of parallel jobs for nested native builds (empty = use generator default)")
|
||||
|
||||
set(_native_parallel_args --parallel)
|
||||
if(NOT OLLAMA_BUILD_PARALLEL STREQUAL "")
|
||||
list(APPEND _native_parallel_args ${OLLAMA_BUILD_PARALLEL})
|
||||
endif()
|
||||
|
||||
set(OLLAMA_NATIVE_BUILD_TOOL_COMMAND
|
||||
${CMAKE_COMMAND} --build <BINARY_DIR>)
|
||||
${CMAKE_COMMAND} --build <BINARY_DIR> ${_native_parallel_args})
|
||||
set(OLLAMA_NATIVE_BUILD_TARGET_ARG --target)
|
||||
if(CMAKE_GENERATOR MATCHES "Makefiles")
|
||||
set(OLLAMA_NATIVE_BUILD_TOOL_COMMAND
|
||||
@@ -236,6 +246,67 @@ function(ollama_cache_arg_is_set name output)
|
||||
endif()
|
||||
endfunction()
|
||||
|
||||
function(ollama_backend_cuda_major backend output)
|
||||
if("${backend}" MATCHES "^cuda_v([0-9]+)$")
|
||||
set(${output} "${CMAKE_MATCH_1}" PARENT_SCOPE)
|
||||
else()
|
||||
set(${output} "" PARENT_SCOPE)
|
||||
endif()
|
||||
endfunction()
|
||||
|
||||
function(ollama_find_windows_cuda_root major output)
|
||||
if(NOT WIN32 OR "${major}" STREQUAL "")
|
||||
set(${output} "" PARENT_SCOPE)
|
||||
return()
|
||||
endif()
|
||||
|
||||
execute_process(
|
||||
COMMAND ${CMAKE_COMMAND} -E environment
|
||||
OUTPUT_VARIABLE _environment)
|
||||
string(REPLACE "\r\n" "\n" _environment "${_environment}")
|
||||
string(REPLACE "\r" "\n" _environment "${_environment}")
|
||||
string(REGEX MATCHALL "CUDA_PATH_V${major}_[0-9]+=[^\n]*" _matches "${_environment}")
|
||||
|
||||
set(_best_minor -1)
|
||||
set(_best_root "")
|
||||
foreach(_entry IN LISTS _matches)
|
||||
if(_entry MATCHES "^CUDA_PATH_V${major}_([0-9]+)=(.*)$")
|
||||
set(_minor "${CMAKE_MATCH_1}")
|
||||
set(_root "${CMAKE_MATCH_2}")
|
||||
if(_minor GREATER _best_minor)
|
||||
set(_best_minor ${_minor})
|
||||
set(_best_root "${_root}")
|
||||
endif()
|
||||
endif()
|
||||
endforeach()
|
||||
|
||||
if(_best_root STREQUAL "" AND DEFINED ENV{CUDA_PATH})
|
||||
set(_cuda_path "$ENV{CUDA_PATH}")
|
||||
if(EXISTS "${_cuda_path}/version.json")
|
||||
file(READ "${_cuda_path}/version.json" _version_json)
|
||||
if(_version_json MATCHES "\"cuda\"[ \t\r\n]*:[ \t\r\n]*\"${major}\\.")
|
||||
set(_best_root "${_cuda_path}")
|
||||
endif()
|
||||
endif()
|
||||
endif()
|
||||
|
||||
set(${output} "${_best_root}" PARENT_SCOPE)
|
||||
endfunction()
|
||||
|
||||
function(ollama_append_cuda_toolkit_args output backend)
|
||||
# If CUDAToolkit_ROOT is already explicitly set, just forward it.
|
||||
ollama_append_cache_arg_if_set(${output} CUDAToolkit_ROOT)
|
||||
if(NOT DEFINED CUDAToolkit_ROOT OR "${CUDAToolkit_ROOT}" STREQUAL "")
|
||||
# Auto-discover CUDA toolkit for the requested backend version on Windows.
|
||||
ollama_backend_cuda_major("${backend}" _cuda_major)
|
||||
ollama_find_windows_cuda_root("${_cuda_major}" _cuda_root)
|
||||
if(NOT "${_cuda_root}" STREQUAL "")
|
||||
ollama_escape_cmake_list("${_cuda_root}" _value)
|
||||
set(${output} ${${output}} "-DCUDAToolkit_ROOT=${_value}" PARENT_SCOPE)
|
||||
endif()
|
||||
endif()
|
||||
endfunction()
|
||||
|
||||
function(ollama_llama_cuda_preset backend output)
|
||||
ollama_cache_arg_is_set(CMAKE_CUDA_ARCHITECTURES _has_cuda_arch)
|
||||
if(_has_cuda_arch)
|
||||
@@ -327,12 +398,28 @@ function(ollama_add_llama_server_build name)
|
||||
-DCMAKE_OSX_DEPLOYMENT_TARGET=${CMAKE_OSX_DEPLOYMENT_TARGET})
|
||||
endif()
|
||||
endif()
|
||||
# Visual Studio requires -T toolset override to select the correct CUDA toolkit.
|
||||
# MSBuild's CUDA integration ignores -DCUDAToolkit_ROOT for nvcc selection.
|
||||
# Prefer user-specified CUDAToolkit_ROOT before falling back to auto-discovery.
|
||||
set(_generator_args)
|
||||
if(WIN32 AND CMAKE_GENERATOR MATCHES "Visual Studio")
|
||||
set(_cuda_root "${CUDAToolkit_ROOT}")
|
||||
if("${_cuda_root}" STREQUAL "")
|
||||
ollama_backend_cuda_major("${name}" _cuda_major)
|
||||
ollama_find_windows_cuda_root("${_cuda_major}" _cuda_root)
|
||||
endif()
|
||||
if(NOT "${_cuda_root}" STREQUAL "")
|
||||
list(APPEND _generator_args -T cuda=${_cuda_root})
|
||||
endif()
|
||||
endif()
|
||||
set(_configure_command ${CMAKE_COMMAND}
|
||||
${_generator_args}
|
||||
-S ${CMAKE_SOURCE_DIR}/llama/server
|
||||
-B <BINARY_DIR>
|
||||
${_cmake_args})
|
||||
if(ARG_PRESET)
|
||||
set(_configure_command ${CMAKE_COMMAND}
|
||||
${_generator_args}
|
||||
-S ${CMAKE_SOURCE_DIR}/llama/server
|
||||
--preset ${ARG_PRESET}
|
||||
-B <BINARY_DIR>
|
||||
@@ -544,6 +631,7 @@ if(OLLAMA_HAVE_LLAMA_SERVER)
|
||||
set(_cuda_args)
|
||||
ollama_append_cache_arg_if_set(_cuda_args CMAKE_CUDA_ARCHITECTURES)
|
||||
ollama_append_cache_arg_if_set(_cuda_args CMAKE_CUDA_FLAGS)
|
||||
ollama_append_cuda_toolkit_args(_cuda_args ${_backend})
|
||||
ollama_add_llama_server_build(${_backend}
|
||||
PRESET ${_cuda_preset}
|
||||
RUNNER_DIR ${_backend}
|
||||
@@ -555,6 +643,7 @@ if(OLLAMA_HAVE_LLAMA_SERVER)
|
||||
set(_cuda_args)
|
||||
ollama_append_cache_arg_if_set(_cuda_args CMAKE_CUDA_ARCHITECTURES)
|
||||
ollama_append_cache_arg_if_set(_cuda_args CMAKE_CUDA_FLAGS)
|
||||
ollama_append_cuda_toolkit_args(_cuda_args ${_backend})
|
||||
ollama_add_llama_server_build(${_backend}
|
||||
PRESET ${_cuda_preset}
|
||||
RUNNER_DIR ${_backend}
|
||||
@@ -664,14 +753,15 @@ foreach(_backend IN LISTS OLLAMA_MLX_BACKENDS)
|
||||
endif()
|
||||
ollama_check_metal_toolchain(_metal_version)
|
||||
ollama_macos_sdk_major_version(_ollama_mlx_sdk_major)
|
||||
if(_ollama_mlx_sdk_major AND _ollama_mlx_sdk_major GREATER_EQUAL 26)
|
||||
if(_ollama_mlx_sdk_major
|
||||
AND _ollama_mlx_sdk_major VERSION_GREATER_EQUAL 26.2)
|
||||
ollama_add_mlx_build(metal_v4
|
||||
PRESET mlx_metal_v4
|
||||
RUNNER_DIR mlx_metal_v4)
|
||||
list(APPEND _mlx_targets ollama-mlx-metal_v4)
|
||||
else()
|
||||
message(FATAL_ERROR
|
||||
"OLLAMA_MLX_BACKENDS=metal_v4 requires the macOS 26 SDK. "
|
||||
"OLLAMA_MLX_BACKENDS=metal_v4 requires the macOS 26.2 SDK. "
|
||||
"Install a newer Xcode or use OLLAMA_MLX_BACKENDS=metal_v3.")
|
||||
endif()
|
||||
else()
|
||||
|
||||
+114
-31
@@ -102,7 +102,8 @@ install(RUNTIME_DEPENDENCY_SET mlx_runtime_deps
|
||||
LIBRARY DESTINATION ${OLLAMA_INSTALL_DIR} COMPONENT MLX_VENDOR
|
||||
)
|
||||
|
||||
if(TARGET jaccl)
|
||||
get_target_property(_MLX_LINK_LIBRARIES mlx LINK_LIBRARIES)
|
||||
if(TARGET jaccl AND "jaccl" IN_LIST _MLX_LINK_LIBRARIES)
|
||||
install(TARGETS jaccl
|
||||
RUNTIME DESTINATION ${OLLAMA_INSTALL_DIR} COMPONENT MLX
|
||||
LIBRARY DESTINATION ${OLLAMA_INSTALL_DIR} COMPONENT MLX
|
||||
@@ -123,29 +124,53 @@ endif()
|
||||
# --component MLX. Headers are installed alongside libmlx in OLLAMA_INSTALL_DIR.
|
||||
#
|
||||
# Layout:
|
||||
# ${OLLAMA_INSTALL_DIR}/include/cccl/{cuda,nv}/ - CCCL headers
|
||||
# ${OLLAMA_INSTALL_DIR}/include/*.h - CUDA toolkit headers
|
||||
# ${OLLAMA_INSTALL_DIR}/include/cccl/ - CCCL headers
|
||||
# ${OLLAMA_INSTALL_DIR}/include/{cute,cutlass}/ - CUTLASS/CUTE headers
|
||||
# ${OLLAMA_INSTALL_DIR}/include/ - CUDA runtime/core headers
|
||||
#
|
||||
# MLX's jit_module.cpp resolves CCCL via
|
||||
# current_binary_dir()[.parent_path()] / "include" / "cccl"
|
||||
# On Linux, MLX's jit_module.cpp resolves CCCL via
|
||||
# current_binary_dir().parent_path() / "include" / "cccl", so we create a
|
||||
# symlink from lib/ollama/include -> ${OLLAMA_RUNNER_DIR}/include.
|
||||
# MLX's jit_module.cpp resolves JIT support headers from the backend-local
|
||||
# include directory. On Linux it also probes current_binary_dir().parent_path()
|
||||
# / "include", so we create a symlink from lib/ollama/include to the backend
|
||||
# include directory for archive packaging.
|
||||
# This will need refinement if we add multiple CUDA versions for MLX in the future.
|
||||
# CUDA runtime headers are found via CUDA_PATH env var (set by mlxrunner).
|
||||
if(EXISTS ${CMAKE_BINARY_DIR}/_deps/cccl-src/include/cuda)
|
||||
install(DIRECTORY ${CMAKE_BINARY_DIR}/_deps/cccl-src/include/cuda
|
||||
DESTINATION ${OLLAMA_INSTALL_DIR}/include/cccl
|
||||
set(_mlx_jit_cccl_include_dir "")
|
||||
if(CUDAToolkit_FOUND)
|
||||
foreach(_dir ${CUDAToolkit_INCLUDE_DIRS})
|
||||
if(EXISTS "${_dir}/cccl/cuda/std")
|
||||
set(_mlx_jit_cccl_include_dir "${_dir}/cccl")
|
||||
break()
|
||||
endif()
|
||||
endforeach()
|
||||
endif()
|
||||
if(NOT _mlx_jit_cccl_include_dir AND EXISTS ${CMAKE_BINARY_DIR}/_deps/cccl-src/include/cuda)
|
||||
set(_mlx_jit_cccl_include_dir "${CMAKE_BINARY_DIR}/_deps/cccl-src/include")
|
||||
endif()
|
||||
if(_mlx_jit_cccl_include_dir)
|
||||
foreach(_cccl_dir cuda nv cub thrust)
|
||||
if(EXISTS "${_mlx_jit_cccl_include_dir}/${_cccl_dir}")
|
||||
install(DIRECTORY "${_mlx_jit_cccl_include_dir}/${_cccl_dir}"
|
||||
DESTINATION ${OLLAMA_INSTALL_DIR}/include/cccl
|
||||
COMPONENT MLX)
|
||||
endif()
|
||||
endforeach()
|
||||
endif()
|
||||
if(EXISTS ${CMAKE_BINARY_DIR}/_deps/cutlass-src/include/cute)
|
||||
install(DIRECTORY ${CMAKE_BINARY_DIR}/_deps/cutlass-src/include/cute
|
||||
DESTINATION ${OLLAMA_INSTALL_DIR}/include
|
||||
COMPONENT MLX)
|
||||
install(DIRECTORY ${CMAKE_BINARY_DIR}/_deps/cccl-src/include/nv
|
||||
DESTINATION ${OLLAMA_INSTALL_DIR}/include/cccl
|
||||
install(DIRECTORY ${CMAKE_BINARY_DIR}/_deps/cutlass-src/include/cutlass
|
||||
DESTINATION ${OLLAMA_INSTALL_DIR}/include
|
||||
COMPONENT MLX)
|
||||
endif()
|
||||
|
||||
# Install minimal CUDA toolkit headers needed by MLX JIT kernels.
|
||||
# These are the transitive closure of includes from mlx/backend/cuda/device/*.cuh.
|
||||
# Install CUDA runtime/core headers needed by MLX JIT kernels.
|
||||
# NVIDIA's NVRTC bundled-header model is CUDA Runtime + CCCL, not the entire
|
||||
# toolkit include tree. Keep CCCL coherent above, include CUTLASS/CUTE above,
|
||||
# and avoid shipping unrelated SDK headers such as NPP, CUPTI, cuRAND, NVML,
|
||||
# cuBLAS, cuSPARSE, and cuSOLVER.
|
||||
# The Go mlxrunner sets CUDA_PATH to OLLAMA_INSTALL_DIR so MLX finds them at
|
||||
# $CUDA_PATH/include/*.h via NVRTC --include-path.
|
||||
# $CUDA_PATH/include via NVRTC --include-path.
|
||||
if(CUDAToolkit_FOUND)
|
||||
# CUDAToolkit_INCLUDE_DIRS may be a semicolon-separated list
|
||||
# (e.g. ".../include;.../include/cccl"). Find the entry that
|
||||
@@ -161,39 +186,97 @@ if(CUDAToolkit_FOUND)
|
||||
message(WARNING "Could not find cuda_runtime_api.h in CUDAToolkit_INCLUDE_DIRS: ${CUDAToolkit_INCLUDE_DIRS}")
|
||||
else()
|
||||
set(_dst "${OLLAMA_INSTALL_DIR}/include")
|
||||
set(_MLX_JIT_CUDA_HEADERS
|
||||
|
||||
set(_mlx_jit_cuda_headers
|
||||
builtin_types.h
|
||||
channel_descriptor.h
|
||||
common_functions.h
|
||||
cooperative_groups.h
|
||||
cuComplex.h
|
||||
cuda.h
|
||||
cudaTypedefs.h
|
||||
cuda_awbarrier.h
|
||||
cuda_awbarrier_helpers.h
|
||||
cuda_awbarrier_primitives.h
|
||||
cuda_bf16.h
|
||||
cuda_bf16.hpp
|
||||
cuda_device_runtime_api.h
|
||||
cuda_fp16.h
|
||||
cuda_fp16.hpp
|
||||
cuda_fp4.h
|
||||
cuda_fp4.hpp
|
||||
cuda_fp6.h
|
||||
cuda_fp6.hpp
|
||||
cuda_fp8.h
|
||||
cuda_fp8.hpp
|
||||
cuda_fp16.h
|
||||
cuda_fp16.hpp
|
||||
cuda_occupancy.h
|
||||
cuda_pipeline.h
|
||||
cuda_pipeline_helpers.h
|
||||
cuda_pipeline_primitives.h
|
||||
cuda_runtime.h
|
||||
cuda_runtime_api.h
|
||||
cuda_stdint.h
|
||||
cudart_platform.h
|
||||
device_atomic_functions.h
|
||||
device_atomic_functions.hpp
|
||||
device_double_functions.h
|
||||
device_functions.h
|
||||
device_launch_parameters.h
|
||||
device_types.h
|
||||
driver_functions.h
|
||||
driver_types.h
|
||||
fatbinary_section.h
|
||||
host_config.h
|
||||
host_defines.h
|
||||
library_types.h
|
||||
math_constants.h
|
||||
math_functions.h
|
||||
mma.h
|
||||
nvrtc_device_runtime.h
|
||||
sm_20_atomic_functions.h
|
||||
sm_20_atomic_functions.hpp
|
||||
sm_20_intrinsics.h
|
||||
sm_20_intrinsics.hpp
|
||||
sm_30_intrinsics.h
|
||||
sm_30_intrinsics.hpp
|
||||
sm_32_atomic_functions.h
|
||||
sm_32_atomic_functions.hpp
|
||||
sm_32_intrinsics.h
|
||||
sm_32_intrinsics.hpp
|
||||
sm_35_atomic_functions.h
|
||||
sm_35_intrinsics.h
|
||||
sm_60_atomic_functions.h
|
||||
sm_60_atomic_functions.hpp
|
||||
sm_61_intrinsics.h
|
||||
sm_61_intrinsics.hpp
|
||||
surface_indirect_functions.h
|
||||
surface_types.h
|
||||
target
|
||||
texture_indirect_functions.h
|
||||
texture_types.h
|
||||
vector_functions.h
|
||||
vector_functions.hpp
|
||||
vector_types.h
|
||||
)
|
||||
foreach(_hdr ${_MLX_JIT_CUDA_HEADERS})
|
||||
install(FILES "${_cuda_inc}/${_hdr}"
|
||||
vector_types.h)
|
||||
set(_mlx_jit_cuda_header_paths "")
|
||||
foreach(_header IN LISTS _mlx_jit_cuda_headers)
|
||||
if(EXISTS "${_cuda_inc}/${_header}")
|
||||
list(APPEND _mlx_jit_cuda_header_paths "${_cuda_inc}/${_header}")
|
||||
endif()
|
||||
endforeach()
|
||||
if(_mlx_jit_cuda_header_paths)
|
||||
install(FILES ${_mlx_jit_cuda_header_paths}
|
||||
DESTINATION ${_dst}
|
||||
COMPONENT MLX)
|
||||
endif()
|
||||
|
||||
foreach(_runtime_dir cooperative_groups crt)
|
||||
if(EXISTS "${_cuda_inc}/${_runtime_dir}")
|
||||
install(DIRECTORY "${_cuda_inc}/${_runtime_dir}"
|
||||
DESTINATION ${_dst}
|
||||
COMPONENT MLX)
|
||||
endif()
|
||||
endforeach()
|
||||
# Subdirectory headers.
|
||||
install(DIRECTORY "${_cuda_inc}/cooperative_groups"
|
||||
DESTINATION ${_dst}
|
||||
COMPONENT MLX
|
||||
FILES_MATCHING PATTERN "*.h")
|
||||
install(FILES "${_cuda_inc}/crt/host_defines.h"
|
||||
DESTINATION "${_dst}/crt"
|
||||
COMPONENT MLX)
|
||||
|
||||
if(NOT WIN32 AND NOT APPLE)
|
||||
install(CODE "
|
||||
set(_link \"${CMAKE_INSTALL_PREFIX}/${OLLAMA_LIB_DIR}/include\")
|
||||
|
||||
@@ -55,7 +55,7 @@
|
||||
"inherits": [ "default" ],
|
||||
"binaryDir": "${sourceDir}/../../build/metal-v4",
|
||||
"cacheVariables": {
|
||||
"CMAKE_OSX_DEPLOYMENT_TARGET": "26.0",
|
||||
"CMAKE_OSX_DEPLOYMENT_TARGET": "26.2",
|
||||
"OLLAMA_RUNNER_DIR": "mlx_metal_v4"
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,863 @@
|
||||
package cmd
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"os"
|
||||
"runtime"
|
||||
"slices"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/spf13/cobra"
|
||||
|
||||
coreagent "github.com/ollama/ollama/agent"
|
||||
agenttools "github.com/ollama/ollama/agent/tools"
|
||||
"github.com/ollama/ollama/api"
|
||||
"github.com/ollama/ollama/cmd/config"
|
||||
"github.com/ollama/ollama/cmd/launch"
|
||||
agentchat "github.com/ollama/ollama/cmd/tui/chat"
|
||||
"github.com/ollama/ollama/format"
|
||||
internalcloud "github.com/ollama/ollama/internal/cloud"
|
||||
"github.com/ollama/ollama/internal/modelref"
|
||||
"github.com/ollama/ollama/types/model"
|
||||
)
|
||||
|
||||
type agentTUIOptions struct {
|
||||
Model string
|
||||
OpenModelPicker bool
|
||||
System string
|
||||
Format string
|
||||
Options map[string]any
|
||||
Think *api.ThinkValue
|
||||
KeepAlive *api.Duration
|
||||
ContextWindowTokens int
|
||||
AllowAllTools bool
|
||||
ToolsDisabled bool
|
||||
MultiModal bool
|
||||
}
|
||||
|
||||
func registerAgentFlags(cmd *cobra.Command) {
|
||||
cmd.Flags().String("model", "", "Model to use")
|
||||
cmd.Flags().String("keepalive", "", "Duration to keep a model loaded (e.g. 5m)")
|
||||
cmd.Flags().String("format", "", "Response format (e.g. json)")
|
||||
cmd.Flags().String("think", "", "Enable thinking mode: true/false or high/medium/low for supported models")
|
||||
cmd.Flags().Lookup("think").NoOptDefVal = "true"
|
||||
cmd.Flags().Bool("auto-approve-tools", false, "Allow agent tools to run without prompting")
|
||||
cmd.Flags().Bool("yolo", false, "Alias for --auto-approve-tools")
|
||||
cmd.Flags().Bool("no-tools", false, "Disable agent tools")
|
||||
}
|
||||
|
||||
func AgentHandler(cmd *cobra.Command, _ []string) error {
|
||||
opts := agentTUIOptions{
|
||||
Model: strings.TrimSpace(config.LastModel()),
|
||||
Options: map[string]any{},
|
||||
}
|
||||
thinkExplicit, err := applyAgentFlags(cmd, &opts)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if strings.TrimSpace(opts.Model) == "" {
|
||||
opts.OpenModelPicker = true
|
||||
} else if cmd.Flags().Lookup("model") == nil || !cmd.Flags().Lookup("model").Changed {
|
||||
opts.OpenModelPicker = true
|
||||
}
|
||||
|
||||
client, err := api.ClientFromEnvironment()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if opts.OpenModelPicker {
|
||||
modelName, err := selectAgentModel(cmd.Context(), client, opts.Model)
|
||||
if errors.Is(err, launch.ErrCancelled) {
|
||||
return nil
|
||||
}
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
opts.Model = modelName
|
||||
opts.OpenModelPicker = false
|
||||
}
|
||||
|
||||
if strings.TrimSpace(opts.Model) != "" {
|
||||
info, err := prepareAgentModel(cmd, client, &opts, thinkExplicit)
|
||||
if err != nil {
|
||||
if handleCloudAuthorizationError(err) {
|
||||
return nil
|
||||
}
|
||||
return err
|
||||
}
|
||||
opts.System = info.System
|
||||
if err := saveLastAgentModel(opts.Model); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
if err := GenerateAgentTUI(cmd, client, opts); err != nil {
|
||||
if handleCloudAuthorizationError(err) {
|
||||
return nil
|
||||
}
|
||||
return fmt.Errorf("error running agent: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func applyAgentFlags(cmd *cobra.Command, opts *agentTUIOptions) (bool, error) {
|
||||
if flag := cmd.Flags().Lookup("model"); flag != nil && flag.Changed {
|
||||
modelName, err := cmd.Flags().GetString("model")
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
modelName = strings.TrimSpace(modelName)
|
||||
if modelName == "" {
|
||||
return false, errors.New("--model cannot be empty")
|
||||
}
|
||||
opts.Model = modelName
|
||||
opts.OpenModelPicker = false
|
||||
}
|
||||
|
||||
format, err := cmd.Flags().GetString("format")
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
opts.Format = format
|
||||
|
||||
thinkExplicit := false
|
||||
thinkFlag := cmd.Flags().Lookup("think")
|
||||
if thinkFlag != nil && thinkFlag.Changed {
|
||||
thinkExplicit = true
|
||||
thinkStr, err := cmd.Flags().GetString("think")
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
switch thinkStr {
|
||||
case "", "true":
|
||||
opts.Think = &api.ThinkValue{Value: true}
|
||||
case "false":
|
||||
opts.Think = &api.ThinkValue{Value: false}
|
||||
case "high", "medium", "low", "max":
|
||||
opts.Think = &api.ThinkValue{Value: thinkStr}
|
||||
default:
|
||||
return false, fmt.Errorf("invalid value for --think: %q (must be true, false, high, medium, low, or max)", thinkStr)
|
||||
}
|
||||
}
|
||||
|
||||
keepAlive, err := cmd.Flags().GetString("keepalive")
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
if keepAlive != "" {
|
||||
d, err := time.ParseDuration(keepAlive)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
opts.KeepAlive = &api.Duration{Duration: d}
|
||||
}
|
||||
|
||||
autoApprove, err := cmd.Flags().GetBool("auto-approve-tools")
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
yolo, err := cmd.Flags().GetBool("yolo")
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
opts.AllowAllTools = autoApprove || yolo
|
||||
toolsDisabled, err := cmd.Flags().GetBool("no-tools")
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
opts.ToolsDisabled = toolsDisabled
|
||||
return thinkExplicit, nil
|
||||
}
|
||||
|
||||
func saveLastAgentModel(model string) error {
|
||||
model = strings.TrimSpace(model)
|
||||
if model == "" {
|
||||
return nil
|
||||
}
|
||||
return config.SetLastModel(model)
|
||||
}
|
||||
|
||||
func prepareAgentModel(cmd *cobra.Command, client *api.Client, opts *agentTUIOptions, thinkExplicit bool) (*api.ShowResponse, error) {
|
||||
requestedCloud := modelref.HasExplicitCloudSource(opts.Model)
|
||||
info, err := func() (*api.ShowResponse, error) {
|
||||
info, err := client.Show(cmd.Context(), &api.ShowRequest{Model: opts.Model})
|
||||
var se api.StatusError
|
||||
if errors.As(err, &se) && se.StatusCode == http.StatusNotFound {
|
||||
if requestedCloud {
|
||||
return nil, err
|
||||
}
|
||||
if err := PullHandler(cmd, []string{opts.Model}); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return client.Show(cmd.Context(), &api.ShowRequest{Model: opts.Model})
|
||||
}
|
||||
return info, err
|
||||
}()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
ensureCloudStub(cmd.Context(), client, opts.Model)
|
||||
opts.Think, err = inferThinkingOption(&info.Capabilities, &runOptions{Model: opts.Model, Think: opts.Think}, thinkExplicit)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
opts.MultiModal = showResponseSupportsMultimodal(info)
|
||||
opts.ContextWindowTokens = showResponseContextWindow(info)
|
||||
return info, nil
|
||||
}
|
||||
|
||||
func GenerateAgentTUI(cmd *cobra.Command, client *api.Client, opts agentTUIOptions) error {
|
||||
cwd := agentWorkingDir()
|
||||
contextWindowForModel := func(ctx context.Context, model string, fallback int) int {
|
||||
return agentContextWindowForModel(ctx, client, model, fallback)
|
||||
}
|
||||
|
||||
skillCatalog, err := coreagent.LoadDefaultSkills(cwd)
|
||||
if err != nil {
|
||||
return fmt.Errorf("load agent skills: %w", err)
|
||||
}
|
||||
for _, diagnostic := range skillCatalog.Diagnostics() {
|
||||
fmt.Fprintf(os.Stderr, "\033[1mwarning:\033[0m ignored invalid agent skill: %v\n", diagnostic)
|
||||
}
|
||||
var registry *coreagent.Registry
|
||||
registryForModel := func(ctx context.Context, model string) *coreagent.Registry {
|
||||
return agentToolsRegistry(ctx, client, model, skillCatalog)
|
||||
}
|
||||
if opts.Model != "" {
|
||||
registry = agentToolsRegistry(cmd.Context(), client, opts.Model, skillCatalog)
|
||||
}
|
||||
systemPrompt := agentSystemPromptWithWorkingDir(opts.Model, opts.System, agentSkillSystemContext(skillCatalog, registry, opts.ToolsDisabled), cwd)
|
||||
|
||||
_, err = agentchat.Run(cmd.Context(), agentchat.Options{
|
||||
Model: opts.Model,
|
||||
Client: client,
|
||||
Tools: registry,
|
||||
ToolRegistryForModel: registryForModel,
|
||||
ToolsDisabled: opts.ToolsDisabled,
|
||||
MultiModalForModel: func(ctx context.Context, model string) bool {
|
||||
return agentModelSupportsMultimodal(ctx, client, model)
|
||||
},
|
||||
ModelOptions: func(ctx context.Context) ([]agentchat.ModelOption, error) {
|
||||
return agentModelOptions(ctx, client)
|
||||
},
|
||||
OnModelSelected: func(_ context.Context, model string) error {
|
||||
return config.SetLastModel(model)
|
||||
},
|
||||
SystemPromptForModel: func(ctx context.Context, model string, registry *coreagent.Registry, toolsDisabled bool) string {
|
||||
return agentSystemPromptWithWorkingDir(model, agentSystemFromShow(ctx, client, model), agentSkillSystemContext(skillCatalog, registry, toolsDisabled), cwd)
|
||||
},
|
||||
Skills: skillCatalog,
|
||||
SystemPrompt: systemPrompt,
|
||||
WorkingDir: cwd,
|
||||
Format: opts.Format,
|
||||
Options: opts.Options,
|
||||
Think: opts.Think,
|
||||
KeepAlive: opts.KeepAlive,
|
||||
MultiModal: opts.MultiModal,
|
||||
AllowAllTools: opts.AllowAllTools,
|
||||
ContextWindowTokens: opts.ContextWindowTokens,
|
||||
Compactor: &coreagent.SimpleCompactor{
|
||||
Client: client,
|
||||
Options: coreagent.CompactionOptions{ContextWindowTokens: opts.ContextWindowTokens},
|
||||
},
|
||||
ContextWindowTokensForModel: func(ctx context.Context, model string, fallback int) int {
|
||||
return contextWindowForModel(ctx, model, fallback)
|
||||
},
|
||||
PreloadModel: func(ctx context.Context, model string, think *api.ThinkValue) (int, error) {
|
||||
return preloadAgentModelIfLocal(ctx, client, opts, model, think)
|
||||
},
|
||||
CheckCloudModel: func(ctx context.Context, model, requiredPlan string) error {
|
||||
return ensureCloudModelAccess(ctx, client, model, requiredPlan)
|
||||
},
|
||||
OpenBrowser: launch.OpenBrowser,
|
||||
PollCloudAuth: func(ctx context.Context) (string, bool, error) {
|
||||
user, err := client.Whoami(ctx)
|
||||
if err != nil {
|
||||
return "", false, err
|
||||
}
|
||||
if user == nil || user.Name == "" {
|
||||
return "", false, nil
|
||||
}
|
||||
return user.Name, true, nil
|
||||
},
|
||||
})
|
||||
return err
|
||||
}
|
||||
|
||||
func agentSkillSystemContext(catalog *coreagent.SkillCatalog, registry *coreagent.Registry, toolsDisabled bool) string {
|
||||
if toolsDisabled || registry == nil {
|
||||
return ""
|
||||
}
|
||||
if _, ok := registry.Get("skill"); !ok {
|
||||
return ""
|
||||
}
|
||||
return catalog.SystemContext()
|
||||
}
|
||||
|
||||
func selectAgentModel(ctx context.Context, client *api.Client, current string) (string, error) {
|
||||
models, err := agentModelOptions(ctx, client)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if len(models) == 0 {
|
||||
return "", errors.New("no models available, run 'ollama pull <model>' first")
|
||||
}
|
||||
|
||||
items := agentSelectionItems(models)
|
||||
switch {
|
||||
case launch.DefaultSingleSelectorWithUpdates != nil:
|
||||
return launch.DefaultSingleSelectorWithUpdates("Select model to run:", items, current, nil)
|
||||
case launch.DefaultSingleSelector != nil:
|
||||
return launch.DefaultSingleSelector("Select model to run:", items, current)
|
||||
default:
|
||||
return "", errors.New("no selector configured")
|
||||
}
|
||||
}
|
||||
|
||||
func agentSelectionItems(models []agentchat.ModelOption) []launch.SelectionItem {
|
||||
items := make([]launch.SelectionItem, 0, len(models))
|
||||
for _, model := range models {
|
||||
items = append(items, launch.SelectionItem{
|
||||
Name: model.Name,
|
||||
Description: agentSelectionDescription(model),
|
||||
Recommended: model.Recommended,
|
||||
AvailabilityBadge: model.AvailabilityBadge,
|
||||
})
|
||||
}
|
||||
return items
|
||||
}
|
||||
|
||||
func agentSelectionDescription(model agentchat.ModelOption) string {
|
||||
return strings.TrimSpace(model.Description)
|
||||
}
|
||||
|
||||
var agentGetwd = os.Getwd
|
||||
|
||||
func agentWorkingDir() string {
|
||||
cwd, err := agentGetwd()
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
return cwd
|
||||
}
|
||||
|
||||
func agentSystemPromptWithWorkingDir(modelName string, modelSystem string, extra string, workingDir string) string {
|
||||
return agentSystemPromptAtWithWorkingDir(time.Now(), modelName, modelSystem, extra, workingDir)
|
||||
}
|
||||
|
||||
func agentSystemPromptAtWithWorkingDir(now time.Time, modelName string, modelSystem string, extra string, workingDir string) string {
|
||||
var parts []string
|
||||
parts = append(parts, agentDefaultSystemPromptWithWorkingDir(now, modelName, workingDir))
|
||||
if strings.TrimSpace(modelSystem) != "" {
|
||||
parts = append(parts, strings.TrimSpace(modelSystem))
|
||||
}
|
||||
if strings.TrimSpace(extra) != "" {
|
||||
parts = append(parts, strings.TrimSpace(extra))
|
||||
}
|
||||
return strings.Join(parts, "\n\n")
|
||||
}
|
||||
|
||||
func agentDefaultSystemPromptWithWorkingDir(now time.Time, modelName string, workingDir string) string {
|
||||
date := now.Format("Monday, January 2, 2006")
|
||||
shellName := "bash"
|
||||
if runtime.GOOS == "windows" {
|
||||
shellName = "PowerShell"
|
||||
}
|
||||
parts := []string{
|
||||
"You are running in Ollama, in a harness to help the user accomplish tasks, and the model is " + modelName + ".",
|
||||
"",
|
||||
"Current date: " + date + ".",
|
||||
"",
|
||||
}
|
||||
parts = append(parts,
|
||||
"Be concise, practical, and action-oriented. Use tools when they materially help. Verify current or fast-changing facts with web tools when available; otherwise state uncertainty.",
|
||||
"",
|
||||
"Use "+shellName+" carefully. Prefer read-only inspection first. Stay within the current working directory unless explicitly asked. Surface intent before risky actions such as writes, deletes, moves, installs, git state changes, service changes, sudo, secrets access, network scripts, or commands outside the working directory. Request approval when required and do not work around denied approvals.",
|
||||
"",
|
||||
"Tell the user about meaningful changes, verification, failures, blockers, assumptions, and risks. Summarize routine tool output instead of dumping it.",
|
||||
)
|
||||
if workingDir != "" {
|
||||
parts = append(parts, "Current working directory: "+strconv.Quote(workingDir)+".")
|
||||
}
|
||||
return strings.Join(parts, "\n")
|
||||
}
|
||||
|
||||
func agentSystemFromShow(ctx context.Context, client *api.Client, modelName string) string {
|
||||
if client == nil || strings.TrimSpace(modelName) == "" {
|
||||
return ""
|
||||
}
|
||||
resp, err := client.Show(ctx, &api.ShowRequest{Model: modelName})
|
||||
if err != nil {
|
||||
fmt.Fprintf(os.Stderr, "\033[1mwarning:\033[0m could not load model system prompt: %v\n", err)
|
||||
return ""
|
||||
}
|
||||
return resp.System
|
||||
}
|
||||
|
||||
func agentToolsRegistry(ctx context.Context, client *api.Client, modelName string, skillCatalog *coreagent.SkillCatalog) *coreagent.Registry {
|
||||
supportsTools, err := agentModelSupportsTools(ctx, client, modelName)
|
||||
if err != nil {
|
||||
fmt.Fprintf(os.Stderr, "\033[1mwarning:\033[0m could not check model capabilities: %v\n", err)
|
||||
}
|
||||
if !supportsTools {
|
||||
return nil
|
||||
}
|
||||
|
||||
registry := &coreagent.Registry{}
|
||||
if os.Getenv("OLLAMA_AGENT_DISABLE_SHELL") == "" {
|
||||
registry.Register(&agenttools.Bash{})
|
||||
}
|
||||
registry.Register(&agenttools.Read{})
|
||||
registry.Register(&agenttools.Edit{})
|
||||
if len(skillCatalog.List()) > 0 {
|
||||
registry.Register(&agenttools.Skill{Catalog: skillCatalog})
|
||||
}
|
||||
|
||||
if os.Getenv("OLLAMA_AGENT_DISABLE_WEBSEARCH") == "" {
|
||||
if disabled, known := agentCloudStatusDisabled(ctx, client); !known || !disabled {
|
||||
registry.Register(&agenttools.WebSearch{})
|
||||
registry.Register(&agenttools.WebFetch{})
|
||||
} else {
|
||||
fmt.Fprintf(os.Stderr, "%s\n", internalcloud.DisabledError("web search is unavailable"))
|
||||
}
|
||||
}
|
||||
return registry
|
||||
}
|
||||
|
||||
func agentModelSupportsTools(ctx context.Context, client *api.Client, modelName string) (bool, error) {
|
||||
if client == nil || strings.TrimSpace(modelName) == "" {
|
||||
return false, nil
|
||||
}
|
||||
resp, err := client.Show(ctx, &api.ShowRequest{Model: modelName})
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
return slices.Contains(resp.Capabilities, model.CapabilityTools), nil
|
||||
}
|
||||
|
||||
func agentModelSupportsMultimodal(ctx context.Context, client *api.Client, modelName string) bool {
|
||||
if client == nil || strings.TrimSpace(modelName) == "" {
|
||||
return false
|
||||
}
|
||||
resp, err := client.Show(ctx, &api.ShowRequest{Model: modelName})
|
||||
if err != nil {
|
||||
fmt.Fprintf(os.Stderr, "\033[1mwarning:\033[0m could not check model capabilities: %v\n", err)
|
||||
return false
|
||||
}
|
||||
return showResponseSupportsMultimodal(resp)
|
||||
}
|
||||
|
||||
func showResponseSupportsMultimodal(resp *api.ShowResponse) bool {
|
||||
if resp == nil {
|
||||
return false
|
||||
}
|
||||
if slices.Contains(resp.Capabilities, model.CapabilityVision) || slices.Contains(resp.Capabilities, model.CapabilityAudio) {
|
||||
return true
|
||||
}
|
||||
if len(resp.ProjectorInfo) != 0 {
|
||||
return true
|
||||
}
|
||||
for key := range resp.ModelInfo {
|
||||
if strings.Contains(key, ".vision.") {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func agentContextWindowForModel(ctx context.Context, client *api.Client, modelName string, fallback int) int {
|
||||
if client == nil || strings.TrimSpace(modelName) == "" {
|
||||
return fallback
|
||||
}
|
||||
if tokens := loadedContextWindowForModel(ctx, client, modelName); tokens > 0 {
|
||||
return tokens
|
||||
}
|
||||
if modelref.HasExplicitCloudSource(modelName) {
|
||||
if tokens := agentRecommendationContextWindowForModel(ctx, client, modelName); tokens > 0 {
|
||||
return tokens
|
||||
}
|
||||
}
|
||||
resp, err := client.Show(ctx, &api.ShowRequest{Model: modelName})
|
||||
if err != nil {
|
||||
return fallback
|
||||
}
|
||||
if tokens := showResponseContextWindow(resp); tokens > 0 {
|
||||
return tokens
|
||||
}
|
||||
return fallback
|
||||
}
|
||||
|
||||
func agentRecommendationContextWindowForModel(ctx context.Context, client *api.Client, modelName string) int {
|
||||
if client == nil {
|
||||
return 0
|
||||
}
|
||||
recs, err := client.ModelRecommendationsExperimental(ctx)
|
||||
if err != nil || recs == nil {
|
||||
return 0
|
||||
}
|
||||
return contextWindowFromRecommendations(modelName, recs.Recommendations)
|
||||
}
|
||||
|
||||
func contextWindowFromRecommendations(modelName string, recommendations []api.ModelRecommendation) int {
|
||||
for _, rec := range recommendations {
|
||||
if rec.ContextLength <= 0 {
|
||||
continue
|
||||
}
|
||||
if sameModelRef(modelName, rec.Model) {
|
||||
return rec.ContextLength
|
||||
}
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
func sameModelRef(a, b string) bool {
|
||||
a = comparableModelRef(a)
|
||||
b = comparableModelRef(b)
|
||||
if strings.EqualFold(a, b) {
|
||||
return true
|
||||
}
|
||||
pa, errA := modelref.ParseRef(a)
|
||||
pb, errB := modelref.ParseRef(b)
|
||||
if errA != nil || errB != nil {
|
||||
return false
|
||||
}
|
||||
if !strings.EqualFold(pa.Base, pb.Base) {
|
||||
return false
|
||||
}
|
||||
return pa.Source == pb.Source ||
|
||||
pa.Source == modelref.ModelSourceUnspecified ||
|
||||
pb.Source == modelref.ModelSourceUnspecified
|
||||
}
|
||||
|
||||
func comparableModelRef(value string) string {
|
||||
value = strings.TrimSpace(value)
|
||||
if strings.HasSuffix(strings.ToLower(value), ":latest") {
|
||||
return strings.TrimSpace(value[:len(value)-len(":latest")])
|
||||
}
|
||||
return value
|
||||
}
|
||||
|
||||
func showResponseContextWindow(resp *api.ShowResponse) int {
|
||||
if resp == nil {
|
||||
return 0
|
||||
}
|
||||
if resp.Details.ContextLength > 0 {
|
||||
return resp.Details.ContextLength
|
||||
}
|
||||
if n, ok := numericModelInfo(resp.ModelInfo["general.context_length"]); ok {
|
||||
return n
|
||||
}
|
||||
best := 0
|
||||
for key, value := range resp.ModelInfo {
|
||||
if key != "context_length" && !strings.HasSuffix(key, ".context_length") {
|
||||
continue
|
||||
}
|
||||
if n, ok := numericModelInfo(value); ok && n > best {
|
||||
best = n
|
||||
}
|
||||
}
|
||||
return best
|
||||
}
|
||||
|
||||
func numericModelInfo(value any) (int, bool) {
|
||||
switch v := value.(type) {
|
||||
case int:
|
||||
return v, v > 0
|
||||
case int32:
|
||||
return int(v), v > 0
|
||||
case int64:
|
||||
return int(v), v > 0
|
||||
case uint:
|
||||
return int(v), v > 0
|
||||
case uint32:
|
||||
return int(v), v > 0
|
||||
case uint64:
|
||||
return int(v), v > 0
|
||||
case float64:
|
||||
return int(v), v > 0
|
||||
case string:
|
||||
n, err := strconv.Atoi(strings.TrimSpace(v))
|
||||
return n, err == nil && n > 0
|
||||
default:
|
||||
return 0, false
|
||||
}
|
||||
}
|
||||
|
||||
func preloadAgentModelIfLocal(ctx context.Context, client *api.Client, opts agentTUIOptions, modelName string, think *api.ThinkValue) (int, error) {
|
||||
modelName = strings.TrimSpace(modelName)
|
||||
if client == nil || modelName == "" {
|
||||
return 0, nil
|
||||
}
|
||||
if modelref.HasExplicitCloudSource(modelName) {
|
||||
return 0, nil
|
||||
}
|
||||
info, err := client.Show(ctx, &api.ShowRequest{Model: modelName})
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
if info.RemoteHost != "" {
|
||||
return 0, nil
|
||||
}
|
||||
if err := client.Generate(ctx, &api.GenerateRequest{
|
||||
Model: modelName,
|
||||
KeepAlive: opts.KeepAlive,
|
||||
Options: opts.Options,
|
||||
Think: think,
|
||||
}, func(api.GenerateResponse) error {
|
||||
return nil
|
||||
}); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return loadedContextWindowForModel(ctx, client, modelName), nil
|
||||
}
|
||||
|
||||
func loadedContextWindowForModel(ctx context.Context, client *api.Client, modelName string) int {
|
||||
if client == nil || strings.TrimSpace(modelName) == "" {
|
||||
return 0
|
||||
}
|
||||
resp, err := client.ListRunning(ctx)
|
||||
if err != nil {
|
||||
return 0
|
||||
}
|
||||
return processContextWindowForModel(modelName, resp)
|
||||
}
|
||||
|
||||
func processContextWindowForModel(modelName string, resp *api.ProcessResponse) int {
|
||||
if resp == nil {
|
||||
return 0
|
||||
}
|
||||
for _, running := range resp.Models {
|
||||
if running.ContextLength <= 0 {
|
||||
continue
|
||||
}
|
||||
if sameModelRef(modelName, running.Name) || sameModelRef(modelName, running.Model) {
|
||||
return running.ContextLength
|
||||
}
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
func agentModelOptions(ctx context.Context, client *api.Client) ([]agentchat.ModelOption, error) {
|
||||
if client == nil {
|
||||
return nil, errors.New("model picker requires an API client")
|
||||
}
|
||||
|
||||
list, err := client.List(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
seen := make(map[string]struct{})
|
||||
var options []agentchat.ModelOption
|
||||
add := func(name, description string, recommended bool, requiredPlan string, cloud bool) {
|
||||
name = strings.TrimSpace(name)
|
||||
if name == "" {
|
||||
return
|
||||
}
|
||||
key := strings.ToLower(name)
|
||||
if _, ok := seen[key]; ok {
|
||||
return
|
||||
}
|
||||
seen[key] = struct{}{}
|
||||
options = append(options, agentchat.ModelOption{
|
||||
Name: name,
|
||||
Description: strings.TrimSpace(description),
|
||||
Recommended: recommended,
|
||||
RequiredPlan: requiredPlan,
|
||||
Cloud: cloud,
|
||||
})
|
||||
}
|
||||
|
||||
if disabled, known := agentCloudStatusDisabled(ctx, client); !known || !disabled {
|
||||
if recs, err := client.ModelRecommendationsExperimental(ctx); err == nil {
|
||||
for _, rec := range recs.Recommendations {
|
||||
name := strings.TrimSpace(rec.Model)
|
||||
if !modelref.HasExplicitCloudSource(name) {
|
||||
continue
|
||||
}
|
||||
add(name, agentRecommendationDescription(rec), true, strings.TrimSpace(rec.RequiredPlan), true)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
local := slices.Clone(list.Models)
|
||||
slices.SortStableFunc(local, func(a, b api.ListModelResponse) int {
|
||||
return strings.Compare(strings.ToLower(a.Name), strings.ToLower(b.Name))
|
||||
})
|
||||
for _, model := range local {
|
||||
name := strings.TrimSpace(model.Name)
|
||||
if name == "" {
|
||||
name = strings.TrimSpace(model.Model)
|
||||
}
|
||||
name = strings.TrimSuffix(name, ":latest")
|
||||
if modelref.HasExplicitCloudSource(name) {
|
||||
add(name, agentCloudModelDescription(model), false, "", true)
|
||||
continue
|
||||
}
|
||||
add(name, agentLocalModelDescription(model), false, "", false)
|
||||
}
|
||||
|
||||
badges, signInURLs := cloudAvailabilityBadges(ctx, client, options)
|
||||
for i := range options {
|
||||
options[i].AvailabilityBadge = badges[options[i].Name]
|
||||
options[i].SignInURL = signInURLs[options[i].Name]
|
||||
}
|
||||
return options, nil
|
||||
}
|
||||
|
||||
func cloudAvailabilityBadges(ctx context.Context, client *api.Client, options []agentchat.ModelOption) (map[string]string, map[string]string) {
|
||||
badges := make(map[string]string)
|
||||
signInURLs := make(map[string]string)
|
||||
hasCloud := false
|
||||
for _, opt := range options {
|
||||
if opt.Cloud {
|
||||
hasCloud = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if !hasCloud {
|
||||
return badges, signInURLs
|
||||
}
|
||||
|
||||
if disabled, known := agentCloudStatusDisabled(ctx, client); known && disabled {
|
||||
return badges, signInURLs
|
||||
}
|
||||
|
||||
whoamiCtx, cancel := context.WithTimeout(ctx, 3*time.Second)
|
||||
defer cancel()
|
||||
user, err := client.Whoami(whoamiCtx)
|
||||
if err != nil {
|
||||
var authErr api.AuthorizationError
|
||||
signInURL := ""
|
||||
if errors.As(err, &authErr) && (authErr.StatusCode == http.StatusUnauthorized || authErr.SigninURL != "") {
|
||||
if authErr.SigninURL != "" {
|
||||
signInURL = authErr.SigninURL
|
||||
}
|
||||
} else {
|
||||
return badges, signInURLs
|
||||
}
|
||||
for _, opt := range options {
|
||||
if opt.Cloud {
|
||||
badges[opt.Name] = "Sign in required"
|
||||
if signInURL != "" {
|
||||
signInURLs[opt.Name] = signInURL
|
||||
}
|
||||
}
|
||||
}
|
||||
return badges, signInURLs
|
||||
}
|
||||
|
||||
signedIn := user != nil && user.Name != ""
|
||||
for _, opt := range options {
|
||||
if !opt.Cloud {
|
||||
continue
|
||||
}
|
||||
if !signedIn {
|
||||
badges[opt.Name] = "Sign in required"
|
||||
} else if opt.RequiredPlan != "" && !launch.PlanSatisfies(user.Plan, opt.RequiredPlan) {
|
||||
badges[opt.Name] = "Upgrade required"
|
||||
}
|
||||
}
|
||||
return badges, signInURLs
|
||||
}
|
||||
|
||||
func agentRecommendationDescription(rec api.ModelRecommendation) string {
|
||||
var parts []string
|
||||
if description := strings.TrimSpace(rec.Description); description != "" {
|
||||
parts = append(parts, description)
|
||||
} else {
|
||||
parts = append(parts, "cloud")
|
||||
}
|
||||
if rec.ContextLength > 0 {
|
||||
parts = append(parts, format.HumanNumber(uint64(rec.ContextLength))+" ctx")
|
||||
}
|
||||
return strings.Join(parts, " - ")
|
||||
}
|
||||
|
||||
func agentLocalModelDescription(model api.ListModelResponse) string {
|
||||
desc := agentModelArchDescription(model)
|
||||
if desc == "" {
|
||||
return "local"
|
||||
}
|
||||
return "local - " + desc
|
||||
}
|
||||
|
||||
func agentCloudModelDescription(model api.ListModelResponse) string {
|
||||
return agentModelArchDescription(model)
|
||||
}
|
||||
|
||||
func agentModelArchDescription(model api.ListModelResponse) string {
|
||||
var details []string
|
||||
if model.Details.Family != "" {
|
||||
details = append(details, model.Details.Family)
|
||||
}
|
||||
if ps := humanizedParameterSize(model.Details.ParameterSize); ps != "" {
|
||||
details = append(details, ps)
|
||||
}
|
||||
if model.Details.QuantizationLevel != "" {
|
||||
details = append(details, model.Details.QuantizationLevel)
|
||||
}
|
||||
var parts []string
|
||||
if len(details) > 0 {
|
||||
parts = append(parts, strings.Join(details, " "))
|
||||
}
|
||||
if model.Details.ContextLength > 0 {
|
||||
parts = append(parts, format.HumanNumber(uint64(model.Details.ContextLength))+" ctx")
|
||||
}
|
||||
return strings.Join(parts, " - ")
|
||||
}
|
||||
|
||||
func humanizedParameterSize(s string) string {
|
||||
s = strings.TrimSpace(s)
|
||||
if s == "" {
|
||||
return ""
|
||||
}
|
||||
if f, err := strconv.ParseFloat(s, 64); err == nil {
|
||||
return format.HumanNumber(uint64(f))
|
||||
}
|
||||
return s
|
||||
}
|
||||
|
||||
func agentCloudStatusDisabled(ctx context.Context, client *api.Client) (disabled bool, known bool) {
|
||||
if internalcloud.Disabled() {
|
||||
return true, true
|
||||
}
|
||||
|
||||
status, err := client.CloudStatusExperimental(ctx)
|
||||
if err != nil {
|
||||
var statusErr api.StatusError
|
||||
if errors.As(err, &statusErr) && statusErr.StatusCode == http.StatusNotFound {
|
||||
return false, false
|
||||
}
|
||||
return false, false
|
||||
}
|
||||
return status.Cloud.Disabled, true
|
||||
}
|
||||
|
||||
func ensureCloudModelAccess(ctx context.Context, client *api.Client, modelName, requiredPlan string) error {
|
||||
if client == nil {
|
||||
return errors.New("no API client available")
|
||||
}
|
||||
if disabled, known := agentCloudStatusDisabled(ctx, client); known && disabled {
|
||||
return errors.New("remote inference is unavailable")
|
||||
}
|
||||
user, err := client.Whoami(ctx)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if user != nil && user.Name != "" {
|
||||
if requiredPlan != "" && !launch.PlanSatisfies(user.Plan, requiredPlan) {
|
||||
return fmt.Errorf("plan upgrade required: %s needs plan %s, you have %s", modelName, requiredPlan, user.Plan)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
return fmt.Errorf("%s requires sign in", modelName)
|
||||
}
|
||||
@@ -0,0 +1,186 @@
|
||||
package cmd
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strconv"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/spf13/cobra"
|
||||
|
||||
coreagent "github.com/ollama/ollama/agent"
|
||||
agenttools "github.com/ollama/ollama/agent/tools"
|
||||
"github.com/ollama/ollama/api"
|
||||
"github.com/ollama/ollama/cmd/config"
|
||||
agentchat "github.com/ollama/ollama/cmd/tui/chat"
|
||||
)
|
||||
|
||||
func TestAgentSystemPromptIncludesSessionWorkingDirOnce(t *testing.T) {
|
||||
workingDir := t.TempDir()
|
||||
prompt := agentSystemPromptAtWithWorkingDir(
|
||||
time.Date(2026, time.July, 14, 0, 0, 0, 0, time.UTC),
|
||||
"test-model",
|
||||
"model instruction",
|
||||
"caller instruction",
|
||||
workingDir,
|
||||
)
|
||||
|
||||
workingDirInstruction := "Current working directory: " + strconv.Quote(workingDir) + "."
|
||||
if got := strings.Count(prompt, workingDirInstruction); got != 1 {
|
||||
t.Fatalf("working directory instruction count = %d, want 1:\n%s", got, prompt)
|
||||
}
|
||||
for _, want := range []string{"model instruction", "caller instruction"} {
|
||||
if !strings.Contains(prompt, want) {
|
||||
t.Fatalf("prompt missing %q:\n%s", want, prompt)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestAgentWorkingDirIgnoresGetwdFailure(t *testing.T) {
|
||||
original := agentGetwd
|
||||
agentGetwd = func() (string, error) {
|
||||
return "", errors.New("getwd failed")
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
agentGetwd = original
|
||||
})
|
||||
|
||||
if got := agentWorkingDir(); got != "" {
|
||||
t.Fatalf("working directory = %q, want empty on getwd failure", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAgentSystemPromptIncludesSkillCatalog(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
if err := os.Mkdir(filepath.Join(dir, "release-notes"), 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := os.WriteFile(filepath.Join(dir, "release-notes", "SKILL.md"), []byte("---\nname: release-notes\ndescription: Draft releases.\n---\nUse bullets."), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
catalog, err := coreagent.DiscoverSkills(dir)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
got := agentSystemPromptAtWithWorkingDir(time.Date(2026, 7, 14, 0, 0, 0, 0, time.UTC), "model", "", catalog.SystemContext(), "")
|
||||
if !strings.Contains(got, "release-notes: Draft releases.") || !strings.Contains(got, "normal approval rules") {
|
||||
t.Fatalf("system prompt missing skill context: %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAgentSkillSystemContextRequiresAvailableEnabledSkillTool(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
if err := os.Mkdir(filepath.Join(dir, "release-notes"), 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := os.WriteFile(filepath.Join(dir, "release-notes", "SKILL.md"), []byte("---\nname: release-notes\ndescription: Draft releases.\n---\nUse bullets."), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
catalog, err := coreagent.DiscoverSkills(dir)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
registry := &coreagent.Registry{}
|
||||
registry.Register(&agenttools.Skill{Catalog: catalog})
|
||||
|
||||
if got := agentSkillSystemContext(catalog, registry, false); !strings.Contains(got, "release-notes: Draft releases.") {
|
||||
t.Fatalf("enabled skill context = %q", got)
|
||||
}
|
||||
if got := agentSkillSystemContext(catalog, registry, true); got != "" {
|
||||
t.Fatalf("disabled tools should omit skill context, got %q", got)
|
||||
}
|
||||
if got := agentSkillSystemContext(catalog, &coreagent.Registry{}, false); got != "" {
|
||||
t.Fatalf("unavailable skill tool should omit skill context, got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAgentSelectionItemsUseLaunchSections(t *testing.T) {
|
||||
items := agentSelectionItems([]agentchat.ModelOption{
|
||||
{Name: "glm-5.2:cloud", Description: "cloud", Recommended: true, Cloud: true},
|
||||
{Name: "llama3.2", Description: "local"},
|
||||
})
|
||||
|
||||
if len(items) != 2 {
|
||||
t.Fatalf("items = %d, want 2", len(items))
|
||||
}
|
||||
if !items[0].Recommended {
|
||||
t.Fatalf("cloud recommendation should be pinned: %#v", items[0])
|
||||
}
|
||||
if items[1].Recommended {
|
||||
t.Fatalf("local selected model should stay in launch More section: %#v", items[1])
|
||||
}
|
||||
if items[1].Description != "local" {
|
||||
t.Fatalf("selected model description = %q, want plain description", items[1].Description)
|
||||
}
|
||||
}
|
||||
|
||||
func TestContextWindowFromRecommendationsMatchesCloudModel(t *testing.T) {
|
||||
got := contextWindowFromRecommendations("glm-5.2:cloud", []api.ModelRecommendation{
|
||||
{Model: "gemma4:cloud", ContextLength: 32768},
|
||||
{Model: "glm-5.2:cloud", ContextLength: 1048576},
|
||||
})
|
||||
if got != 1048576 {
|
||||
t.Fatalf("context window = %d, want 1048576", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestShowResponseContextWindowReadsArchitectureContextLength(t *testing.T) {
|
||||
got := showResponseContextWindow(&api.ShowResponse{
|
||||
ModelInfo: map[string]any{
|
||||
"qwen3.context_length": uint32(262144),
|
||||
"qwen3.rope.scaling.original_context_length": uint32(32768),
|
||||
},
|
||||
})
|
||||
if got != 262144 {
|
||||
t.Fatalf("context window = %d, want 262144", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestProcessContextWindowForModelMatchesLatestAlias(t *testing.T) {
|
||||
got := processContextWindowForModel("ornith", &api.ProcessResponse{
|
||||
Models: []api.ProcessModelResponse{
|
||||
{Name: "other:latest", Model: "other:latest", ContextLength: 32768},
|
||||
{Name: "ornith:latest", Model: "ornith:latest", ContextLength: 262144},
|
||||
},
|
||||
})
|
||||
if got != 262144 {
|
||||
t.Fatalf("context window = %d, want 262144", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSaveLastAgentModel(t *testing.T) {
|
||||
setCmdTestHome(t, t.TempDir())
|
||||
|
||||
if err := saveLastAgentModel(" qwen3:8b "); err != nil {
|
||||
t.Fatalf("saveLastAgentModel returned error: %v", err)
|
||||
}
|
||||
if got := config.LastModel(); got != "qwen3:8b" {
|
||||
t.Fatalf("last model = %q, want qwen3:8b", got)
|
||||
}
|
||||
|
||||
if err := saveLastAgentModel(" "); err != nil {
|
||||
t.Fatalf("saveLastAgentModel blank returned error: %v", err)
|
||||
}
|
||||
if got := config.LastModel(); got != "qwen3:8b" {
|
||||
t.Fatalf("blank save changed last model to %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestApplyAgentFlagsNoTools(t *testing.T) {
|
||||
cmd := &cobra.Command{}
|
||||
registerAgentFlags(cmd)
|
||||
if err := cmd.Flags().Set("no-tools", "true"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
var opts agentTUIOptions
|
||||
if _, err := applyAgentFlags(cmd, &opts); err != nil {
|
||||
t.Fatalf("applyAgentFlags returned error: %v", err)
|
||||
}
|
||||
if !opts.ToolsDisabled {
|
||||
t.Fatal("--no-tools should disable tools")
|
||||
}
|
||||
}
|
||||
+130
-66
@@ -27,6 +27,7 @@ import (
|
||||
"strings"
|
||||
"sync/atomic"
|
||||
"syscall"
|
||||
"text/tabwriter"
|
||||
"time"
|
||||
|
||||
"github.com/containerd/console"
|
||||
@@ -55,7 +56,6 @@ import (
|
||||
"github.com/ollama/ollama/types/model"
|
||||
"github.com/ollama/ollama/types/syncmap"
|
||||
"github.com/ollama/ollama/version"
|
||||
xcmd "github.com/ollama/ollama/x/cmd"
|
||||
xcreate "github.com/ollama/ollama/x/create"
|
||||
xcreateclient "github.com/ollama/ollama/x/create/client"
|
||||
"github.com/ollama/ollama/x/imagegen"
|
||||
@@ -96,6 +96,8 @@ func init() {
|
||||
}
|
||||
|
||||
launch.DefaultConfirmPrompt = tui.RunConfirmWithOptions
|
||||
|
||||
launch.DefaultSpinner = tui.RunSpinner
|
||||
}
|
||||
|
||||
func runTUISingleSelector(title string, items []launch.SelectionItem, current string, updates <-chan []launch.SelectionItem) (string, error) {
|
||||
@@ -885,11 +887,6 @@ func RunHandler(cmd *cobra.Command, args []string) error {
|
||||
return imagegen.RunCLI(cmd, name, opts.Prompt, interactive, opts.KeepAlive)
|
||||
}
|
||||
|
||||
// Check for experimental flag
|
||||
isExperimental, _ := cmd.Flags().GetBool("experimental")
|
||||
yoloMode, _ := cmd.Flags().GetBool("experimental-yolo")
|
||||
enableWebsearch, _ := cmd.Flags().GetBool("experimental-websearch")
|
||||
|
||||
if interactive {
|
||||
if err := loadOrUnloadModel(cmd, &opts); err != nil {
|
||||
var sErr api.AuthorizationError
|
||||
@@ -916,11 +913,6 @@ func RunHandler(cmd *cobra.Command, args []string) error {
|
||||
}
|
||||
}
|
||||
|
||||
// Use experimental agent loop with tools
|
||||
if isExperimental {
|
||||
return xcmd.GenerateInteractive(cmd, opts.Model, opts.WordWrap, opts.Options, opts.Think, opts.HideThinking, opts.KeepAlive, yoloMode, enableWebsearch)
|
||||
}
|
||||
|
||||
return generateInteractive(cmd, opts)
|
||||
}
|
||||
if err := generate(cmd, opts); err != nil {
|
||||
@@ -986,6 +978,100 @@ func SignoutHandler(cmd *cobra.Command, args []string) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func UsageHandler(cmd *cobra.Command, args []string) error {
|
||||
out := cmd.OutOrStdout()
|
||||
|
||||
client, err := api.ClientFromEnvironment()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
usage, err := client.Usage(cmd.Context())
|
||||
if err != nil {
|
||||
var aErr api.AuthorizationError
|
||||
if errors.As(err, &aErr) && aErr.StatusCode == http.StatusUnauthorized {
|
||||
fmt.Fprintln(out, "You need to be signed in to Ollama to view usage.")
|
||||
fmt.Fprintln(out)
|
||||
if aErr.SigninURL != "" {
|
||||
_ = browser.OpenURL(aErr.SigninURL)
|
||||
fmt.Fprintf(out, ConnectInstructions, aErr.SigninURL)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
return err
|
||||
}
|
||||
|
||||
fmt.Fprintln(out, "Usage")
|
||||
details := tabwriter.NewWriter(out, 0, 4, 2, ' ', 0)
|
||||
fmt.Fprintf(details, " Period\t%s to %s\n", usage.Activity.Period.StartingAt.Format("2006-01-02"), usage.Activity.Period.EndingAt.Format("2006-01-02"))
|
||||
fmt.Fprintf(details, " Spend\t$%s\n", usage.Activity.Cost)
|
||||
if err := details.Flush(); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if len(usage.Activity.Models) == 0 && usageLimitEmpty(usage.Limits.Session) && usageLimitEmpty(usage.Limits.Weekly) {
|
||||
fmt.Fprintln(out)
|
||||
fmt.Fprintln(out, "No usage recorded for this period.")
|
||||
return nil
|
||||
}
|
||||
|
||||
if len(usage.Activity.Models) > 0 {
|
||||
fmt.Fprintln(out)
|
||||
fmt.Fprintln(out, "Activity")
|
||||
table := tabwriter.NewWriter(out, 0, 4, 2, ' ', 0)
|
||||
fmt.Fprintln(table, " Model\tRequests\tSpend")
|
||||
for _, m := range usage.Activity.Models {
|
||||
fmt.Fprintf(table, " %s\t%d\t$%s\n", usageModelName(m.Name), m.RequestCount, m.Cost)
|
||||
}
|
||||
if err := table.Flush(); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
if err := writeUsageLimit(out, "Session", usage.Limits.Session); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := writeUsageLimit(out, "Weekly", usage.Limits.Weekly); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func usageLimitEmpty(limit api.UsageLimit) bool {
|
||||
return limit.Usage == 0 && len(limit.Models) == 0
|
||||
}
|
||||
|
||||
func usageModelName(name string) string {
|
||||
switch name {
|
||||
case "web search":
|
||||
return "Web Search"
|
||||
case "web fetch":
|
||||
return "Web Fetch"
|
||||
default:
|
||||
return name
|
||||
}
|
||||
}
|
||||
|
||||
func writeUsageLimit(out io.Writer, name string, limit api.UsageLimit) error {
|
||||
if usageLimitEmpty(limit) {
|
||||
return nil
|
||||
}
|
||||
|
||||
fmt.Fprintln(out)
|
||||
fmt.Fprintln(out, name)
|
||||
table := tabwriter.NewWriter(out, 0, 4, 2, ' ', 0)
|
||||
fmt.Fprintf(table, " Used\t%.1f%%\n", limit.Usage*100)
|
||||
if len(limit.Models) > 0 {
|
||||
fmt.Fprintln(table, " Model\tRequests")
|
||||
}
|
||||
for _, m := range limit.Models {
|
||||
fmt.Fprintf(table, " %s\t%d\n", usageModelName(m.Name), m.RequestCount)
|
||||
}
|
||||
return table.Flush()
|
||||
}
|
||||
|
||||
func PushHandler(cmd *cobra.Command, args []string) error {
|
||||
client, err := api.ClientFromEnvironment()
|
||||
if err != nil {
|
||||
@@ -2143,72 +2229,32 @@ func ensureServerRunning(ctx context.Context) error {
|
||||
}
|
||||
|
||||
func launchInteractiveModel(cmd *cobra.Command, modelName string) error {
|
||||
opts := runOptions{
|
||||
Model: modelName,
|
||||
WordWrap: os.Getenv("TERM") == "xterm-256color",
|
||||
Options: map[string]any{},
|
||||
ShowConnect: true,
|
||||
}
|
||||
|
||||
client, err := api.ClientFromEnvironment()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
requestedCloud := modelref.HasExplicitCloudSource(modelName)
|
||||
|
||||
info, err := func() (*api.ShowResponse, error) {
|
||||
showReq := &api.ShowRequest{Name: modelName}
|
||||
info, err := client.Show(cmd.Context(), showReq)
|
||||
var se api.StatusError
|
||||
if errors.As(err, &se) && se.StatusCode == http.StatusNotFound {
|
||||
if requestedCloud {
|
||||
return nil, err
|
||||
}
|
||||
if err := PullHandler(cmd, []string{modelName}); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return client.Show(cmd.Context(), &api.ShowRequest{Name: modelName})
|
||||
}
|
||||
return info, err
|
||||
}()
|
||||
opts := agentTUIOptions{
|
||||
Model: modelName,
|
||||
Options: map[string]any{},
|
||||
}
|
||||
info, err := prepareAgentModel(cmd, client, &opts, false)
|
||||
if err != nil {
|
||||
if handleCloudAuthorizationError(err) {
|
||||
return nil
|
||||
}
|
||||
return err
|
||||
}
|
||||
opts.System = info.System
|
||||
|
||||
ensureCloudStub(cmd.Context(), client, modelName)
|
||||
|
||||
opts.Think, err = inferThinkingOption(&info.Capabilities, &opts, false)
|
||||
if err != nil {
|
||||
if err := saveLastAgentModel(opts.Model); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
audioCapable := slices.Contains(info.Capabilities, model.CapabilityAudio)
|
||||
opts.MultiModal = slices.Contains(info.Capabilities, model.CapabilityVision) || audioCapable
|
||||
|
||||
// TODO: remove the projector info and vision info checks below,
|
||||
// these are left in for backwards compatibility with older servers
|
||||
// that don't have the capabilities field in the model info
|
||||
if len(info.ProjectorInfo) != 0 {
|
||||
opts.MultiModal = true
|
||||
}
|
||||
for k := range info.ModelInfo {
|
||||
if strings.Contains(k, ".vision.") {
|
||||
opts.MultiModal = true
|
||||
break
|
||||
if err := GenerateAgentTUI(cmd, client, opts); err != nil {
|
||||
if handleCloudAuthorizationError(err) {
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
applyShowResponseToRunOptions(&opts, info)
|
||||
|
||||
if err := loadOrUnloadModel(cmd, &opts); err != nil {
|
||||
return fmt.Errorf("error loading model: %w", err)
|
||||
}
|
||||
if err := generateInteractive(cmd, opts); err != nil {
|
||||
return fmt.Errorf("error running model: %w", err)
|
||||
return fmt.Errorf("error running agent: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -2324,7 +2370,7 @@ func runLauncherAction(cmd *cobra.Command, action tui.TUIAction, deps launcherDe
|
||||
|
||||
func launcherActionExitsLoop(integration string) bool {
|
||||
switch integration {
|
||||
case "codex-app", "vscode":
|
||||
case "chatgpt", "codex-app", "vscode":
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
@@ -2413,9 +2459,6 @@ func NewCLI() *cobra.Command {
|
||||
runCmd.Flags().Bool("hidethinking", false, "Hide thinking output (if provided)")
|
||||
runCmd.Flags().Bool("truncate", false, "For embedding models: truncate inputs exceeding context length (default: true). Set --truncate=false to error instead")
|
||||
runCmd.Flags().Int("dimensions", 0, "Truncate output embeddings to specified dimension (embedding models only)")
|
||||
runCmd.Flags().Bool("experimental", false, "Enable experimental agent loop with tools")
|
||||
runCmd.Flags().Bool("experimental-yolo", false, "Skip all tool approval prompts (use with caution)")
|
||||
runCmd.Flags().Bool("experimental-websearch", false, "Enable web search tool in experimental mode")
|
||||
|
||||
// Image generation flags (width, height, steps, seed, etc.)
|
||||
imagegen.RegisterFlags(runCmd)
|
||||
@@ -2423,6 +2466,15 @@ func NewCLI() *cobra.Command {
|
||||
runCmd.Flags().Bool("imagegen", false, "Use the imagegen runner for LLM inference")
|
||||
runCmd.Flags().MarkHidden("imagegen")
|
||||
|
||||
agentCmd := &cobra.Command{
|
||||
Use: "agent",
|
||||
Short: "Run an agent",
|
||||
Args: cobra.ExactArgs(0),
|
||||
PreRunE: checkServerHeartbeat,
|
||||
RunE: AgentHandler,
|
||||
}
|
||||
registerAgentFlags(agentCmd)
|
||||
|
||||
stopCmd := &cobra.Command{
|
||||
Use: "stop MODEL",
|
||||
Short: "Stop a running model",
|
||||
@@ -2493,6 +2545,14 @@ func NewCLI() *cobra.Command {
|
||||
RunE: SignoutHandler,
|
||||
}
|
||||
|
||||
usageCmd := &cobra.Command{
|
||||
Use: "usage",
|
||||
Short: "Show your ollama.com usage",
|
||||
Args: cobra.ExactArgs(0),
|
||||
PreRunE: checkServerHeartbeat,
|
||||
RunE: UsageHandler,
|
||||
}
|
||||
|
||||
listCmd := &cobra.Command{
|
||||
Use: "list",
|
||||
Aliases: []string{"ls"},
|
||||
@@ -2553,9 +2613,11 @@ func NewCLI() *cobra.Command {
|
||||
createCmd,
|
||||
showCmd,
|
||||
runCmd,
|
||||
agentCmd,
|
||||
stopCmd,
|
||||
pullCmd,
|
||||
pushCmd,
|
||||
usageCmd,
|
||||
listCmd,
|
||||
psCmd,
|
||||
copyCmd,
|
||||
@@ -2600,6 +2662,7 @@ func NewCLI() *cobra.Command {
|
||||
createCmd,
|
||||
showCmd,
|
||||
runCmd,
|
||||
agentCmd,
|
||||
stopCmd,
|
||||
pullCmd,
|
||||
pushCmd,
|
||||
@@ -2607,6 +2670,7 @@ func NewCLI() *cobra.Command {
|
||||
loginCmd,
|
||||
signoutCmd,
|
||||
logoutCmd,
|
||||
usageCmd,
|
||||
listCmd,
|
||||
psCmd,
|
||||
copyCmd,
|
||||
|
||||
@@ -249,7 +249,7 @@ func TestRunLauncherAction_GUIAppsExitTUILoop(t *testing.T) {
|
||||
cmd := &cobra.Command{}
|
||||
cmd.SetContext(context.Background())
|
||||
|
||||
for _, integration := range []string{"codex-app", "vscode"} {
|
||||
for _, integration := range []string{"chatgpt", "vscode"} {
|
||||
continueLoop, err := runLauncherAction(cmd, tui.TUIAction{Kind: tui.TUIActionLaunchIntegration, Integration: integration}, launcherDeps{
|
||||
resolveRunModel: unexpectedRunModelResolution(t),
|
||||
launchIntegration: func(ctx context.Context, req launch.IntegrationLaunchRequest) error {
|
||||
|
||||
@@ -1398,6 +1398,102 @@ func TestListHandler(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestUsageHandler(t *testing.T) {
|
||||
startsAt := time.Date(2026, time.June, 29, 0, 0, 0, 0, time.UTC)
|
||||
endsAt := time.Date(2026, time.July, 27, 0, 0, 0, 0, time.UTC)
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
statusCode int
|
||||
response any
|
||||
want string
|
||||
}{
|
||||
{
|
||||
name: "activity and limits",
|
||||
statusCode: http.StatusOK,
|
||||
response: api.UsageResponse{
|
||||
Activity: api.UsageActivity{
|
||||
Cost: "12.34000",
|
||||
Period: api.UsagePeriod{
|
||||
Type: "last_4_weeks",
|
||||
StartingAt: startsAt,
|
||||
EndingAt: endsAt,
|
||||
},
|
||||
Models: []api.UsageModel{{Name: "gpt-oss:120b", RequestCount: 42, Cost: "12.34000"}},
|
||||
},
|
||||
Limits: api.UsageLimits{
|
||||
Session: api.UsageLimit{Usage: 0.006, Models: []api.UsageModel{{Name: "web search", RequestCount: 1}}},
|
||||
},
|
||||
},
|
||||
want: "Usage\n" +
|
||||
" Period 2026-06-29 to 2026-07-27\n" +
|
||||
" Spend $12.34000\n\n" +
|
||||
"Activity\n" +
|
||||
" Model Requests Spend\n" +
|
||||
" gpt-oss:120b 42 $12.34000\n\n" +
|
||||
"Session\n" +
|
||||
" Used 0.6%\n" +
|
||||
" Model Requests\n" +
|
||||
" Web Search 1\n",
|
||||
},
|
||||
{
|
||||
name: "no usage",
|
||||
statusCode: http.StatusOK,
|
||||
response: api.UsageResponse{
|
||||
Activity: api.UsageActivity{
|
||||
Cost: "0.00000",
|
||||
Period: api.UsagePeriod{Type: "last_4_weeks", StartingAt: startsAt, EndingAt: endsAt},
|
||||
Models: []api.UsageModel{},
|
||||
},
|
||||
Limits: api.UsageLimits{
|
||||
Session: api.UsageLimit{Models: []api.UsageModel{}},
|
||||
Weekly: api.UsageLimit{Models: []api.UsageModel{}},
|
||||
},
|
||||
},
|
||||
want: "Usage\n" +
|
||||
" Period 2026-06-29 to 2026-07-27\n" +
|
||||
" Spend $0.00000\n\n" +
|
||||
"No usage recorded for this period.\n",
|
||||
},
|
||||
{
|
||||
name: "not signed in",
|
||||
statusCode: http.StatusUnauthorized,
|
||||
response: map[string]string{"error": "unauthorized"},
|
||||
want: "You need to be signed in to Ollama to view usage.\n\n",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodGet || r.URL.Path != "/api/usage" {
|
||||
t.Fatalf("request = %s %s, want GET /api/usage", r.Method, r.URL.Path)
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
w.WriteHeader(tt.statusCode)
|
||||
if err := json.NewEncoder(w).Encode(tt.response); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
t.Setenv("OLLAMA_HOST", server.URL)
|
||||
|
||||
cmd := &cobra.Command{}
|
||||
cmd.SetContext(t.Context())
|
||||
var out bytes.Buffer
|
||||
cmd.SetOut(&out)
|
||||
|
||||
if err := UsageHandler(cmd, nil); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if got := out.String(); got != tt.want {
|
||||
t.Errorf("unexpected output (-want +got):\n%s", cmp.Diff(tt.want, got))
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestCreateHandler(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
|
||||
@@ -0,0 +1,189 @@
|
||||
package filedata
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"regexp"
|
||||
"slices"
|
||||
"strings"
|
||||
|
||||
"github.com/ollama/ollama/api"
|
||||
)
|
||||
|
||||
type File struct {
|
||||
Path string
|
||||
Data api.ImageData
|
||||
}
|
||||
|
||||
func NormalizePath(fp string) string {
|
||||
fp = strings.Trim(fp, "\"")
|
||||
fp = strings.NewReplacer(
|
||||
"\\ ", " ",
|
||||
"\\(", "(",
|
||||
"\\)", ")",
|
||||
"\\[", "[",
|
||||
"\\]", "]",
|
||||
"\\{", "{",
|
||||
"\\}", "}",
|
||||
"\\$", "$",
|
||||
"\\&", "&",
|
||||
"\\;", ";",
|
||||
"\\'", "'",
|
||||
"\\\\", "\\",
|
||||
"\\*", "*",
|
||||
"\\?", "?",
|
||||
"\\~", "~",
|
||||
).Replace(fp)
|
||||
|
||||
if u, err := url.Parse(fp); err == nil && strings.EqualFold(u.Scheme, "file") {
|
||||
return normalizeFileURL(u)
|
||||
} else if normalized, ok := normalizeMalformedFileURL(fp); ok {
|
||||
return normalized
|
||||
}
|
||||
|
||||
return fp
|
||||
}
|
||||
|
||||
// fileExtractRe matches file:// URLs and filesystem paths ending in image/audio
|
||||
// extensions. Hoisted to package scope so the per-keystroke slash-completion
|
||||
// path (chat.slashInputIsMultimodalFile -> ExtractNames) doesn't recompile it
|
||||
// on every call.
|
||||
var fileExtractRe = regexp.MustCompile(`(?:file://\S+?\.(?i:jpg|jpeg|png|webp|wav)\b)|(?:(?:[a-zA-Z]:)?(?:\./|\.\\|/|\\)[\S\\ ]+?\.(?i:jpg|jpeg|png|webp|wav)\b)`)
|
||||
|
||||
func ExtractNames(input string) []string {
|
||||
return fileExtractRe.FindAllString(input, -1)
|
||||
}
|
||||
|
||||
func Extract(input string) (string, []api.ImageData, error) {
|
||||
cleaned, files, err := ExtractWithFiles(input)
|
||||
if err != nil {
|
||||
return "", nil, err
|
||||
}
|
||||
data := make([]api.ImageData, 0, len(files))
|
||||
for _, file := range files {
|
||||
data = append(data, file.Data)
|
||||
}
|
||||
return cleaned, data, nil
|
||||
}
|
||||
|
||||
func ExtractWithFiles(input string) (string, []File, error) {
|
||||
filePaths := ExtractNames(input)
|
||||
var files []File
|
||||
|
||||
for _, fp := range filePaths {
|
||||
nfp := NormalizePath(fp)
|
||||
data, err := GetData(nfp)
|
||||
if errors.Is(err, os.ErrNotExist) {
|
||||
continue
|
||||
} else if err != nil {
|
||||
return "", nil, fmt.Errorf("couldn't process file %q: %w", nfp, err)
|
||||
}
|
||||
input = strings.ReplaceAll(input, "'"+nfp+"'", "")
|
||||
input = strings.ReplaceAll(input, "'"+fp+"'", "")
|
||||
input = strings.ReplaceAll(input, `"`+nfp+`"`, "")
|
||||
input = strings.ReplaceAll(input, `"`+fp+`"`, "")
|
||||
input = strings.ReplaceAll(input, fp, "")
|
||||
files = append(files, File{Path: nfp, Data: data})
|
||||
}
|
||||
return strings.TrimSpace(input), files, nil
|
||||
}
|
||||
|
||||
func GetData(filePath string) ([]byte, error) {
|
||||
file, err := os.Open(filePath)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer file.Close()
|
||||
|
||||
buf := make([]byte, 512)
|
||||
_, err = file.Read(buf)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
contentType := http.DetectContentType(buf)
|
||||
allowedTypes := []string{"image/jpeg", "image/jpg", "image/png", "image/webp", "audio/wave"}
|
||||
if !slices.Contains(allowedTypes, contentType) {
|
||||
return nil, fmt.Errorf("invalid file type: %s", contentType)
|
||||
}
|
||||
|
||||
info, err := file.Stat()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
var maxSize int64 = 100 * 1024 * 1024
|
||||
if info.Size() > maxSize {
|
||||
return nil, errors.New("file size exceeds maximum limit (100MB)")
|
||||
}
|
||||
|
||||
buf = make([]byte, info.Size())
|
||||
_, err = file.Seek(0, 0)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
_, err = io.ReadFull(file, buf)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return buf, nil
|
||||
}
|
||||
|
||||
func Kind(path string) string {
|
||||
if strings.EqualFold(filepath.Ext(path), ".wav") {
|
||||
return "audio"
|
||||
}
|
||||
return "image"
|
||||
}
|
||||
|
||||
func normalizeFileURL(u *url.URL) string {
|
||||
path := u.Path
|
||||
if unescaped, err := url.PathUnescape(path); err == nil {
|
||||
path = unescaped
|
||||
}
|
||||
host := u.Host
|
||||
if unescaped, err := url.PathUnescape(host); err == nil {
|
||||
host = unescaped
|
||||
}
|
||||
if len(host) >= 2 && host[1] == ':' && isASCIIAlpha(host[0]) {
|
||||
return filepath.Clean(filepath.FromSlash(host + path))
|
||||
}
|
||||
if len(path) >= 4 && path[0] == '/' && path[2] == ':' && isASCIIAlpha(path[1]) {
|
||||
path = path[1:]
|
||||
}
|
||||
if u.Host != "" && !strings.EqualFold(u.Host, "localhost") {
|
||||
return `\\` + u.Host + filepath.FromSlash(path)
|
||||
}
|
||||
return filepath.FromSlash(path)
|
||||
}
|
||||
|
||||
func normalizeMalformedFileURL(raw string) (string, bool) {
|
||||
const prefix = "file://"
|
||||
if !strings.HasPrefix(strings.ToLower(raw), prefix) {
|
||||
return "", false
|
||||
}
|
||||
|
||||
path := raw[len(prefix):]
|
||||
if unescaped, err := url.PathUnescape(path); err == nil {
|
||||
path = unescaped
|
||||
}
|
||||
path = strings.TrimPrefix(path, "localhost")
|
||||
if len(path) >= 3 && path[0] == '/' && path[2] == ':' && isASCIIAlpha(path[1]) {
|
||||
path = path[1:]
|
||||
}
|
||||
if len(path) >= 2 && path[1] == ':' && isASCIIAlpha(path[0]) {
|
||||
return filepath.Clean(filepath.FromSlash(path)), true
|
||||
}
|
||||
return "", false
|
||||
}
|
||||
|
||||
func isASCIIAlpha(b byte) bool {
|
||||
return (b >= 'a' && b <= 'z') || (b >= 'A' && b <= 'Z')
|
||||
}
|
||||
@@ -0,0 +1,223 @@
|
||||
package filedata
|
||||
|
||||
import (
|
||||
"net/url"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestNormalizePathMalformedWindowsFileURL(t *testing.T) {
|
||||
got := NormalizePath(`file://C:%5CUsers%5Cjdoe%5CPictures%5Cimg.png`)
|
||||
want := filepath.Clean(`C:\Users\jdoe\Pictures\img.png`)
|
||||
if got != want {
|
||||
t.Fatalf("path = %q, want %q", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNormalizePathTwoSlashWindowsFileURL(t *testing.T) {
|
||||
got := NormalizePath(`file://C:/Users/jdoe/Pictures/img.png`)
|
||||
want := filepath.Clean(`C:/Users/jdoe/Pictures/img.png`)
|
||||
if got != want {
|
||||
t.Fatalf("path = %q, want %q", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNormalizePathLocalhostWindowsFileURL(t *testing.T) {
|
||||
got := NormalizePath(`file://localhost/C:/Users/jdoe/Pictures/img.png`)
|
||||
want := filepath.Clean(`C:/Users/jdoe/Pictures/img.png`)
|
||||
if got != want {
|
||||
t.Fatalf("path = %q, want %q", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestExtractNames(t *testing.T) {
|
||||
// Unix style paths
|
||||
input := ` some preamble
|
||||
./relative\ path/one.png inbetween1 ./not a valid two.jpg inbetween2 ./1.svg
|
||||
/unescaped space /three.jpeg inbetween3 /valid\ path/dir/four.png "./quoted with spaces/five.JPG
|
||||
/unescaped space /six.webp inbetween6 /valid\ path/dir/seven.WEBP`
|
||||
res := ExtractNames(input)
|
||||
if len(res) != 7 {
|
||||
t.Fatalf("len = %d, want 7", len(res))
|
||||
}
|
||||
assertContains(t, res[0], "one.png")
|
||||
assertContains(t, res[1], "two.jpg")
|
||||
assertContains(t, res[2], "three.jpeg")
|
||||
assertContains(t, res[3], "four.png")
|
||||
assertContains(t, res[4], "five.JPG")
|
||||
assertContains(t, res[5], "six.webp")
|
||||
assertContains(t, res[6], "seven.WEBP")
|
||||
assertNotContains(t, res[4], "\"")
|
||||
for _, r := range res {
|
||||
assertNotContains(t, r, "inbetween1")
|
||||
}
|
||||
assertNotContainsSlice(t, res, "./1.svg")
|
||||
}
|
||||
|
||||
func TestExtractNamesWindowsPaths(t *testing.T) {
|
||||
input := ` some preamble
|
||||
c:/users/jdoe/one.png inbetween1 c:/program files/someplace/two.jpg inbetween2
|
||||
/absolute/nospace/three.jpeg inbetween3 /absolute/with space/four.png inbetween4
|
||||
./relative\ path/five.JPG inbetween5 "./relative with/spaces/six.png inbetween6
|
||||
d:\path with\spaces\seven.JPEG inbetween7 c:\users\jdoe\eight.png inbetween8
|
||||
d:\program files\someplace\nine.png inbetween9 "E:\program files\someplace\ten.PNG
|
||||
c:/users/jdoe/eleven.webp inbetween11 c:/program files/someplace/twelve.WebP inbetween12
|
||||
d:\path with\spaces\thirteen.WEBP some ending
|
||||
`
|
||||
res := ExtractNames(input)
|
||||
if len(res) != 13 {
|
||||
t.Fatalf("len = %d, want 13", len(res))
|
||||
}
|
||||
assertNotContainsSlice(t, res, "inbetween2")
|
||||
assertContains(t, res[0], "one.png")
|
||||
assertContains(t, res[0], "c:")
|
||||
assertContains(t, res[1], "two.jpg")
|
||||
assertContains(t, res[1], "c:")
|
||||
assertContains(t, res[2], "three.jpeg")
|
||||
assertContains(t, res[3], "four.png")
|
||||
assertContains(t, res[4], "five.JPG")
|
||||
assertContains(t, res[5], "six.png")
|
||||
assertContains(t, res[6], "seven.JPEG")
|
||||
assertContains(t, res[6], "d:")
|
||||
assertContains(t, res[7], "eight.png")
|
||||
assertContains(t, res[7], "c:")
|
||||
assertContains(t, res[8], "nine.png")
|
||||
assertContains(t, res[8], "d:")
|
||||
assertContains(t, res[9], "ten.PNG")
|
||||
assertContains(t, res[9], "E:")
|
||||
assertContains(t, res[10], "eleven.webp")
|
||||
assertContains(t, res[10], "c:")
|
||||
assertContains(t, res[11], "twelve.WebP")
|
||||
assertContains(t, res[11], "c:")
|
||||
assertContains(t, res[12], "thirteen.WEBP")
|
||||
assertContains(t, res[12], "d:")
|
||||
}
|
||||
|
||||
func TestExtractNamesDragDropPaths(t *testing.T) {
|
||||
input := `file:///Users/jdoe/Pictures/one.png file://localhost/C:/Users/jdoe/Pictures/two.webp file:///C:/Users/jdoe/Pictures/three.jpg .\relative\four.png`
|
||||
res := ExtractNames(input)
|
||||
if len(res) != 4 {
|
||||
t.Fatalf("len = %d, want 4", len(res))
|
||||
}
|
||||
assertContains(t, res[0], "file:///Users/jdoe/Pictures/one.png")
|
||||
assertContains(t, res[1], "file://localhost/C:/Users/jdoe/Pictures/two.webp")
|
||||
assertContains(t, res[2], "file:///C:/Users/jdoe/Pictures/three.jpg")
|
||||
assertContains(t, res[3], `.\relative\four.png`)
|
||||
}
|
||||
|
||||
func TestNormalizePathFileURL(t *testing.T) {
|
||||
got := NormalizePath("file:///C:/Users/jdoe/Pictures/img.png")
|
||||
want := filepath.FromSlash("C:/Users/jdoe/Pictures/img.png")
|
||||
if got != want {
|
||||
t.Fatalf("path = %q, want %q", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestExtractRemovesQuotedFilepath(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
fp := filepath.Join(dir, "img.jpg")
|
||||
data := make([]byte, 600)
|
||||
copy(data, []byte{
|
||||
0xff, 0xd8, 0xff, 0xe0, 0x00, 0x10, 'J', 'F', 'I', 'F',
|
||||
0x00, 0x01, 0x01, 0x01, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00,
|
||||
0xff, 0xd9,
|
||||
})
|
||||
if err := os.WriteFile(fp, data, 0o600); err != nil {
|
||||
t.Fatalf("failed to write test image: %v", err)
|
||||
}
|
||||
|
||||
input := "before '" + fp + "' after"
|
||||
cleaned, imgs, err := Extract(input)
|
||||
if err != nil {
|
||||
t.Fatalf("err: %v", err)
|
||||
}
|
||||
if len(imgs) != 1 {
|
||||
t.Fatalf("imgs = %d, want 1", len(imgs))
|
||||
}
|
||||
if cleaned != "before after" {
|
||||
t.Fatalf("cleaned = %q, want %q", cleaned, "before after")
|
||||
}
|
||||
}
|
||||
|
||||
func TestExtractFileURL(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
fp := filepath.Join(dir, "img.png")
|
||||
data := make([]byte, 600)
|
||||
copy(data, []byte{0x89, 'P', 'N', 'G', '\r', '\n', 0x1a, '\n'})
|
||||
if err := os.WriteFile(fp, data, 0o600); err != nil {
|
||||
t.Fatalf("failed to write test image: %v", err)
|
||||
}
|
||||
|
||||
fileURL := (&url.URL{Scheme: "file", Path: fp}).String()
|
||||
cleaned, imgs, err := Extract("before " + fileURL + " after")
|
||||
if err != nil {
|
||||
t.Fatalf("err: %v", err)
|
||||
}
|
||||
if len(imgs) != 1 {
|
||||
t.Fatalf("imgs = %d, want 1", len(imgs))
|
||||
}
|
||||
if cleaned != "before after" {
|
||||
t.Fatalf("cleaned = %q, want %q", cleaned, "before after")
|
||||
}
|
||||
}
|
||||
|
||||
func TestExtractWAV(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
fp := filepath.Join(dir, "sample.wav")
|
||||
data := make([]byte, 600)
|
||||
copy(data[:44], []byte{
|
||||
'R', 'I', 'F', 'F',
|
||||
0x58, 0x02, 0x00, 0x00,
|
||||
'W', 'A', 'V', 'E',
|
||||
'f', 'm', 't', ' ',
|
||||
0x10, 0x00, 0x00, 0x00,
|
||||
0x01, 0x00,
|
||||
0x01, 0x00,
|
||||
0x80, 0x3e, 0x00, 0x00,
|
||||
0x00, 0x7d, 0x00, 0x00,
|
||||
0x02, 0x00,
|
||||
0x10, 0x00,
|
||||
'd', 'a', 't', 'a',
|
||||
0x34, 0x02, 0x00, 0x00,
|
||||
})
|
||||
if err := os.WriteFile(fp, data, 0o600); err != nil {
|
||||
t.Fatalf("failed to write test audio: %v", err)
|
||||
}
|
||||
|
||||
input := "before " + fp + " after"
|
||||
cleaned, imgs, err := Extract(input)
|
||||
if err != nil {
|
||||
t.Fatalf("err: %v", err)
|
||||
}
|
||||
if len(imgs) != 1 {
|
||||
t.Fatalf("imgs = %d, want 1", len(imgs))
|
||||
}
|
||||
if cleaned != "before after" {
|
||||
t.Fatalf("cleaned = %q, want %q", cleaned, "before after")
|
||||
}
|
||||
}
|
||||
|
||||
func assertContains(t *testing.T, s, want string) {
|
||||
t.Helper()
|
||||
if !strings.Contains(s, want) {
|
||||
t.Fatalf("%q does not contain %q", s, want)
|
||||
}
|
||||
}
|
||||
|
||||
func assertNotContains(t *testing.T, s, want string) {
|
||||
t.Helper()
|
||||
if strings.Contains(s, want) {
|
||||
t.Fatalf("%q unexpectedly contains %q", s, want)
|
||||
}
|
||||
}
|
||||
|
||||
func assertNotContainsSlice(t *testing.T, ss []string, want string) {
|
||||
t.Helper()
|
||||
for _, s := range ss {
|
||||
if strings.Contains(s, want) {
|
||||
t.Fatalf("slice unexpectedly contains %q in %q", want, s)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -20,7 +20,7 @@ const (
|
||||
)
|
||||
|
||||
var (
|
||||
ErrPlanVerificationUnavailable = errors.New("Could not verify your plan. Try again in a moment.")
|
||||
ErrPlanVerificationUnavailable = errors.New("Could not verify Ollama plan. Try again in a moment or use a local model.")
|
||||
errUpgradeCancelled = errors.New("upgrade cancelled")
|
||||
)
|
||||
|
||||
@@ -247,7 +247,7 @@ func (c *launcherClient) ensureCloudModelAccess(ctx context.Context, model strin
|
||||
c.accountState = &state
|
||||
}
|
||||
if state.Status == accountStateUnknown {
|
||||
return ErrPlanVerificationUnavailable
|
||||
return nil
|
||||
}
|
||||
|
||||
if state.Status == accountStateSignedOut {
|
||||
@@ -259,7 +259,7 @@ func (c *launcherClient) ensureCloudModelAccess(ctx context.Context, model strin
|
||||
c.accountState = &state
|
||||
}
|
||||
if state.Status == accountStateUnknown {
|
||||
return ErrPlanVerificationUnavailable
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
+103
-11
@@ -7,6 +7,7 @@ import (
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"github.com/ollama/ollama/envconfig"
|
||||
)
|
||||
@@ -37,17 +38,21 @@ func (c *Claude) findPath() (string, error) {
|
||||
if runtime.GOOS == "windows" {
|
||||
name = "claude.exe"
|
||||
}
|
||||
fallback := filepath.Join(home, ".claude", "local", name)
|
||||
if _, err := os.Stat(fallback); err != nil {
|
||||
return "", err
|
||||
for _, fallback := range []string{
|
||||
filepath.Join(home, ".local", "bin", name),
|
||||
filepath.Join(home, ".claude", "local", name),
|
||||
} {
|
||||
if _, err := os.Stat(fallback); err == nil {
|
||||
return fallback, nil
|
||||
}
|
||||
}
|
||||
return fallback, nil
|
||||
return "", fmt.Errorf("claude binary not found")
|
||||
}
|
||||
|
||||
func (c *Claude) Run(model string, _ []LaunchModel, args []string) error {
|
||||
claudePath, err := c.findPath()
|
||||
claudePath, err := ensureClaudeInstalled()
|
||||
if err != nil {
|
||||
return fmt.Errorf("claude is not installed, install from https://code.claude.com/docs/en/quickstart")
|
||||
return err
|
||||
}
|
||||
|
||||
cmd := exec.Command(claudePath, c.args(model, args)...)
|
||||
@@ -55,17 +60,104 @@ func (c *Claude) Run(model string, _ []LaunchModel, args []string) error {
|
||||
cmd.Stdout = os.Stdout
|
||||
cmd.Stderr = os.Stderr
|
||||
|
||||
env := append(os.Environ(),
|
||||
"ANTHROPIC_BASE_URL="+envconfig.Host().String(),
|
||||
cmd.Env = append(os.Environ(), c.envVars(model)...)
|
||||
return cmd.Run()
|
||||
}
|
||||
|
||||
func (c *Claude) envVars(model string) []string {
|
||||
env := []string{
|
||||
"ANTHROPIC_BASE_URL=" + envconfig.Host().String(),
|
||||
"ANTHROPIC_API_KEY=",
|
||||
"ANTHROPIC_AUTH_TOKEN=ollama",
|
||||
"CLAUDE_CODE_ATTRIBUTION_HEADER=0",
|
||||
)
|
||||
"DISABLE_ERROR_REPORTING=1",
|
||||
"DISABLE_FEEDBACK_COMMAND=1",
|
||||
"CLAUDE_CODE_DISABLE_FEEDBACK_SURVEY=1",
|
||||
}
|
||||
|
||||
env = append(env, c.modelEnvVars(model)...)
|
||||
return env
|
||||
}
|
||||
|
||||
cmd.Env = env
|
||||
return cmd.Run()
|
||||
func ensureClaudeInstalled() (string, error) {
|
||||
if path, err := (&Claude{}).findPath(); err == nil {
|
||||
return path, nil
|
||||
}
|
||||
|
||||
if err := checkClaudeInstallerDependencies(); err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
ok, err := ConfirmPrompt("Claude Code is not installed. Install now?")
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if !ok {
|
||||
return "", fmt.Errorf("claude installation cancelled")
|
||||
}
|
||||
|
||||
bin, args, err := claudeInstallerCommand(runtime.GOOS)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
fmt.Fprintf(os.Stderr, "\nInstalling Claude Code...\n")
|
||||
cmd := exec.Command(bin, args...)
|
||||
cmd.Stdin = os.Stdin
|
||||
cmd.Stdout = os.Stdout
|
||||
cmd.Stderr = os.Stderr
|
||||
if err := cmd.Run(); err != nil {
|
||||
return "", fmt.Errorf("failed to install claude: %w", err)
|
||||
}
|
||||
|
||||
path, err := (&Claude{}).findPath()
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("claude was installed but the binary was not found on PATH\n\nYou may need to restart your shell")
|
||||
}
|
||||
|
||||
fmt.Fprintf(os.Stderr, "%sClaude Code installed successfully%s\n\n", ansiGreen, ansiReset)
|
||||
return path, nil
|
||||
}
|
||||
|
||||
func checkClaudeInstallerDependencies() error {
|
||||
switch runtime.GOOS {
|
||||
case "windows":
|
||||
if _, err := exec.LookPath("powershell"); err != nil {
|
||||
return fmt.Errorf("claude is not installed and required dependencies are missing\n\nInstall the following first:\n PowerShell: https://learn.microsoft.com/powershell/\n\nThen re-run:\n ollama launch claude")
|
||||
}
|
||||
default:
|
||||
var missing []string
|
||||
if _, err := exec.LookPath("curl"); err != nil {
|
||||
missing = append(missing, "curl: https://curl.se/")
|
||||
}
|
||||
if _, err := exec.LookPath("bash"); err != nil {
|
||||
missing = append(missing, "bash: https://www.gnu.org/software/bash/")
|
||||
}
|
||||
if len(missing) > 0 {
|
||||
return fmt.Errorf("claude is not installed and required dependencies are missing\n\nInstall the following first:\n %s\n\nThen re-run:\n ollama launch claude", strings.Join(missing, "\n "))
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func claudeInstallerCommand(goos string) (string, []string, error) {
|
||||
switch goos {
|
||||
case "windows":
|
||||
return "powershell", []string{
|
||||
"-NoProfile",
|
||||
"-ExecutionPolicy",
|
||||
"Bypass",
|
||||
"-Command",
|
||||
"irm https://claude.ai/install.ps1 | iex",
|
||||
}, nil
|
||||
case "darwin", "linux":
|
||||
return "bash", []string{
|
||||
"-c",
|
||||
"curl -fsSL https://claude.ai/install.sh | bash",
|
||||
}, nil
|
||||
default:
|
||||
return "", nil, fmt.Errorf("unsupported platform for claude install: %s", goos)
|
||||
}
|
||||
}
|
||||
|
||||
// modelEnvVars returns Claude Code env vars that route all model tiers through Ollama.
|
||||
|
||||
@@ -1,12 +1,15 @@
|
||||
package launch
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"slices"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/ollama/ollama/envconfig"
|
||||
)
|
||||
|
||||
func TestClaudeIntegration(t *testing.T) {
|
||||
@@ -67,6 +70,28 @@ func TestClaudeFindPath(t *testing.T) {
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("falls back to ~/.local/bin/claude", func(t *testing.T) {
|
||||
tmpDir := t.TempDir()
|
||||
setTestHome(t, tmpDir)
|
||||
t.Setenv("PATH", t.TempDir()) // empty dir, no claude binary
|
||||
|
||||
name := "claude"
|
||||
if runtime.GOOS == "windows" {
|
||||
name = "claude.exe"
|
||||
}
|
||||
fallback := filepath.Join(tmpDir, ".local", "bin", name)
|
||||
os.MkdirAll(filepath.Dir(fallback), 0o755)
|
||||
os.WriteFile(fallback, []byte("#!/bin/sh\n"), 0o755)
|
||||
|
||||
got, err := c.findPath()
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if got != fallback {
|
||||
t.Errorf("findPath() = %q, want %q", got, fallback)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("returns error when neither PATH nor fallback exists", func(t *testing.T) {
|
||||
tmpDir := t.TempDir()
|
||||
setTestHome(t, tmpDir)
|
||||
@@ -79,6 +104,210 @@ func TestClaudeFindPath(t *testing.T) {
|
||||
})
|
||||
}
|
||||
|
||||
func TestEnsureClaudeInstalled(t *testing.T) {
|
||||
withConfirm := func(t *testing.T, fn func(prompt string) (bool, error)) {
|
||||
t.Helper()
|
||||
oldConfirm := DefaultConfirmPrompt
|
||||
DefaultConfirmPrompt = func(prompt string, options ConfirmOptions) (bool, error) {
|
||||
return fn(prompt)
|
||||
}
|
||||
t.Cleanup(func() { DefaultConfirmPrompt = oldConfirm })
|
||||
}
|
||||
|
||||
t.Run("already installed", func(t *testing.T) {
|
||||
setTestHome(t, t.TempDir())
|
||||
tmpDir := t.TempDir()
|
||||
t.Setenv("PATH", tmpDir)
|
||||
writeFakeBinary(t, tmpDir, "claude")
|
||||
|
||||
withConfirm(t, func(prompt string) (bool, error) {
|
||||
t.Fatalf("did not expect prompt, got %q", prompt)
|
||||
return false, nil
|
||||
})
|
||||
|
||||
bin, err := ensureClaudeInstalled()
|
||||
if err != nil {
|
||||
t.Fatalf("ensureClaudeInstalled() error = %v", err)
|
||||
}
|
||||
if filepath.Base(bin) != "claude" && filepath.Base(bin) != "claude.cmd" {
|
||||
t.Fatalf("bin = %q, want claude binary", bin)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("missing dependencies", func(t *testing.T) {
|
||||
setTestHome(t, t.TempDir())
|
||||
t.Setenv("PATH", t.TempDir())
|
||||
|
||||
withConfirm(t, func(prompt string) (bool, error) {
|
||||
t.Fatalf("did not expect prompt, got %q", prompt)
|
||||
return false, nil
|
||||
})
|
||||
|
||||
_, err := ensureClaudeInstalled()
|
||||
if err == nil || !strings.Contains(err.Error(), "required dependencies are missing") {
|
||||
t.Fatalf("expected missing dependency error, got %v", err)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("missing and user declines install", func(t *testing.T) {
|
||||
setTestHome(t, t.TempDir())
|
||||
tmpDir := t.TempDir()
|
||||
t.Setenv("PATH", tmpDir)
|
||||
writeClaudeInstallerDeps(t, tmpDir)
|
||||
|
||||
withConfirm(t, func(prompt string) (bool, error) {
|
||||
if prompt != "Claude Code is not installed. Install now?" {
|
||||
t.Fatalf("unexpected prompt: %q", prompt)
|
||||
}
|
||||
return false, nil
|
||||
})
|
||||
|
||||
_, err := ensureClaudeInstalled()
|
||||
if err == nil || !strings.Contains(err.Error(), "installation cancelled") {
|
||||
t.Fatalf("expected cancellation error, got %v", err)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("missing and user confirms install succeeds", func(t *testing.T) {
|
||||
if runtime.GOOS == "windows" {
|
||||
t.Skip("uses POSIX shell fake binaries")
|
||||
}
|
||||
|
||||
homeDir := t.TempDir()
|
||||
setTestHome(t, homeDir)
|
||||
tmpDir := t.TempDir()
|
||||
t.Setenv("PATH", tmpDir)
|
||||
|
||||
writeFakeBinary(t, tmpDir, "curl")
|
||||
|
||||
installLog := filepath.Join(tmpDir, "bash.log")
|
||||
installedClaude := filepath.Join(homeDir, ".local", "bin", "claude")
|
||||
bashScript := fmt.Sprintf(`#!/bin/sh
|
||||
echo "$@" >> %q
|
||||
if [ "$1" = "-c" ]; then
|
||||
/bin/mkdir -p %q
|
||||
/bin/cat > %q <<'EOS'
|
||||
#!/bin/sh
|
||||
exit 0
|
||||
EOS
|
||||
/bin/chmod +x %q
|
||||
fi
|
||||
exit 0
|
||||
`, installLog, filepath.Dir(installedClaude), installedClaude, installedClaude)
|
||||
if err := os.WriteFile(filepath.Join(tmpDir, "bash"), []byte(bashScript), 0o755); err != nil {
|
||||
t.Fatalf("failed to write fake bash: %v", err)
|
||||
}
|
||||
|
||||
withConfirm(t, func(prompt string) (bool, error) {
|
||||
return true, nil
|
||||
})
|
||||
|
||||
bin, err := ensureClaudeInstalled()
|
||||
if err != nil {
|
||||
t.Fatalf("ensureClaudeInstalled() error = %v", err)
|
||||
}
|
||||
if bin != installedClaude {
|
||||
t.Fatalf("bin = %q, want %q", bin, installedClaude)
|
||||
}
|
||||
|
||||
logData, err := os.ReadFile(installLog)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to read install log: %v", err)
|
||||
}
|
||||
if !strings.Contains(string(logData), "https://claude.ai/install.sh") {
|
||||
t.Fatalf("expected install.sh command in log, got:\n%s", string(logData))
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("install command fails", func(t *testing.T) {
|
||||
if runtime.GOOS == "windows" {
|
||||
t.Skip("uses POSIX shell fake binaries")
|
||||
}
|
||||
|
||||
setTestHome(t, t.TempDir())
|
||||
tmpDir := t.TempDir()
|
||||
t.Setenv("PATH", tmpDir)
|
||||
writeFakeBinary(t, tmpDir, "curl")
|
||||
if err := os.WriteFile(filepath.Join(tmpDir, "bash"), []byte("#!/bin/sh\nexit 1\n"), 0o755); err != nil {
|
||||
t.Fatalf("failed to write fake bash: %v", err)
|
||||
}
|
||||
|
||||
withConfirm(t, func(prompt string) (bool, error) {
|
||||
return true, nil
|
||||
})
|
||||
|
||||
_, err := ensureClaudeInstalled()
|
||||
if err == nil || !strings.Contains(err.Error(), "failed to install claude") {
|
||||
t.Fatalf("expected install failure error, got %v", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func writeClaudeInstallerDeps(t *testing.T, dir string) {
|
||||
t.Helper()
|
||||
if runtime.GOOS == "windows" {
|
||||
writeFakeBinary(t, dir, "powershell")
|
||||
return
|
||||
}
|
||||
writeFakeBinary(t, dir, "curl")
|
||||
writeFakeBinary(t, dir, "bash")
|
||||
}
|
||||
|
||||
func TestClaudeInstallerCommand(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
goos string
|
||||
wantBin string
|
||||
want string
|
||||
wantErr string
|
||||
}{
|
||||
{
|
||||
name: "unix",
|
||||
goos: "linux",
|
||||
wantBin: "bash",
|
||||
want: "curl -fsSL https://claude.ai/install.sh | bash",
|
||||
},
|
||||
{
|
||||
name: "macos",
|
||||
goos: "darwin",
|
||||
wantBin: "bash",
|
||||
want: "curl -fsSL https://claude.ai/install.sh | bash",
|
||||
},
|
||||
{
|
||||
name: "windows",
|
||||
goos: "windows",
|
||||
wantBin: "powershell",
|
||||
want: "irm https://claude.ai/install.ps1 | iex",
|
||||
},
|
||||
{
|
||||
name: "unsupported",
|
||||
goos: "plan9",
|
||||
wantErr: "unsupported platform",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
bin, args, err := claudeInstallerCommand(tt.goos)
|
||||
if tt.wantErr != "" {
|
||||
if err == nil || !strings.Contains(err.Error(), tt.wantErr) {
|
||||
t.Fatalf("expected error containing %q, got %v", tt.wantErr, err)
|
||||
}
|
||||
return
|
||||
}
|
||||
if err != nil {
|
||||
t.Fatalf("claudeInstallerCommand() error = %v", err)
|
||||
}
|
||||
if bin != tt.wantBin {
|
||||
t.Fatalf("bin = %q, want %q", bin, tt.wantBin)
|
||||
}
|
||||
if !slices.Contains(args, tt.want) {
|
||||
t.Fatalf("args = %v, want command containing %q", args, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestClaudeArgs(t *testing.T) {
|
||||
c := &Claude{}
|
||||
|
||||
@@ -93,6 +322,7 @@ func TestClaudeArgs(t *testing.T) {
|
||||
{"with model and verbose", "llama3.2", []string{"--verbose"}, []string{"--model", "llama3.2", "--verbose"}},
|
||||
{"empty model with help", "", []string{"--help"}, []string{"--help"}},
|
||||
{"with allowed tools", "llama3.2", []string{"--allowedTools", "Read,Write,Bash"}, []string{"--model", "llama3.2", "--allowedTools", "Read,Write,Bash"}},
|
||||
{"with channels", "llama3.2", []string{"--channels", "plugin:telegram@claude-plugins-official"}, []string{"--model", "llama3.2", "--channels", "plugin:telegram@claude-plugins-official"}},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
@@ -105,6 +335,45 @@ func TestClaudeArgs(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestClaudeEnvVars(t *testing.T) {
|
||||
c := &Claude{}
|
||||
|
||||
envMap := func(envs []string) map[string]string {
|
||||
m := make(map[string]string)
|
||||
for _, e := range envs {
|
||||
k, v, _ := strings.Cut(e, "=")
|
||||
m[k] = v
|
||||
}
|
||||
return m
|
||||
}
|
||||
|
||||
got := envMap(c.envVars("llama3.2"))
|
||||
for key, want := range map[string]string{
|
||||
"ANTHROPIC_BASE_URL": envconfig.Host().String(),
|
||||
"ANTHROPIC_API_KEY": "",
|
||||
"ANTHROPIC_AUTH_TOKEN": "ollama",
|
||||
"CLAUDE_CODE_ATTRIBUTION_HEADER": "0",
|
||||
"DISABLE_ERROR_REPORTING": "1",
|
||||
"DISABLE_FEEDBACK_COMMAND": "1",
|
||||
"CLAUDE_CODE_DISABLE_FEEDBACK_SURVEY": "1",
|
||||
"ANTHROPIC_DEFAULT_OPUS_MODEL": "llama3.2",
|
||||
"ANTHROPIC_DEFAULT_SONNET_MODEL": "llama3.2",
|
||||
"ANTHROPIC_DEFAULT_HAIKU_MODEL": "llama3.2",
|
||||
"CLAUDE_CODE_SUBAGENT_MODEL": "llama3.2",
|
||||
} {
|
||||
if got[key] != want {
|
||||
t.Errorf("%s = %q, want %q", key, got[key], want)
|
||||
}
|
||||
}
|
||||
|
||||
// Both variables disable Claude Code feature-flag evaluation, which keeps Channels unavailable.
|
||||
for _, key := range []string{"CLAUDE_CODE_DISABLE_NONESSENTIAL_TRAFFIC", "DISABLE_TELEMETRY"} {
|
||||
if _, ok := got[key]; ok {
|
||||
t.Errorf("%s must not be set by Ollama", key)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestClaudeModelEnvVars(t *testing.T) {
|
||||
c := &Claude{}
|
||||
|
||||
|
||||
+167
-42
@@ -2,6 +2,7 @@ package launch
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"os"
|
||||
"os/exec"
|
||||
@@ -16,13 +17,14 @@ import (
|
||||
)
|
||||
|
||||
const (
|
||||
chatGPTIntegrationName = "chatgpt"
|
||||
codexAppIntegrationName = "codex-app"
|
||||
codexAppProfileName = "ollama-launch-codex-app"
|
||||
codexAppBundleID = "com.openai.codex"
|
||||
codexAppModelCatalogFilename = "ollama-launch-models.json"
|
||||
codexAppRestoreHint = "To restore your usual Codex profile, run: ollama launch codex-app --restore"
|
||||
codexAppConfigurationSuccess = "Codex App profile changed to Ollama."
|
||||
codexAppRestoreSuccess = "Codex App restored to your usual profile."
|
||||
codexAppRestoreHint = "To restore your usual ChatGPT profile, run: ollama launch chatgpt --restore"
|
||||
codexAppConfigurationSuccess = "ChatGPT profile changed to Ollama."
|
||||
codexAppRestoreSuccess = "ChatGPT restored to your usual profile."
|
||||
)
|
||||
|
||||
var (
|
||||
@@ -49,7 +51,7 @@ var (
|
||||
// model while leaving model discovery and switching to Codex's Ollama provider.
|
||||
type CodexApp struct{}
|
||||
|
||||
func (c *CodexApp) String() string { return "Codex App" }
|
||||
func (c *CodexApp) String() string { return "ChatGPT" }
|
||||
|
||||
func (c *CodexApp) Supported() error { return codexAppSupported() }
|
||||
|
||||
@@ -68,7 +70,7 @@ func (c *CodexApp) Configure(model string) error {
|
||||
func (c *CodexApp) ConfigureWithModels(primary string, models []LaunchModel) error {
|
||||
primary = strings.TrimSpace(primary)
|
||||
if primary == "" {
|
||||
return fmt.Errorf("codex-app requires a model")
|
||||
return fmt.Errorf("chatgpt requires a model")
|
||||
}
|
||||
|
||||
configPath, err := codexConfigPath()
|
||||
@@ -106,7 +108,10 @@ func (c *CodexApp) CurrentModel() string {
|
||||
if parsed.RootString(codexRootModelProviderKey) == profileName {
|
||||
baseURL := parsed.ProviderString(profileName, "base_url")
|
||||
if codexNormalizeURL(baseURL) == codexNormalizeURL(codexBaseURL()) && codexAppCatalogHealthy(parsed, profileName) {
|
||||
return strings.TrimSpace(parsed.RootString(codexRootModelKey))
|
||||
model := strings.TrimSpace(parsed.RootString(codexRootModelKey))
|
||||
if codexAppCatalogContainsModel(model) {
|
||||
return model
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -125,7 +130,11 @@ func (c *CodexApp) CurrentModel() string {
|
||||
if !codexAppCatalogHealthy(parsed, profileName) {
|
||||
return ""
|
||||
}
|
||||
return strings.TrimSpace(parsed.ProfileString(profileName, codexRootModelKey))
|
||||
model := strings.TrimSpace(parsed.ProfileString(profileName, codexRootModelKey))
|
||||
if !codexAppCatalogContainsModel(model) {
|
||||
return ""
|
||||
}
|
||||
return model
|
||||
}
|
||||
|
||||
func codexAppManagedProfileNames() []string {
|
||||
@@ -169,6 +178,40 @@ func codexAppCatalogHealthy(config codexParsedConfig, profileName string) bool {
|
||||
return len(catalog.Models) > 0
|
||||
}
|
||||
|
||||
// codexAppCatalogContainsModel reports whether model appears as a slug in the
|
||||
// Ollama-managed model catalog. When the configured model is not in the catalog
|
||||
// the user has drifted away from the launch-managed model (e.g. by selecting a
|
||||
// built-in OpenAI model in the Codex App UI), and the launch config should be
|
||||
// treated as inactive.
|
||||
func codexAppCatalogContainsModel(model string) bool {
|
||||
if strings.TrimSpace(model) == "" {
|
||||
return false
|
||||
}
|
||||
catalogPath, err := codexAppModelCatalogPath()
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
data, err := os.ReadFile(catalogPath)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
var catalog struct {
|
||||
Models []struct {
|
||||
Slug string `json:"slug"`
|
||||
} `json:"models"`
|
||||
}
|
||||
if err := json.Unmarshal(data, &catalog); err != nil {
|
||||
return false
|
||||
}
|
||||
target := codexAppCatalogModelKey(model)
|
||||
for _, m := range catalog.Models {
|
||||
if codexAppCatalogModelKey(m.Slug) == target {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func writeCodexAppConfig(configPath, model, modelCatalogPath string) error {
|
||||
baseURL := codexBaseURL()
|
||||
|
||||
@@ -209,10 +252,10 @@ func writeCodexAppConfig(configPath, model, modelCatalogPath string) error {
|
||||
|
||||
func codexValidateAppConfigText(config codexParsedConfig, model, modelCatalogPath, baseURL string) error {
|
||||
if got, ok := config.RootStringOK(codexRootProfileKey); ok {
|
||||
return fmt.Errorf("generated Codex App config still contains legacy profile = %q", got)
|
||||
return fmt.Errorf("generated ChatGPT config still contains legacy profile = %q", got)
|
||||
}
|
||||
if config.Exists("profiles", codexAppProfileName) {
|
||||
return fmt.Errorf("generated Codex App config still contains legacy profiles.%s table", codexAppProfileName)
|
||||
return fmt.Errorf("generated ChatGPT config still contains legacy profiles.%s table", codexAppProfileName)
|
||||
}
|
||||
for _, check := range []struct {
|
||||
path []string
|
||||
@@ -226,14 +269,14 @@ func codexValidateAppConfigText(config codexParsedConfig, model, modelCatalogPat
|
||||
{[]string{"model_providers", codexAppProfileName, "wire_api"}, "responses"},
|
||||
} {
|
||||
if got, ok := config.String(check.path...); !ok || got != check.want {
|
||||
return fmt.Errorf("generated Codex App config missing %s = %q", strings.Join(check.path, "."), check.want)
|
||||
return fmt.Errorf("generated ChatGPT config missing %s = %q", strings.Join(check.path, "."), check.want)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *CodexApp) Onboard() error {
|
||||
return config.MarkIntegrationOnboarded(codexAppIntegrationName)
|
||||
return config.MarkIntegrationOnboarded(chatGPTIntegrationName)
|
||||
}
|
||||
|
||||
func (c *CodexApp) RequiresInteractiveOnboarding() bool {
|
||||
@@ -257,9 +300,9 @@ func (c *CodexApp) Run(_ string, _ []LaunchModel, args []string) error {
|
||||
return err
|
||||
}
|
||||
if len(args) > 0 {
|
||||
return fmt.Errorf("codex-app does not accept extra arguments")
|
||||
return fmt.Errorf("chatgpt does not accept extra arguments")
|
||||
}
|
||||
return codexAppLaunchOrRestart("Restart Codex to use Ollama?", nil)
|
||||
return codexAppLaunchOrRestart("Restart ChatGPT to use Ollama?", nil)
|
||||
}
|
||||
|
||||
func (c *CodexApp) Restore() error {
|
||||
@@ -283,7 +326,7 @@ func (c *CodexApp) Restore() error {
|
||||
if err := codexAppRemoveOwnedCatalog(); err != nil {
|
||||
return codexAppRestoreFailure(configPath, err)
|
||||
}
|
||||
return codexAppLaunchOrRestart("Restart Codex to use your usual profile?", nil)
|
||||
return codexAppLaunchOrRestart("Restart ChatGPT to use your usual profile?", nil)
|
||||
}
|
||||
return codexAppRestoreFailure(configPath, err)
|
||||
}
|
||||
@@ -319,11 +362,11 @@ func (c *CodexApp) Restore() error {
|
||||
if err := removeCodexAppRestoreState(); err != nil {
|
||||
return codexAppRestoreFailure(configPath, err)
|
||||
}
|
||||
return codexAppLaunchOrRestart("Restart Codex to use your usual profile?", nil)
|
||||
return codexAppLaunchOrRestart("Restart ChatGPT to use your usual profile?", nil)
|
||||
}
|
||||
|
||||
func codexAppRestoreFailure(configPath string, err error) error {
|
||||
return fmt.Errorf("restore Codex App config: %w\n\nRestore did not complete. Check these files before retrying:\n Codex config: %s\n Restore state: %s\n Model catalog: %s\n Backups: %s",
|
||||
return fmt.Errorf("restore ChatGPT config: %w\n\nRestore did not complete. Check these files before retrying:\n Codex config: %s\n Restore state: %s\n Model catalog: %s\n Backups: %s",
|
||||
err,
|
||||
configPath,
|
||||
codexAppRestoreStatePath(),
|
||||
@@ -337,7 +380,7 @@ func codexAppSupported() error {
|
||||
case "darwin", "windows":
|
||||
return nil
|
||||
default:
|
||||
return fmt.Errorf("Codex App launch is only supported on macOS and Windows")
|
||||
return fmt.Errorf("ChatGPT launch is only supported on macOS and Windows")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -381,7 +424,7 @@ func codexAppModelCatalogPathForConfig(configPath string) string {
|
||||
|
||||
func writeCodexAppModelCatalog(path, primary string, models []LaunchModel) error {
|
||||
if len(models) == 0 {
|
||||
return fmt.Errorf("codex-app model catalog cannot be empty")
|
||||
return fmt.Errorf("chatgpt model catalog cannot be empty")
|
||||
}
|
||||
|
||||
baseInstructions := codexAppBaseInstructions()
|
||||
@@ -544,9 +587,12 @@ func codexAppAppPath() string {
|
||||
}
|
||||
|
||||
func codexAppDarwinAppCandidates() []string {
|
||||
candidates := []string{"/Applications/Codex.app"}
|
||||
candidates := []string{"/Applications/ChatGPT.app", "/Applications/Codex.app"}
|
||||
if home, err := os.UserHomeDir(); err == nil {
|
||||
candidates = append(candidates, filepath.Join(home, "Applications", "Codex.app"))
|
||||
candidates = append(candidates,
|
||||
filepath.Join(home, "Applications", "ChatGPT.app"),
|
||||
filepath.Join(home, "Applications", "Codex.app"),
|
||||
)
|
||||
}
|
||||
return candidates
|
||||
}
|
||||
@@ -558,6 +604,11 @@ func codexAppWindowsAppCandidates() []string {
|
||||
}
|
||||
|
||||
candidates := []string{
|
||||
filepath.Join(local, "Programs", "ChatGPT", "ChatGPT.exe"),
|
||||
filepath.Join(local, "Programs", "OpenAI ChatGPT", "ChatGPT.exe"),
|
||||
filepath.Join(local, "ChatGPT", "ChatGPT.exe"),
|
||||
filepath.Join(local, "OpenAI ChatGPT", "ChatGPT.exe"),
|
||||
filepath.Join(local, "OpenAI", "ChatGPT", "ChatGPT.exe"),
|
||||
filepath.Join(local, "Programs", "Codex", "Codex.exe"),
|
||||
filepath.Join(local, "Programs", "OpenAI Codex", "Codex.exe"),
|
||||
filepath.Join(local, "Codex", "Codex.exe"),
|
||||
@@ -566,6 +617,11 @@ func codexAppWindowsAppCandidates() []string {
|
||||
filepath.Join(local, "openai-codex-electron", "Codex.exe"),
|
||||
}
|
||||
for _, pattern := range []string{
|
||||
filepath.Join(local, "Programs", "ChatGPT", "app-*", "ChatGPT.exe"),
|
||||
filepath.Join(local, "Programs", "OpenAI ChatGPT", "app-*", "ChatGPT.exe"),
|
||||
filepath.Join(local, "ChatGPT", "app-*", "ChatGPT.exe"),
|
||||
filepath.Join(local, "OpenAI ChatGPT", "app-*", "ChatGPT.exe"),
|
||||
filepath.Join(local, "OpenAI", "ChatGPT", "app-*", "ChatGPT.exe"),
|
||||
filepath.Join(local, "Programs", "Codex", "app-*", "Codex.exe"),
|
||||
filepath.Join(local, "Programs", "OpenAI Codex", "app-*", "Codex.exe"),
|
||||
filepath.Join(local, "Codex", "app-*", "Codex.exe"),
|
||||
@@ -628,22 +684,54 @@ func codexAppLaunchOrRestart(prompt string, launchArgs []string) error {
|
||||
return err
|
||||
}
|
||||
if !restart {
|
||||
fmt.Fprintln(os.Stderr, "\nQuit and reopen Codex when you're ready for the profile change to take effect.")
|
||||
fmt.Fprintln(os.Stderr, "\nQuit and reopen ChatGPT when you're ready for the profile change to take effect.")
|
||||
return nil
|
||||
}
|
||||
|
||||
if err := codexAppQuitApp(); err != nil {
|
||||
return fmt.Errorf("quit Codex: %w", err)
|
||||
// A single spinner and cancellation channel span the entire restart flow
|
||||
// (quit, wait, force-quit, wait, reopen) so that one Ctrl+C aborts the
|
||||
// whole sequence rather than just the currently-active wait. The bubbletea
|
||||
// spinner closes Cancelled() from its raw-mode Ctrl+C handler; the ANSI
|
||||
// fallback relies on SIGINT terminating the process directly.
|
||||
sp := StartSpinner(codexAppRestartMessage)
|
||||
defer sp.Stop()
|
||||
cancelled := sp.Cancelled()
|
||||
isCancelled := func() bool {
|
||||
if cancelled == nil {
|
||||
return false
|
||||
}
|
||||
select {
|
||||
case <-cancelled:
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
if err := codexAppQuitApp(); err != nil {
|
||||
return fmt.Errorf("quit ChatGPT: %w", err)
|
||||
}
|
||||
if isCancelled() {
|
||||
return ErrCancelled
|
||||
}
|
||||
gracefulErr := waitForCodexAppGracefulExit(codexAppExitTimeout, cancelled)
|
||||
if isCancelled() {
|
||||
return ErrCancelled
|
||||
}
|
||||
if errors.Is(gracefulErr, ErrCancelled) {
|
||||
return gracefulErr
|
||||
}
|
||||
gracefulErr := waitForCodexAppGracefulExit(codexAppExitTimeout)
|
||||
if gracefulErr != nil && !codexAppForceQuitSupported() {
|
||||
return gracefulErr
|
||||
}
|
||||
if codexAppForceQuitSupported() && codexAppIsRunning() {
|
||||
if forceErr := codexAppForceQuit(); forceErr != nil {
|
||||
return fmt.Errorf("force stop Codex: %w", forceErr)
|
||||
if isCancelled() {
|
||||
return ErrCancelled
|
||||
}
|
||||
if err := waitForCodexAppExit(codexAppForceExitTimeout); err != nil {
|
||||
if forceErr := codexAppForceQuit(); forceErr != nil {
|
||||
return fmt.Errorf("force stop ChatGPT: %w", forceErr)
|
||||
}
|
||||
if err := waitForCodexAppExit(codexAppForceExitTimeout, cancelled); err != nil {
|
||||
return err
|
||||
}
|
||||
} else if gracefulErr != nil {
|
||||
@@ -651,6 +739,10 @@ func codexAppLaunchOrRestart(prompt string, launchArgs []string) error {
|
||||
return gracefulErr
|
||||
}
|
||||
}
|
||||
if isCancelled() {
|
||||
return ErrCancelled
|
||||
}
|
||||
sp.Stop()
|
||||
if restartAppID != "" {
|
||||
return codexAppOpenStart(restartAppID)
|
||||
}
|
||||
@@ -664,8 +756,8 @@ func codexAppForceQuitSupported() bool {
|
||||
return codexAppGOOS == "darwin" || codexAppGOOS == "windows"
|
||||
}
|
||||
|
||||
func waitForCodexAppGracefulExit(timeout time.Duration) error {
|
||||
return waitForCodexAppCondition(timeout, func() bool {
|
||||
func waitForCodexAppGracefulExit(timeout time.Duration, cancel <-chan struct{}) error {
|
||||
return waitForCodexAppCondition(timeout, cancel, func() bool {
|
||||
if codexAppGOOS == "windows" {
|
||||
return !codexAppHasWindow()
|
||||
}
|
||||
@@ -673,21 +765,41 @@ func waitForCodexAppGracefulExit(timeout time.Duration) error {
|
||||
})
|
||||
}
|
||||
|
||||
func waitForCodexAppExit(timeout time.Duration) error {
|
||||
return waitForCodexAppCondition(timeout, func() bool {
|
||||
func waitForCodexAppExit(timeout time.Duration, cancel <-chan struct{}) error {
|
||||
return waitForCodexAppCondition(timeout, cancel, func() bool {
|
||||
return !codexAppIsRunning()
|
||||
})
|
||||
}
|
||||
|
||||
func waitForCodexAppCondition(timeout time.Duration, done func() bool) error {
|
||||
// codexAppRestartMessage is the label shown next to the animated spinner while
|
||||
// the ChatGPT desktop app is quitting before being reopened.
|
||||
const codexAppRestartMessage = "Restarting ChatGPT..."
|
||||
|
||||
// waitForCodexAppCondition polls done at a 200ms cadence until it reports the
|
||||
// app has exited or timeout elapses. It watches cancel (closed by the spinner
|
||||
// when the user hits Ctrl+C) and returns ErrCancelled if the flow is aborted.
|
||||
// The spinner itself is owned by the caller so a single spinner spans the
|
||||
// whole restart sequence. When timeout is zero the loop never runs, so
|
||||
// force-quit paths that short-circuit the graceful wait return immediately.
|
||||
func waitForCodexAppCondition(timeout time.Duration, cancel <-chan struct{}, done func() bool) error {
|
||||
deadline := time.Now().Add(timeout)
|
||||
for time.Now().Before(deadline) {
|
||||
if cancel != nil {
|
||||
select {
|
||||
case <-cancel:
|
||||
return ErrCancelled
|
||||
default:
|
||||
}
|
||||
}
|
||||
if done() {
|
||||
return nil
|
||||
}
|
||||
codexAppSleep(200 * time.Millisecond)
|
||||
}
|
||||
return fmt.Errorf("Codex did not quit; quit it manually and re-run the command")
|
||||
if done() {
|
||||
return nil
|
||||
}
|
||||
return fmt.Errorf("ChatGPT did not quit; quit it manually and re-run the command")
|
||||
}
|
||||
|
||||
func defaultCodexAppOpenApp(args []string) error {
|
||||
@@ -710,7 +822,7 @@ func defaultCodexAppOpenApp(args []string) error {
|
||||
if appID := codexAppStartID(); appID != "" {
|
||||
return codexAppOpenStart(appID)
|
||||
}
|
||||
return fmt.Errorf("Codex executable was not found; open Codex manually once and re-run 'ollama launch codex-app'")
|
||||
return fmt.Errorf("ChatGPT was not found; install it from https://chatgpt.com/download, then re-run 'ollama launch chatgpt'")
|
||||
case "darwin":
|
||||
if path := codexAppAppPath(); path != "" {
|
||||
cmd := exec.Command("open", path)
|
||||
@@ -747,14 +859,17 @@ func defaultCodexAppOpenStartAppID(appID string) error {
|
||||
|
||||
func defaultCodexAppQuitApp() error {
|
||||
if codexAppGOOS == "windows" {
|
||||
script := `Get-Process Codex -ErrorAction SilentlyContinue | Where-Object { $_.MainWindowHandle -ne 0 } | ForEach-Object { [void]$_.CloseMainWindow() }`
|
||||
script := `Get-Process ChatGPT,Codex -ErrorAction SilentlyContinue | Where-Object { $_.MainWindowHandle -ne 0 } | ForEach-Object { [void]$_.CloseMainWindow() }`
|
||||
return exec.Command("powershell.exe", "-NoProfile", "-Command", script).Run()
|
||||
}
|
||||
|
||||
scriptErr := exec.Command("osascript", "-e", `tell application "Codex" to quit`).Run()
|
||||
scriptErr := exec.Command("osascript", "-e", `tell application "ChatGPT" to quit`).Run()
|
||||
if scriptErr != nil {
|
||||
scriptErr = exec.Command("osascript", "-e", `tell application id "`+codexAppBundleID+`" to quit`).Run()
|
||||
}
|
||||
if scriptErr != nil {
|
||||
scriptErr = exec.Command("osascript", "-e", `tell application "Codex" to quit`).Run()
|
||||
}
|
||||
return scriptErr
|
||||
}
|
||||
|
||||
@@ -793,7 +908,7 @@ func defaultCodexAppHasOpenWindow() bool {
|
||||
if codexAppGOOS != "windows" {
|
||||
return codexAppIsRunning()
|
||||
}
|
||||
script := `(Get-Process Codex -ErrorAction SilentlyContinue | Where-Object { $_.MainWindowHandle -ne 0 } | Select-Object -First 1).Id`
|
||||
script := `(Get-Process ChatGPT,Codex -ErrorAction SilentlyContinue | Where-Object { $_.MainWindowHandle -ne 0 } | Select-Object -First 1).Id`
|
||||
out, err := exec.Command("powershell.exe", "-NoProfile", "-Command", script).Output()
|
||||
return err == nil && strings.TrimSpace(string(out)) != ""
|
||||
}
|
||||
@@ -803,7 +918,11 @@ func defaultCodexAppIsRunning() bool {
|
||||
case "windows":
|
||||
return len(codexAppMatchingProcessIDs()) > 0
|
||||
case "darwin":
|
||||
out, err := exec.Command("osascript", "-e", `tell application "System Events" to exists process "Codex"`).Output()
|
||||
out, err := exec.Command("osascript", "-e", `tell application "System Events" to exists process "ChatGPT"`).Output()
|
||||
if err == nil && strings.TrimSpace(string(out)) == "true" {
|
||||
return true
|
||||
}
|
||||
out, err = exec.Command("osascript", "-e", `tell application "System Events" to exists process "Codex"`).Output()
|
||||
if err == nil && strings.TrimSpace(string(out)) == "true" {
|
||||
return true
|
||||
}
|
||||
@@ -845,7 +964,7 @@ func codexAppMatchingProcessIDs() []int {
|
||||
}
|
||||
|
||||
func codexAppWindowsMatchingProcessIDs() []int {
|
||||
script := fmt.Sprintf(`$current = %d; Get-CimInstance Win32_Process -Filter "Name = 'Codex.exe' OR Name = 'codex.exe'" | Where-Object { $_.ProcessId -ne $current -and ((($_.Name -ieq 'Codex.exe') -and (($null -eq $_.CommandLine) -or ($_.CommandLine -notlike '* --type=*'))) -or (($_.Name -ieq 'codex.exe') -and ($_.CommandLine -like '*app-server*'))) } | Select-Object -ExpandProperty ProcessId`, os.Getpid())
|
||||
script := fmt.Sprintf(`$current = %d; Get-CimInstance Win32_Process -Filter "Name = 'Codex.exe' OR Name = 'codex.exe' OR Name = 'ChatGPT.exe' OR Name = 'chatgpt.exe'" | Where-Object { $_.ProcessId -ne $current -and ((($_.Name -ieq 'Codex.exe' -or $_.Name -ieq 'ChatGPT.exe') -and (($null -eq $_.CommandLine) -or ($_.CommandLine -notlike '* --type=*'))) -or ((($_.Name -ieq 'codex.exe') -or ($_.Name -ieq 'chatgpt.exe')) -and ($_.CommandLine -like '*app-server*'))) } | Select-Object -ExpandProperty ProcessId`, os.Getpid())
|
||||
out, err := exec.Command("powershell.exe", "-NoProfile", "-Command", script).Output()
|
||||
if err != nil {
|
||||
return nil
|
||||
@@ -865,7 +984,7 @@ func defaultCodexAppRunningAppPath() string {
|
||||
if codexAppGOOS != "windows" {
|
||||
return ""
|
||||
}
|
||||
script := `(Get-Process Codex -ErrorAction SilentlyContinue | Where-Object { $_.MainWindowHandle -ne 0 -and $_.Path } | Select-Object -First 1 -ExpandProperty Path)`
|
||||
script := `(Get-Process ChatGPT,Codex -ErrorAction SilentlyContinue | Where-Object { $_.MainWindowHandle -ne 0 -and $_.Path } | Select-Object -First 1 -ExpandProperty Path)`
|
||||
out, err := exec.Command("powershell.exe", "-NoProfile", "-Command", script).Output()
|
||||
if err != nil {
|
||||
return ""
|
||||
@@ -877,7 +996,7 @@ func defaultCodexAppStartAppID() string {
|
||||
if codexAppGOOS != "windows" {
|
||||
return ""
|
||||
}
|
||||
script := `(Get-StartApps Codex | Where-Object { $_.Name -eq 'Codex' -or $_.Name -like 'Codex*' } | Select-Object -First 1 -ExpandProperty AppID)`
|
||||
script := `(Get-StartApps | Where-Object { $_.Name -eq 'ChatGPT' -or $_.Name -like 'ChatGPT*' -or $_.Name -eq 'Codex' -or $_.Name -like 'Codex*' } | Select-Object -First 1 -ExpandProperty AppID)`
|
||||
out, err := exec.Command("powershell.exe", "-NoProfile", "-Command", script).Output()
|
||||
if err != nil {
|
||||
return ""
|
||||
@@ -895,7 +1014,7 @@ func defaultCodexAppCanOpenBundleID() bool {
|
||||
}
|
||||
|
||||
func codexAppProcessMatches(command string) bool {
|
||||
if strings.Contains(command, `\Codex.exe`) && strings.Contains(command, " --type=") {
|
||||
if (strings.Contains(command, `\Codex.exe`) || strings.Contains(command, `\ChatGPT.exe`)) && strings.Contains(command, " --type=") {
|
||||
return false
|
||||
}
|
||||
for _, pattern := range codexAppProcessPatterns() {
|
||||
@@ -908,8 +1027,14 @@ func codexAppProcessMatches(command string) bool {
|
||||
|
||||
func codexAppProcessPatterns() []string {
|
||||
return []string{
|
||||
"ChatGPT.app/Contents/MacOS/ChatGPT",
|
||||
"ChatGPT.app/Contents/Resources/codex app-server",
|
||||
"Codex.app/Contents/MacOS/Codex",
|
||||
"Codex.app/Contents/Resources/codex app-server",
|
||||
`\ChatGPT.exe`,
|
||||
`resources\chatgpt.exe app-server`,
|
||||
`resources\chatgpt.exe" app-server`,
|
||||
`resources\chatgpt.exe" "app-server`,
|
||||
`\Codex.exe`,
|
||||
`resources\codex.exe app-server`,
|
||||
`resources\codex.exe" app-server`,
|
||||
|
||||
@@ -2,9 +2,11 @@ package launch
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"slices"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
@@ -157,6 +159,27 @@ func TestCodexAppInstalledUsesMacBundleIDFallback(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestChatGPTMissingAppGivesDownloadRecovery(t *testing.T) {
|
||||
withCodexAppPlatform(t, "darwin")
|
||||
|
||||
oldCanOpenID := codexAppCanOpenID
|
||||
oldStat := codexAppStat
|
||||
codexAppCanOpenID = func() bool { return false }
|
||||
codexAppStat = func(string) (os.FileInfo, error) { return nil, os.ErrNotExist }
|
||||
t.Cleanup(func() {
|
||||
codexAppCanOpenID = oldCanOpenID
|
||||
codexAppStat = oldStat
|
||||
})
|
||||
|
||||
err := EnsureIntegrationInstalled(chatGPTIntegrationName, &CodexApp{})
|
||||
if err == nil {
|
||||
t.Fatal("expected missing ChatGPT install error")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "chatgpt is not installed") || !strings.Contains(err.Error(), "https://chatgpt.com/download") {
|
||||
t.Fatalf("missing-app error = %q, want ChatGPT download recovery", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCodexAppConfigureActivatesOllamaProviderWithoutLegacyProfile(t *testing.T) {
|
||||
tmpDir := t.TempDir()
|
||||
setTestHome(t, tmpDir)
|
||||
@@ -291,6 +314,58 @@ func TestCodexAppConfigureUsesAppSpecificProfileWithoutTouchingCLIProfile(t *tes
|
||||
assertBackupContains(t, filepath.Join(fileutil.BackupDir(), codexAppIntegrationName, "config.toml.*"), `profile = "default"`)
|
||||
}
|
||||
|
||||
func TestCodexAppConfigureIsIdempotentAndPreservesUnrelatedProvider(t *testing.T) {
|
||||
tmpDir := t.TempDir()
|
||||
setTestHome(t, tmpDir)
|
||||
t.Setenv("OLLAMA_HOST", "http://127.0.0.1:9999")
|
||||
|
||||
configPath := filepath.Join(tmpDir, ".codex", "config.toml")
|
||||
if err := os.MkdirAll(filepath.Dir(configPath), 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
existing := "[model_providers.custom]\n" +
|
||||
`name = "Custom"` + "\n" +
|
||||
`base_url = "https://example.invalid/v1"` + "\n" +
|
||||
`env_key = "CUSTOM_API_KEY"` + "\n"
|
||||
if err := os.WriteFile(configPath, []byte(existing), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
app := &CodexApp{}
|
||||
models := testLaunchModels("llama3.2", "qwen3:8b")
|
||||
if err := app.ConfigureWithModels("llama3.2", models); err != nil {
|
||||
t.Fatalf("first ConfigureWithModels returned error: %v", err)
|
||||
}
|
||||
first, err := os.ReadFile(configPath)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := app.ConfigureWithModels("llama3.2", models); err != nil {
|
||||
t.Fatalf("second ConfigureWithModels returned error: %v", err)
|
||||
}
|
||||
second, err := os.ReadFile(configPath)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if string(second) != string(first) {
|
||||
t.Fatalf("rerun changed generated config:\nfirst:\n%s\nsecond:\n%s", first, second)
|
||||
}
|
||||
|
||||
parsed, err := codexParseConfig(string(second))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if got := parsed.ProviderString("custom", "env_key"); got != "CUSTOM_API_KEY" {
|
||||
t.Fatalf("custom provider env_key = %q, want preserved value", got)
|
||||
}
|
||||
if parsed.Exists("model_providers", codexAppProfileName, "env_key") {
|
||||
t.Fatalf("managed local Ollama provider should not require an API key:\n%s", second)
|
||||
}
|
||||
if got := app.CurrentModel(); got != "llama3.2" {
|
||||
t.Fatalf("CurrentModel = %q, want llama3.2", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCodexCLIConfigRefreshLeavesCodexAppConfigActive(t *testing.T) {
|
||||
tmpDir := t.TempDir()
|
||||
setTestHome(t, tmpDir)
|
||||
@@ -626,6 +701,60 @@ func TestCodexAppCurrentModelRequiresHealthyCatalog(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestCodexAppCurrentModelDetectsDriftedModel(t *testing.T) {
|
||||
tmpDir := t.TempDir()
|
||||
setTestHome(t, tmpDir)
|
||||
t.Setenv("OLLAMA_HOST", "http://127.0.0.1:11434")
|
||||
|
||||
catalogPath := mustWriteCodexAppTestCatalog(t, "llama3.2")
|
||||
configPath := filepath.Join(tmpDir, ".codex", "config.toml")
|
||||
if err := os.MkdirAll(filepath.Dir(configPath), 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
content := "" +
|
||||
`model = "gpt-5.5"` + "\n" +
|
||||
fmt.Sprintf(`model_provider = %q`, codexAppProfileName) + "\n\n" +
|
||||
fmt.Sprintf(`model_catalog_json = %q`, catalogPath) + "\n\n" +
|
||||
codexProviderHeaderFor(codexAppProfileName) + "\n" +
|
||||
`name = "Ollama"` + "\n" +
|
||||
`base_url = "http://127.0.0.1:11434/v1/"` + "\n" +
|
||||
`wire_api = "responses"` + "\n"
|
||||
if err := os.WriteFile(configPath, []byte(content), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
if got := (&CodexApp{}).CurrentModel(); got != "" {
|
||||
t.Fatalf("CurrentModel = %q, want empty when model has drifted from the Ollama catalog", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCodexAppCurrentModelAcceptsLatestSuffixDrift(t *testing.T) {
|
||||
tmpDir := t.TempDir()
|
||||
setTestHome(t, tmpDir)
|
||||
t.Setenv("OLLAMA_HOST", "http://127.0.0.1:11434")
|
||||
|
||||
catalogPath := mustWriteCodexAppTestCatalog(t, "llama3.2")
|
||||
configPath := filepath.Join(tmpDir, ".codex", "config.toml")
|
||||
if err := os.MkdirAll(filepath.Dir(configPath), 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
content := "" +
|
||||
`model = "llama3.2:latest"` + "\n" +
|
||||
fmt.Sprintf(`model_provider = %q`, codexAppProfileName) + "\n\n" +
|
||||
fmt.Sprintf(`model_catalog_json = %q`, catalogPath) + "\n\n" +
|
||||
codexProviderHeaderFor(codexAppProfileName) + "\n" +
|
||||
`name = "Ollama"` + "\n" +
|
||||
`base_url = "http://127.0.0.1:11434/v1/"` + "\n" +
|
||||
`wire_api = "responses"` + "\n"
|
||||
if err := os.WriteFile(configPath, []byte(content), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
if got := (&CodexApp{}).CurrentModel(); got != "llama3.2:latest" {
|
||||
t.Fatalf("CurrentModel = %q, want llama3.2:latest (:latest suffix should not be treated as drift)", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCodexAppConfigurePopulatesCatalogFromEnrichedModels(t *testing.T) {
|
||||
tmpDir := t.TempDir()
|
||||
setTestHome(t, tmpDir)
|
||||
@@ -1392,6 +1521,61 @@ func TestCodexAppRunWaitsForGracefulExitBeforeReopening(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestCodexAppRunCtrlCAbortsEntireRestartFlow(t *testing.T) {
|
||||
withCodexAppPlatform(t, "darwin")
|
||||
restoreConfirm := withLaunchConfirmPolicy(launchConfirmPolicy{yes: true})
|
||||
defer restoreConfirm()
|
||||
|
||||
oldSleep := codexAppSleep
|
||||
oldDefaultSpinner := DefaultSpinner
|
||||
t.Cleanup(func() {
|
||||
codexAppSleep = oldSleep
|
||||
DefaultSpinner = oldDefaultSpinner
|
||||
})
|
||||
|
||||
// Simulate the user pressing Ctrl+C during the graceful-exit wait: the
|
||||
// shared spinner's cancellation channel is closed on the first poll,
|
||||
// which only happens inside the wait loop.
|
||||
cancel := make(chan struct{})
|
||||
var spinnerStopped bool
|
||||
codexAppSleep = func(time.Duration) {
|
||||
select {
|
||||
case <-cancel:
|
||||
default:
|
||||
close(cancel)
|
||||
}
|
||||
}
|
||||
DefaultSpinner = func(string) *Spinner {
|
||||
return NewSpinner(func() { spinnerStopped = true }, cancel)
|
||||
}
|
||||
|
||||
var calls []string
|
||||
withCodexAppProcessHooks(t,
|
||||
func() bool { return true }, // app stays "running" so the wait polls
|
||||
func() error { calls = append(calls, "quit"); return nil },
|
||||
func() error { calls = append(calls, "open"); return nil },
|
||||
)
|
||||
codexAppExitTimeout = 5 * time.Second
|
||||
codexAppForceQuit = func() error {
|
||||
calls = append(calls, "force")
|
||||
return nil
|
||||
}
|
||||
|
||||
err := (&CodexApp{}).Run("qwen3.5", nil, nil)
|
||||
if !errors.Is(err, ErrCancelled) {
|
||||
t.Fatalf("Run error = %v, want ErrCancelled", err)
|
||||
}
|
||||
if !spinnerStopped {
|
||||
t.Fatal("expected the shared spinner to be stopped on cancel")
|
||||
}
|
||||
// The flow must abort after quit: no force-quit, no reopen, despite the app
|
||||
// still being "running" (which would otherwise trigger the force-quit path).
|
||||
want := []string{"quit"}
|
||||
if !slices.Equal(calls, want) {
|
||||
t.Fatalf("calls = %v, want the whole flow to abort after quit: %v", calls, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCodexAppRunForceStopsMacAfterGracefulTimeout(t *testing.T) {
|
||||
withCodexAppPlatform(t, "darwin")
|
||||
restoreConfirm := withLaunchConfirmPolicy(launchConfirmPolicy{yes: true})
|
||||
@@ -1445,7 +1629,7 @@ func TestCodexAppRunReturnsMacForceStopError(t *testing.T) {
|
||||
}
|
||||
|
||||
err := (&CodexApp{}).Run("qwen3.5", nil, nil)
|
||||
if err == nil || !strings.Contains(err.Error(), "force stop Codex") || !strings.Contains(err.Error(), "operation not permitted") {
|
||||
if err == nil || !strings.Contains(err.Error(), "force stop ChatGPT") || !strings.Contains(err.Error(), "operation not permitted") {
|
||||
t.Fatalf("Run error = %v, want force stop failure", err)
|
||||
}
|
||||
}
|
||||
@@ -1587,7 +1771,7 @@ func TestCodexAppRunReturnsWindowsForceStopError(t *testing.T) {
|
||||
}
|
||||
|
||||
err := (&CodexApp{}).Run("qwen3.5", nil, nil)
|
||||
if err == nil || !strings.Contains(err.Error(), "force stop Codex") || !strings.Contains(err.Error(), "access denied") {
|
||||
if err == nil || !strings.Contains(err.Error(), "force stop ChatGPT") || !strings.Contains(err.Error(), "access denied") {
|
||||
t.Fatalf("Run error = %v, want force stop failure", err)
|
||||
}
|
||||
}
|
||||
@@ -1602,6 +1786,8 @@ func TestCodexAppRunRejectsExtraArgs(t *testing.T) {
|
||||
|
||||
func TestCodexAppProcessMatchesMainAndAppServer(t *testing.T) {
|
||||
for _, command := range []string{
|
||||
"/Applications/ChatGPT.app/Contents/MacOS/ChatGPT",
|
||||
"/Applications/ChatGPT.app/Contents/Resources/codex app-server --analytics-default-enabled",
|
||||
"/Applications/Codex.app/Contents/MacOS/Codex",
|
||||
"/Applications/Codex.app/Contents/Resources/codex app-server --analytics-default-enabled",
|
||||
`C:\Users\parth\AppData\Local\Programs\Codex\Codex.exe`,
|
||||
@@ -1614,6 +1800,7 @@ func TestCodexAppProcessMatchesMainAndAppServer(t *testing.T) {
|
||||
}
|
||||
|
||||
for _, command := range []string{
|
||||
"/Applications/ChatGPT.app/Contents/Frameworks/ChatGPT Helper.app/Contents/MacOS/ChatGPT Helper",
|
||||
"/Applications/Codex.app/Contents/Frameworks/Codex Helper.app/Contents/MacOS/Codex Helper",
|
||||
"/Applications/Codex.app/Contents/Frameworks/Electron Framework.framework/Helpers/chrome_crashpad_handler",
|
||||
`"C:\Program Files\WindowsApps\OpenAI.Codex_26.429.8261.0_x64__2p2nqsd0c76g0\app\Codex.exe" --type=renderer --user-data-dir="C:\Users\parth\AppData\Roaming\Codex"`,
|
||||
@@ -1625,6 +1812,24 @@ func TestCodexAppProcessMatchesMainAndAppServer(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestCodexAppCandidatesIncludeChatGPT(t *testing.T) {
|
||||
withCodexAppPlatform(t, "darwin")
|
||||
candidates := codexAppDarwinAppCandidates()
|
||||
if len(candidates) == 0 || candidates[0] != "/Applications/ChatGPT.app" {
|
||||
t.Fatalf("darwin candidates = %v, want ChatGPT first", candidates)
|
||||
}
|
||||
if !slices.Contains(candidates, "/Applications/Codex.app") {
|
||||
t.Fatalf("darwin candidates = %v, want legacy Codex app", candidates)
|
||||
}
|
||||
|
||||
withCodexAppPlatform(t, "windows")
|
||||
local := filepath.Join(t.TempDir(), "LocalAppData")
|
||||
t.Setenv("LOCALAPPDATA", local)
|
||||
if candidates := codexAppWindowsAppCandidates(); !slices.Contains(candidates, filepath.Join(local, "Programs", "ChatGPT", "ChatGPT.exe")) {
|
||||
t.Fatalf("windows candidates = %v, want ChatGPT app", candidates)
|
||||
}
|
||||
}
|
||||
|
||||
func catalogSlugs(models []map[string]any) []string {
|
||||
slugs := make([]string, 0, len(models))
|
||||
for _, model := range models {
|
||||
|
||||
+26
-26
@@ -281,7 +281,7 @@ func TestLaunchCmdModelFlagFiltersDisabledCloudFromSavedConfig(t *testing.T) {
|
||||
case "/api/status":
|
||||
fmt.Fprintf(w, `{"cloud":{"disabled":true,"source":"config"}}`)
|
||||
case "/api/show":
|
||||
fmt.Fprintf(w, `{"model":"llama3.2"}`)
|
||||
fmt.Fprintf(w, `{"model":"sample-model"}`)
|
||||
default:
|
||||
w.WriteHeader(http.StatusNotFound)
|
||||
}
|
||||
@@ -294,7 +294,7 @@ func TestLaunchCmdModelFlagFiltersDisabledCloudFromSavedConfig(t *testing.T) {
|
||||
defer restore()
|
||||
|
||||
cmd := LaunchCmd(func(cmd *cobra.Command, args []string) error { return nil }, func(cmd *cobra.Command) {})
|
||||
cmd.SetArgs([]string{"stubeditor", "--model", "llama3.2"})
|
||||
cmd.SetArgs([]string{"stubeditor", "--model", "sample-model"})
|
||||
if err := cmd.Execute(); err != nil {
|
||||
t.Fatalf("launch command failed: %v", err)
|
||||
}
|
||||
@@ -303,14 +303,14 @@ func TestLaunchCmdModelFlagFiltersDisabledCloudFromSavedConfig(t *testing.T) {
|
||||
if err != nil {
|
||||
t.Fatalf("failed to reload integration config: %v", err)
|
||||
}
|
||||
if diff := cmp.Diff([]string{"llama3.2"}, saved.Models); diff != "" {
|
||||
if diff := cmp.Diff([]string{"sample-model"}, saved.Models); diff != "" {
|
||||
t.Fatalf("saved models mismatch (-want +got):\n%s", diff)
|
||||
}
|
||||
if diff := cmp.Diff([][]string{{"llama3.2"}}, stub.edited); diff != "" {
|
||||
if diff := cmp.Diff([][]string{{"sample-model"}}, stub.edited); diff != "" {
|
||||
t.Fatalf("editor models mismatch (-want +got):\n%s", diff)
|
||||
}
|
||||
if stub.ranModel != "llama3.2" {
|
||||
t.Fatalf("expected launch to run with llama3.2, got %q", stub.ranModel)
|
||||
if stub.ranModel != "sample-model" {
|
||||
t.Fatalf("expected launch to run with sample-model, got %q", stub.ranModel)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -325,9 +325,9 @@ func TestLaunchCmdModelFlagClearsDisabledCloudOverride(t *testing.T) {
|
||||
case "/api/experimental/model-recommendations":
|
||||
fmt.Fprint(w, `{"recommendations":[]}`)
|
||||
case "/api/tags":
|
||||
fmt.Fprint(w, `{"models":[{"name":"llama3.2"}]}`)
|
||||
fmt.Fprint(w, `{"models":[{"name":"sample-model"}]}`)
|
||||
case "/api/show":
|
||||
fmt.Fprint(w, `{"model":"llama3.2"}`)
|
||||
fmt.Fprint(w, `{"model":"sample-model"}`)
|
||||
default:
|
||||
w.WriteHeader(http.StatusNotFound)
|
||||
}
|
||||
@@ -347,7 +347,7 @@ func TestLaunchCmdModelFlagClearsDisabledCloudOverride(t *testing.T) {
|
||||
DefaultSingleSelector = func(title string, items []SelectionItem, current string) (string, error) {
|
||||
selectorCalls++
|
||||
gotCurrent = current
|
||||
return "llama3.2", nil
|
||||
return "sample-model", nil
|
||||
}
|
||||
|
||||
cmd := LaunchCmd(func(cmd *cobra.Command, args []string) error { return nil }, func(cmd *cobra.Command) {})
|
||||
@@ -364,7 +364,7 @@ func TestLaunchCmdModelFlagClearsDisabledCloudOverride(t *testing.T) {
|
||||
if gotCurrent != "" {
|
||||
t.Fatalf("expected disabled override to be cleared before selection, got current %q", gotCurrent)
|
||||
}
|
||||
if stub.ranModel != "llama3.2" {
|
||||
if stub.ranModel != "sample-model" {
|
||||
t.Fatalf("expected launch to run with replacement local model, got %q", stub.ranModel)
|
||||
}
|
||||
if !strings.Contains(stderr, "Warning: ignoring --model glm-5:cloud because cloud is disabled") {
|
||||
@@ -375,7 +375,7 @@ func TestLaunchCmdModelFlagClearsDisabledCloudOverride(t *testing.T) {
|
||||
if err != nil {
|
||||
t.Fatalf("failed to reload integration config: %v", err)
|
||||
}
|
||||
if diff := cmp.Diff([]string{"llama3.2"}, saved.Models); diff != "" {
|
||||
if diff := cmp.Diff([]string{"sample-model"}, saved.Models); diff != "" {
|
||||
t.Fatalf("saved models mismatch (-want +got):\n%s", diff)
|
||||
}
|
||||
}
|
||||
@@ -424,7 +424,7 @@ func TestLaunchCmdYes_AutoConfirmsLaunchPromptPath(t *testing.T) {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
switch r.URL.Path {
|
||||
case "/api/show":
|
||||
fmt.Fprint(w, `{"model":"llama3.2"}`)
|
||||
fmt.Fprint(w, `{"model":"sample-model"}`)
|
||||
case "/api/status":
|
||||
w.WriteHeader(http.StatusNotFound)
|
||||
fmt.Fprint(w, `{"error":"not found"}`)
|
||||
@@ -445,16 +445,16 @@ func TestLaunchCmdYes_AutoConfirmsLaunchPromptPath(t *testing.T) {
|
||||
}
|
||||
|
||||
cmd := LaunchCmd(func(cmd *cobra.Command, args []string) error { return nil }, func(cmd *cobra.Command) {})
|
||||
cmd.SetArgs([]string{"stubeditor", "--model", "llama3.2", "--yes"})
|
||||
cmd.SetArgs([]string{"stubeditor", "--model", "sample-model", "--yes"})
|
||||
if err := cmd.Execute(); err != nil {
|
||||
t.Fatalf("launch command with --yes failed: %v", err)
|
||||
}
|
||||
|
||||
if diff := cmp.Diff([][]string{{"llama3.2"}}, stub.edited); diff != "" {
|
||||
if diff := cmp.Diff([][]string{{"sample-model"}}, stub.edited); diff != "" {
|
||||
t.Fatalf("editor models mismatch (-want +got):\n%s", diff)
|
||||
}
|
||||
if stub.ranModel != "llama3.2" {
|
||||
t.Fatalf("expected launch to run with llama3.2, got %q", stub.ranModel)
|
||||
if stub.ranModel != "sample-model" {
|
||||
t.Fatalf("expected launch to run with sample-model, got %q", stub.ranModel)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -513,7 +513,7 @@ func TestLaunchCmdHeadlessWithoutYes_AllowsConfiguredLaunch(t *testing.T) {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
switch r.URL.Path {
|
||||
case "/api/show":
|
||||
fmt.Fprint(w, `{"model":"llama3.2"}`)
|
||||
fmt.Fprint(w, `{"model":"sample-model"}`)
|
||||
case "/api/status":
|
||||
w.WriteHeader(http.StatusNotFound)
|
||||
fmt.Fprint(w, `{"error":"not found"}`)
|
||||
@@ -534,15 +534,15 @@ func TestLaunchCmdHeadlessWithoutYes_AllowsConfiguredLaunch(t *testing.T) {
|
||||
}
|
||||
|
||||
cmd := LaunchCmd(func(cmd *cobra.Command, args []string) error { return nil }, func(cmd *cobra.Command) {})
|
||||
cmd.SetArgs([]string{"stubeditor", "--model", "llama3.2"})
|
||||
cmd.SetArgs([]string{"stubeditor", "--model", "sample-model"})
|
||||
err := cmd.Execute()
|
||||
if err != nil {
|
||||
t.Fatalf("expected launch command to succeed without --yes when an explicit model is provided, got %v", err)
|
||||
}
|
||||
if diff := compareStringSlices(stub.edited, [][]string{{"llama3.2"}}); diff != "" {
|
||||
if diff := compareStringSlices(stub.edited, [][]string{{"sample-model"}}); diff != "" {
|
||||
t.Fatalf("unexpected editor writes (-want +got):\n%s", diff)
|
||||
}
|
||||
if stub.ranModel != "llama3.2" {
|
||||
if stub.ranModel != "sample-model" {
|
||||
t.Fatalf("expected launch to run configured model, got %q", stub.ranModel)
|
||||
}
|
||||
}
|
||||
@@ -551,7 +551,7 @@ func TestLaunchCmdIntegrationArgPromptsForModelWithSavedSelection(t *testing.T)
|
||||
tmpDir := t.TempDir()
|
||||
setLaunchTestHome(t, tmpDir)
|
||||
|
||||
if err := config.SaveIntegration("stubapp", []string{"llama3.2"}); err != nil {
|
||||
if err := config.SaveIntegration("stubapp", []string{"sample-model"}); err != nil {
|
||||
t.Fatalf("failed to seed saved config: %v", err)
|
||||
}
|
||||
|
||||
@@ -560,7 +560,7 @@ func TestLaunchCmdIntegrationArgPromptsForModelWithSavedSelection(t *testing.T)
|
||||
case "/api/experimental/model-recommendations":
|
||||
fmt.Fprint(w, `{"recommendations":[]}`)
|
||||
case "/api/tags":
|
||||
fmt.Fprint(w, `{"models":[{"name":"llama3.2"},{"name":"qwen3:8b"}]}`)
|
||||
fmt.Fprint(w, `{"models":[{"name":"sample-model"},{"name":"qwen3:8b"}]}`)
|
||||
case "/api/show":
|
||||
fmt.Fprint(w, `{"model":"qwen3:8b"}`)
|
||||
default:
|
||||
@@ -589,8 +589,8 @@ func TestLaunchCmdIntegrationArgPromptsForModelWithSavedSelection(t *testing.T)
|
||||
t.Fatalf("launch command failed: %v", err)
|
||||
}
|
||||
|
||||
if gotCurrent != "llama3.2" {
|
||||
t.Fatalf("expected selector current model to be saved model llama3.2, got %q", gotCurrent)
|
||||
if gotCurrent != "sample-model" {
|
||||
t.Fatalf("expected selector current model to be saved model sample-model, got %q", gotCurrent)
|
||||
}
|
||||
if stub.ranModel != "qwen3:8b" {
|
||||
t.Fatalf("expected launch to run selected model qwen3:8b, got %q", stub.ranModel)
|
||||
@@ -611,14 +611,14 @@ func TestLaunchCmdHeadlessYes_IntegrationRequiresModelEvenWhenSaved(t *testing.T
|
||||
withLauncherHooks(t)
|
||||
withInteractiveSession(t, false)
|
||||
|
||||
if err := config.SaveIntegration("stubapp", []string{"llama3.2"}); err != nil {
|
||||
if err := config.SaveIntegration("stubapp", []string{"sample-model"}); err != nil {
|
||||
t.Fatalf("failed to seed saved config: %v", err)
|
||||
}
|
||||
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
switch r.URL.Path {
|
||||
case "/api/show":
|
||||
fmt.Fprint(w, `{"model":"llama3.2"}`)
|
||||
fmt.Fprint(w, `{"model":"sample-model"}`)
|
||||
default:
|
||||
w.WriteHeader(http.StatusNotFound)
|
||||
}
|
||||
|
||||
@@ -0,0 +1,133 @@
|
||||
package launch
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"github.com/ollama/ollama/internal/modelref"
|
||||
)
|
||||
|
||||
var deprecatedLaunchModels = map[string]struct{}{
|
||||
"codellama": {},
|
||||
"qwen2.5": {},
|
||||
"qwen2.5-coder": {},
|
||||
"llama3": {},
|
||||
"llama3.1": {},
|
||||
"llama3.2": {},
|
||||
"llama3.3": {},
|
||||
"mistral": {},
|
||||
"starcoder": {},
|
||||
}
|
||||
|
||||
var deprecatedLaunchModelTags = map[string]map[string]struct{}{
|
||||
"deepseek-r1": {
|
||||
"": {},
|
||||
"latest": {},
|
||||
"1.5b": {},
|
||||
"7b": {},
|
||||
"8b": {},
|
||||
"14b": {},
|
||||
"32b": {},
|
||||
},
|
||||
}
|
||||
|
||||
var errDeprecatedLaunchModelDeclined = fmt.Errorf("%w: deprecated launch model declined", ErrCancelled)
|
||||
|
||||
func isDeprecatedLaunchModel(name string) bool {
|
||||
family, tag := normalizedLaunchModelRef(name)
|
||||
if _, ok := deprecatedLaunchModels[family]; ok {
|
||||
return true
|
||||
}
|
||||
tags, ok := deprecatedLaunchModelTags[family]
|
||||
if !ok {
|
||||
return false
|
||||
}
|
||||
_, ok = tags[tag]
|
||||
return ok
|
||||
}
|
||||
|
||||
func deprecatedLaunchModelPrompt(name, label, commandName, cloudRec, localRec string) string {
|
||||
if !isDeprecatedLaunchModel(name) {
|
||||
return ""
|
||||
}
|
||||
if label = strings.TrimSpace(label); label == "" {
|
||||
label = "ollama launch"
|
||||
}
|
||||
|
||||
var b strings.Builder
|
||||
fmt.Fprintf(&b, "%s does not work well with %s. ", name, label)
|
||||
switch {
|
||||
case cloudRec != "" && localRec != "":
|
||||
fmt.Fprintf(&b, "Try an agent-capable model like %s or %s instead", cloudRec, localRec)
|
||||
case cloudRec != "":
|
||||
fmt.Fprintf(&b, "Try an agent-capable model like %s instead", cloudRec)
|
||||
case localRec != "":
|
||||
fmt.Fprintf(&b, "Try an agent-capable model like %s instead", localRec)
|
||||
default:
|
||||
b.WriteString("Try a newer recommended agent-capable model instead")
|
||||
}
|
||||
if command := launchReplacementCommand(commandName, firstNonEmpty(cloudRec, localRec)); command != "" {
|
||||
fmt.Fprintf(&b, ":\n %s", command)
|
||||
} else {
|
||||
b.WriteString(".")
|
||||
}
|
||||
fmt.Fprintf(&b, "\n\nLaunch with %s anyway?", name)
|
||||
return b.String()
|
||||
}
|
||||
|
||||
func launchReplacementCommand(commandName, model string) string {
|
||||
commandName = strings.TrimSpace(commandName)
|
||||
model = strings.TrimSpace(model)
|
||||
if commandName == "" || model == "" {
|
||||
return ""
|
||||
}
|
||||
return fmt.Sprintf("ollama launch %s --model %s", commandName, model)
|
||||
}
|
||||
|
||||
func firstNonEmpty(values ...string) string {
|
||||
for _, value := range values {
|
||||
if value != "" {
|
||||
return value
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func normalizedLaunchModelRef(name string) (string, string) {
|
||||
name = strings.TrimSpace(strings.ToLower(name))
|
||||
if name == "" {
|
||||
return "", ""
|
||||
}
|
||||
if base, stripped := modelref.StripCloudSourceTag(name); stripped {
|
||||
name = base
|
||||
}
|
||||
if idx := strings.LastIndex(name, "/"); idx >= 0 {
|
||||
name = name[idx+1:]
|
||||
}
|
||||
tag := ""
|
||||
if idx := strings.Index(name, ":"); idx >= 0 {
|
||||
tag = strings.TrimSpace(name[idx+1:])
|
||||
name = name[:idx]
|
||||
}
|
||||
return strings.TrimSpace(name), tag
|
||||
}
|
||||
|
||||
func filterDeprecatedLaunchModelItems(items []ModelItem) []ModelItem {
|
||||
filtered := items[:0]
|
||||
for _, item := range items {
|
||||
if !isDeprecatedLaunchModel(item.Name) {
|
||||
filtered = append(filtered, item)
|
||||
}
|
||||
}
|
||||
return filtered
|
||||
}
|
||||
|
||||
func filterDeprecatedLaunchModelNames(models []string) []string {
|
||||
filtered := models[:0]
|
||||
for _, model := range models {
|
||||
if !isDeprecatedLaunchModel(model) {
|
||||
filtered = append(filtered, model)
|
||||
}
|
||||
}
|
||||
return filtered
|
||||
}
|
||||
@@ -0,0 +1,68 @@
|
||||
package launch
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestLaunchModelDeprecation(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
deprecated bool
|
||||
}{
|
||||
{name: "qwen2.5", deprecated: true},
|
||||
{name: "qwen2.5:14b", deprecated: true},
|
||||
{name: "qwen2.5-coder:32b", deprecated: true},
|
||||
{name: "library/qwen2.5-coder:7b", deprecated: true},
|
||||
{name: "llama3", deprecated: true},
|
||||
{name: "llama3.1:8b", deprecated: true},
|
||||
{name: "llama3.2:latest", deprecated: true},
|
||||
{name: "llama3.3:70b", deprecated: true},
|
||||
{name: "llama3.2:cloud", deprecated: true},
|
||||
{name: "codellama", deprecated: true},
|
||||
{name: "codellama:13b-code", deprecated: true},
|
||||
{name: "library/codellama:7b", deprecated: true},
|
||||
{name: "starcoder", deprecated: true},
|
||||
{name: "starcoder:15b", deprecated: true},
|
||||
{name: "mistral", deprecated: true},
|
||||
{name: "mistral:7b", deprecated: true},
|
||||
{name: "deepseek-r1", deprecated: true},
|
||||
{name: "deepseek-r1:latest", deprecated: true},
|
||||
{name: "deepseek-r1:1.5b", deprecated: true},
|
||||
{name: "deepseek-r1:7b", deprecated: true},
|
||||
{name: "deepseek-r1:8b", deprecated: true},
|
||||
{name: "deepseek-r1:14b", deprecated: true},
|
||||
{name: "deepseek-r1:32b", deprecated: true},
|
||||
{name: "deepseek-r1:32b-cloud", deprecated: true},
|
||||
{name: "qwen3.5", deprecated: false},
|
||||
{name: "qwen3-coder:30b", deprecated: false},
|
||||
{name: "gemma4", deprecated: false},
|
||||
{name: "my-qwen2.5-coder", deprecated: false},
|
||||
{name: "llama3.2-inspired", deprecated: false},
|
||||
{name: "codellama-inspired", deprecated: false},
|
||||
{name: "starcoder2:15b", deprecated: false},
|
||||
{name: "mixtral:8x7b", deprecated: false},
|
||||
{name: "deepseek-r1:70b", deprecated: false},
|
||||
{name: "deepseek-r1:671b", deprecated: false},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if got := isDeprecatedLaunchModel(tt.name); got != tt.deprecated {
|
||||
t.Fatalf("isDeprecatedLaunchModel(%q) = %v, want %v", tt.name, got, tt.deprecated)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestDeprecatedLaunchModelErrorMentionsRecommendedModels(t *testing.T) {
|
||||
prompt := deprecatedLaunchModelPrompt("qwen2.5-coder:32b", "Codex", "codex", "recommended-cloud:cloud", "recommended-local")
|
||||
if prompt == "" {
|
||||
t.Fatal("expected deprecated model prompt")
|
||||
}
|
||||
for _, want := range []string{"qwen2.5-coder:32b does not work well with Codex", "recommended-cloud:cloud", "recommended-local", "ollama launch codex --model recommended-cloud:cloud", "Launch with qwen2.5-coder:32b anyway?"} {
|
||||
if !strings.Contains(prompt, want) {
|
||||
t.Fatalf("prompt %q does not contain %q", prompt, want)
|
||||
}
|
||||
}
|
||||
}
|
||||
+256
-21
@@ -14,6 +14,7 @@ import (
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"golang.org/x/mod/semver"
|
||||
"gopkg.in/yaml.v3"
|
||||
|
||||
"github.com/ollama/ollama/api"
|
||||
@@ -23,7 +24,11 @@ import (
|
||||
)
|
||||
|
||||
const (
|
||||
hermesInstallScript = "curl -fsSL https://raw.githubusercontent.com/NousResearch/hermes-agent/main/scripts/install.sh | bash -s -- --skip-setup"
|
||||
// https://github.com/NousResearch/hermes-agent/releases/tag/v2026.6.5
|
||||
hermesDesktopMinVersion = "v0.16.0"
|
||||
hermesInstallScript = "curl -fsSL https://hermes-agent.nousresearch.com/install.sh | bash -s -- --skip-setup"
|
||||
hermesWindowsInstallURL = "https://hermes-agent.nousresearch.com/install.ps1"
|
||||
hermesWindowsInstallCmd = "& ([scriptblock]::Create((irm " + hermesWindowsInstallURL + "))) -SkipSetup"
|
||||
hermesProviderName = "Ollama"
|
||||
hermesProviderKey = "ollama-launch"
|
||||
hermesLegacyKey = "ollama"
|
||||
@@ -81,6 +86,177 @@ func (h *Hermes) Run(_ string, _ []LaunchModel, args []string) error {
|
||||
return hermesAttachedCommand(bin, args...).Run()
|
||||
}
|
||||
|
||||
type HermesDesktop struct {
|
||||
Hermes
|
||||
}
|
||||
|
||||
func (h *HermesDesktop) String() string { return "Hermes Desktop" }
|
||||
|
||||
func (h *HermesDesktop) Run(_ string, _ []LaunchModel, args []string) error {
|
||||
bin, err := h.binary()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := h.ensureHermesDesktopMinVersion(bin); err != nil {
|
||||
return err
|
||||
}
|
||||
return hermesAttachedCommand(bin, h.launchArgs(args)...).Run()
|
||||
}
|
||||
|
||||
func (h *HermesDesktop) ensureHermesDesktopMinVersion(bin string) error {
|
||||
if hermesGOOS == "windows" {
|
||||
return nil
|
||||
}
|
||||
version := hermesVersionOf(bin)
|
||||
if version == "" {
|
||||
return nil
|
||||
}
|
||||
if semver.Compare(version, hermesDesktopMinVersion) >= 0 {
|
||||
return nil
|
||||
}
|
||||
fmt.Fprintf(os.Stderr, "%sHermes %s is older than the minimum version (%s) for `hermes desktop`; updating...%s\n", ansiGray, version, hermesDesktopMinVersion, ansiReset)
|
||||
if err := hermesAttachedCommand(bin, "update").Run(); err != nil {
|
||||
return fmt.Errorf("failed to update hermes to %s or newer: %w", hermesDesktopMinVersion, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func hermesVersionOf(bin string) string {
|
||||
out, err := hermesCommand(bin, "--version").Output()
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
firstLine := strings.SplitN(strings.TrimSpace(string(out)), "\n", 2)[0]
|
||||
return parseHermesVersion(firstLine)
|
||||
}
|
||||
|
||||
func parseHermesVersion(firstLine string) string {
|
||||
for _, field := range strings.Fields(firstLine) {
|
||||
if semver.IsValid(field) {
|
||||
return field
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (h *HermesDesktop) Onboard() error {
|
||||
return config.MarkIntegrationOnboarded("hermes-desktop")
|
||||
}
|
||||
|
||||
func (h *HermesDesktop) launchArgs(args []string) []string {
|
||||
launchArgs := []string{"desktop"}
|
||||
if h.shouldSkipDesktopBuild(args) {
|
||||
launchArgs = append(launchArgs, "--skip-build")
|
||||
}
|
||||
return append(launchArgs, args...)
|
||||
}
|
||||
|
||||
func (h *HermesDesktop) shouldSkipDesktopBuild(args []string) bool {
|
||||
if hermesDesktopHasFlag(args, "--skip-build", "--source", "--build-only", "--help", "-h") {
|
||||
return false
|
||||
}
|
||||
return h.packagedAppExists()
|
||||
}
|
||||
|
||||
func (h *HermesDesktop) packagedAppExists() bool {
|
||||
for _, root := range hermesDesktopReleaseRoots() {
|
||||
for _, candidate := range hermesDesktopPackagedExecutableCandidates(root) {
|
||||
if _, err := os.Stat(candidate); err == nil {
|
||||
return true
|
||||
}
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// These roots mirror Hermes' own install layout:
|
||||
// install.sh uses ~/.hermes/hermes-agent for user installs and
|
||||
// /usr/local/lib/hermes-agent for new Linux root installs; install.ps1
|
||||
// and the bootstrap installer use %LOCALAPPDATA%\hermes\hermes-agent on
|
||||
// Windows. HERMES_HOME and HERMES_INSTALL_DIR are installer-supported
|
||||
// overrides.
|
||||
func hermesDesktopReleaseRoots() []string {
|
||||
var installRoots []string
|
||||
add := func(path string) {
|
||||
path = strings.TrimSpace(path)
|
||||
if path == "" {
|
||||
return
|
||||
}
|
||||
installRoots = append(installRoots, filepath.Clean(path))
|
||||
}
|
||||
|
||||
if installDir := strings.TrimSpace(os.Getenv("HERMES_INSTALL_DIR")); installDir != "" {
|
||||
add(installDir)
|
||||
}
|
||||
if hermesHome := strings.TrimSpace(os.Getenv("HERMES_HOME")); hermesHome != "" {
|
||||
add(filepath.Join(hermesHome, "hermes-agent"))
|
||||
}
|
||||
|
||||
home, err := hermesUserHome()
|
||||
if err == nil {
|
||||
switch hermesGOOS {
|
||||
case "windows":
|
||||
if localAppData := strings.TrimSpace(os.Getenv("LOCALAPPDATA")); localAppData != "" {
|
||||
add(filepath.Join(localAppData, "hermes", "hermes-agent"))
|
||||
}
|
||||
add(filepath.Join(home, ".hermes", "hermes-agent"))
|
||||
default:
|
||||
add(filepath.Join(home, ".hermes", "hermes-agent"))
|
||||
if hermesGOOS == "linux" {
|
||||
add(filepath.Join(string(filepath.Separator), "usr", "local", "lib", "hermes-agent"))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
seen := make(map[string]bool, len(installRoots))
|
||||
releaseRoots := make([]string, 0, len(installRoots))
|
||||
for _, root := range installRoots {
|
||||
releaseRoot := filepath.Join(root, "apps", "desktop", "release")
|
||||
if seen[releaseRoot] {
|
||||
continue
|
||||
}
|
||||
seen[releaseRoot] = true
|
||||
releaseRoots = append(releaseRoots, releaseRoot)
|
||||
}
|
||||
return releaseRoots
|
||||
}
|
||||
|
||||
func hermesDesktopPackagedExecutableCandidates(releaseRoot string) []string {
|
||||
switch hermesGOOS {
|
||||
case "darwin":
|
||||
matches, err := filepath.Glob(filepath.Join(releaseRoot, "mac*", "Hermes.app", "Contents", "MacOS", "Hermes"))
|
||||
if err != nil {
|
||||
return nil
|
||||
}
|
||||
return matches
|
||||
case "windows":
|
||||
return []string{
|
||||
filepath.Join(releaseRoot, "win-unpacked", "Hermes.exe"),
|
||||
filepath.Join(releaseRoot, "win-ia32-unpacked", "Hermes.exe"),
|
||||
filepath.Join(releaseRoot, "win-arm64-unpacked", "Hermes.exe"),
|
||||
}
|
||||
default:
|
||||
return []string{
|
||||
filepath.Join(releaseRoot, "linux-unpacked", "hermes"),
|
||||
filepath.Join(releaseRoot, "linux-unpacked", "Hermes"),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func hermesDesktopHasFlag(args []string, names ...string) bool {
|
||||
for _, arg := range args {
|
||||
if arg == "--" {
|
||||
return false
|
||||
}
|
||||
for _, name := range names {
|
||||
if arg == name {
|
||||
return true
|
||||
}
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func (h *Hermes) Paths() []string {
|
||||
configPath, err := hermesConfigPath()
|
||||
if err != nil {
|
||||
@@ -183,22 +359,24 @@ func (h *Hermes) installed() bool {
|
||||
}
|
||||
|
||||
func (h *Hermes) ensureInstalled() error {
|
||||
return h.ensureInstalledFor("hermes")
|
||||
}
|
||||
|
||||
func (h *Hermes) ensureInstalledFor(command string) error {
|
||||
if h.installed() {
|
||||
return nil
|
||||
}
|
||||
|
||||
if hermesGOOS == "windows" {
|
||||
return hermesWindowsHint()
|
||||
}
|
||||
|
||||
var missing []string
|
||||
for _, dep := range []string{"bash", "curl", "git"} {
|
||||
if _, err := hermesLookPath(dep); err != nil {
|
||||
missing = append(missing, dep)
|
||||
if hermesGOOS != "windows" {
|
||||
for _, dep := range []string{"bash", "curl", "git"} {
|
||||
if _, err := hermesLookPath(dep); err != nil {
|
||||
missing = append(missing, dep)
|
||||
}
|
||||
}
|
||||
}
|
||||
if len(missing) > 0 {
|
||||
return fmt.Errorf("Hermes is not installed and required dependencies are missing\n\nInstall the following first:\n %s\n\nThen re-run:\n ollama launch hermes", strings.Join(missing, "\n "))
|
||||
return fmt.Errorf("Hermes is not installed and required dependencies are missing\n\nInstall the following first:\n %s\n\nThen re-run:\n ollama launch %s", strings.Join(missing, "\n "), command)
|
||||
}
|
||||
|
||||
ok, err := ConfirmPrompt("Hermes is not installed. Install now?")
|
||||
@@ -210,7 +388,7 @@ func (h *Hermes) ensureInstalled() error {
|
||||
}
|
||||
|
||||
fmt.Fprintf(os.Stderr, "\nInstalling Hermes...\n")
|
||||
if err := hermesAttachedCommand("bash", "-lc", hermesInstallScript).Run(); err != nil {
|
||||
if err := h.runInstallScript(); err != nil {
|
||||
return fmt.Errorf("failed to install hermes: %w", err)
|
||||
}
|
||||
|
||||
@@ -222,6 +400,13 @@ func (h *Hermes) ensureInstalled() error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *Hermes) runInstallScript() error {
|
||||
if hermesGOOS == "windows" {
|
||||
return hermesAttachedCommand("powershell.exe", "-NoProfile", "-ExecutionPolicy", "Bypass", "-Command", hermesWindowsInstallCmd).Run()
|
||||
}
|
||||
return hermesAttachedCommand("bash", "-lc", hermesInstallScript).Run()
|
||||
}
|
||||
|
||||
func (h *Hermes) listModels(defaultModel string) []string {
|
||||
client := hermesOllamaClient()
|
||||
resp, err := client.List(context.Background())
|
||||
@@ -259,7 +444,12 @@ func (h *Hermes) binary() (string, error) {
|
||||
}
|
||||
|
||||
if hermesGOOS == "windows" {
|
||||
return "", hermesWindowsHint()
|
||||
for _, fallback := range hermesWindowsBinaryFallbacks() {
|
||||
if _, err := os.Stat(fallback); err == nil {
|
||||
return fallback, nil
|
||||
}
|
||||
}
|
||||
return "", fmt.Errorf("hermes is not installed")
|
||||
}
|
||||
|
||||
home, err := hermesUserHome()
|
||||
@@ -274,12 +464,63 @@ func (h *Hermes) binary() (string, error) {
|
||||
return "", fmt.Errorf("hermes is not installed")
|
||||
}
|
||||
|
||||
func hermesConfigPath() (string, error) {
|
||||
func hermesWindowsBinaryFallbacks() []string {
|
||||
var roots []string
|
||||
add := func(root string) {
|
||||
root = strings.TrimSpace(root)
|
||||
if root != "" {
|
||||
roots = append(roots, filepath.Clean(root))
|
||||
}
|
||||
}
|
||||
|
||||
add(os.Getenv("HERMES_HOME"))
|
||||
add(os.Getenv("LOCALAPPDATA"))
|
||||
if home, err := hermesUserHome(); err == nil {
|
||||
add(filepath.Join(home, "AppData", "Local"))
|
||||
}
|
||||
|
||||
seen := make(map[string]bool, len(roots))
|
||||
var fallbacks []string
|
||||
for _, root := range roots {
|
||||
if seen[root] {
|
||||
continue
|
||||
}
|
||||
seen[root] = true
|
||||
fallbacks = append(fallbacks, filepath.Join(root, "hermes-agent", "venv", "Scripts", "hermes.exe"))
|
||||
if filepath.Base(root) != "hermes" {
|
||||
fallbacks = append(fallbacks, filepath.Join(root, "hermes", "hermes-agent", "venv", "Scripts", "hermes.exe"))
|
||||
}
|
||||
}
|
||||
return fallbacks
|
||||
}
|
||||
|
||||
func hermesHomePath() (string, error) {
|
||||
if hermesHome := strings.TrimSpace(os.Getenv("HERMES_HOME")); hermesHome != "" {
|
||||
return filepath.Clean(hermesHome), nil
|
||||
}
|
||||
if hermesGOOS == "windows" {
|
||||
if localAppData := strings.TrimSpace(os.Getenv("LOCALAPPDATA")); localAppData != "" {
|
||||
return filepath.Join(localAppData, "hermes"), nil
|
||||
}
|
||||
home, err := hermesUserHome()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return filepath.Join(home, "AppData", "Local", "hermes"), nil
|
||||
}
|
||||
home, err := hermesUserHome()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return filepath.Join(home, ".hermes", "config.yaml"), nil
|
||||
return filepath.Join(home, ".hermes"), nil
|
||||
}
|
||||
|
||||
func hermesConfigPath() (string, error) {
|
||||
home, err := hermesHomePath()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return filepath.Join(home, "config.yaml"), nil
|
||||
}
|
||||
|
||||
func hermesBaseURL() string {
|
||||
@@ -287,11 +528,11 @@ func hermesBaseURL() string {
|
||||
}
|
||||
|
||||
func hermesEnvPath() (string, error) {
|
||||
home, err := hermesUserHome()
|
||||
home, err := hermesHomePath()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return filepath.Join(home, ".hermes", ".env"), nil
|
||||
return filepath.Join(home, ".env"), nil
|
||||
}
|
||||
|
||||
func (h *Hermes) runGatewaySetupPreflight(args []string, runSetup func() error) error {
|
||||
@@ -671,9 +912,3 @@ func hermesAttachedCommand(name string, args ...string) *exec.Cmd {
|
||||
cmd.Stderr = os.Stderr
|
||||
return cmd
|
||||
}
|
||||
|
||||
func hermesWindowsHint() error {
|
||||
return fmt.Errorf("Hermes on Windows requires WSL2. Install WSL with: wsl --install\n" +
|
||||
"Then run 'ollama launch hermes' from inside your WSL shell.\n" +
|
||||
"Docs: https://hermes-agent.nousresearch.com/docs/getting-started/installation/")
|
||||
}
|
||||
+386
-13
@@ -8,6 +8,7 @@ import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"slices"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
@@ -65,6 +66,20 @@ func clearHermesMessagingEnvVars(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func clearHermesDesktopPackageEnvVars(t *testing.T) {
|
||||
t.Helper()
|
||||
for _, key := range []string{"HERMES_INSTALL_DIR", "HERMES_HOME", "LOCALAPPDATA"} {
|
||||
if value, ok := os.LookupEnv(key); ok {
|
||||
t.Setenv(key, value)
|
||||
} else {
|
||||
t.Setenv(key, "")
|
||||
}
|
||||
if err := os.Unsetenv(key); err != nil {
|
||||
t.Fatalf("unset %s: %v", key, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestHermesIntegration(t *testing.T) {
|
||||
h := &Hermes{}
|
||||
|
||||
@@ -408,19 +423,36 @@ func TestHermesConfigureMigratesLegacyManagedAliases(t *testing.T) {
|
||||
func TestHermesPathsUsesLocalConfigPathForNativeWindowsHermes(t *testing.T) {
|
||||
tmpDir := t.TempDir()
|
||||
winHome := filepath.Join(tmpDir, "winhome")
|
||||
localAppData := filepath.Join(tmpDir, "LocalAppData")
|
||||
setTestHome(t, winHome)
|
||||
withHermesPlatform(t, "windows")
|
||||
withHermesUserHome(t, winHome)
|
||||
t.Setenv("PATH", tmpDir)
|
||||
t.Setenv("LOCALAPPDATA", localAppData)
|
||||
writeFakeBinary(t, tmpDir, "hermes")
|
||||
|
||||
got := (&Hermes{}).Paths()
|
||||
want := filepath.Join(winHome, ".hermes", "config.yaml")
|
||||
want := filepath.Join(localAppData, "hermes", "config.yaml")
|
||||
if len(got) != 1 || got[0] != want {
|
||||
t.Fatalf("expected local config path %q, got %v", want, got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestHermesPathsUsesHermesHomeOverride(t *testing.T) {
|
||||
tmpDir := t.TempDir()
|
||||
hermesHome := filepath.Join(tmpDir, "custom-hermes-home")
|
||||
setTestHome(t, filepath.Join(tmpDir, "home"))
|
||||
withHermesPlatform(t, "windows")
|
||||
t.Setenv("HERMES_HOME", hermesHome)
|
||||
t.Setenv("LOCALAPPDATA", filepath.Join(tmpDir, "LocalAppData"))
|
||||
|
||||
got := (&Hermes{}).Paths()
|
||||
want := filepath.Join(hermesHome, "config.yaml")
|
||||
if len(got) != 1 || got[0] != want {
|
||||
t.Fatalf("expected HERMES_HOME config path %q, got %v", want, got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestHermesCurrentModelRequiresHealthyManagedConfig(t *testing.T) {
|
||||
tmpDir := t.TempDir()
|
||||
setTestHome(t, tmpDir)
|
||||
@@ -565,6 +597,314 @@ func TestHermesRunPassthroughArgs(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func writeHermesDesktopPackage(t *testing.T, home string) {
|
||||
t.Helper()
|
||||
writeHermesDesktopExecutable(t,
|
||||
filepath.Join(home, ".hermes", "hermes-agent", "apps", "desktop", "release"),
|
||||
hermesDesktopTestExecutableRelativePath(hermesGOOS),
|
||||
)
|
||||
}
|
||||
|
||||
func writeHermesDesktopExecutable(t *testing.T, releaseRoot, relative string) {
|
||||
t.Helper()
|
||||
path := filepath.Join(releaseRoot, relative)
|
||||
if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := os.WriteFile(path, []byte("#!/bin/sh\nexit 0\n"), 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func hermesDesktopTestExecutableRelativePath(goos string) string {
|
||||
switch goos {
|
||||
case "darwin":
|
||||
return filepath.Join("mac-arm64", "Hermes.app", "Contents", "MacOS", "Hermes")
|
||||
case "windows":
|
||||
return filepath.Join("win-unpacked", "Hermes.exe")
|
||||
default:
|
||||
return filepath.Join("linux-unpacked", "hermes")
|
||||
}
|
||||
}
|
||||
|
||||
func writeHermesDesktopTestBinary(t *testing.T, dir string) {
|
||||
t.Helper()
|
||||
bin := filepath.Join(dir, "hermes")
|
||||
if err := os.WriteFile(bin, []byte("#!/bin/sh\nif [ \"$1\" = \"--version\" ]; then\n printf 'Hermes Agent v0.16.0 (2026.6.5)\\n'\n exit 0\nfi\nprintf '[%s]\\n' \"$*\" >> \"$HOME/hermes-invocations.log\"\n"), 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func readHermesDesktopInvocations(t *testing.T, home string) string {
|
||||
t.Helper()
|
||||
data, err := os.ReadFile(filepath.Join(home, "hermes-invocations.log"))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return strings.TrimSpace(string(data))
|
||||
}
|
||||
|
||||
func TestHermesDesktopRun(t *testing.T) {
|
||||
if runtime.GOOS == "windows" {
|
||||
t.Skip("uses a POSIX shell test binary")
|
||||
}
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
goos string
|
||||
args []string
|
||||
hasPackage bool
|
||||
clearPkgEnv bool
|
||||
want string
|
||||
}{
|
||||
{
|
||||
name: "desktop subcommand",
|
||||
goos: "darwin",
|
||||
args: []string{"--foreground"},
|
||||
clearPkgEnv: true,
|
||||
want: "[desktop --foreground]",
|
||||
},
|
||||
{
|
||||
name: "skip build when packaged app exists",
|
||||
goos: runtime.GOOS,
|
||||
args: []string{"--cwd", "/tmp/project"},
|
||||
hasPackage: true,
|
||||
want: "[desktop --skip-build --cwd /tmp/project]",
|
||||
},
|
||||
{
|
||||
name: "explicit skip build",
|
||||
goos: runtime.GOOS,
|
||||
args: []string{"--skip-build"},
|
||||
hasPackage: true,
|
||||
want: "[desktop --skip-build]",
|
||||
},
|
||||
{
|
||||
name: "source mode",
|
||||
goos: runtime.GOOS,
|
||||
args: []string{"--source"},
|
||||
hasPackage: true,
|
||||
want: "[desktop --source]",
|
||||
},
|
||||
{
|
||||
name: "build only",
|
||||
goos: runtime.GOOS,
|
||||
args: []string{"--build-only"},
|
||||
hasPackage: true,
|
||||
want: "[desktop --build-only]",
|
||||
},
|
||||
{
|
||||
name: "help",
|
||||
goos: runtime.GOOS,
|
||||
args: []string{"--help"},
|
||||
hasPackage: true,
|
||||
want: "[desktop --help]",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
tmpDir := t.TempDir()
|
||||
setTestHome(t, tmpDir)
|
||||
withLauncherHooks(t)
|
||||
withInteractiveSession(t, true)
|
||||
withHermesPlatform(t, tt.goos)
|
||||
clearHermesMessagingEnvVars(t)
|
||||
if tt.clearPkgEnv {
|
||||
clearHermesDesktopPackageEnvVars(t)
|
||||
}
|
||||
t.Setenv("PATH", tmpDir+string(os.PathListSeparator)+os.Getenv("PATH"))
|
||||
if tt.hasPackage {
|
||||
writeHermesDesktopPackage(t, tmpDir)
|
||||
}
|
||||
writeHermesDesktopTestBinary(t, tmpDir)
|
||||
|
||||
DefaultConfirmPrompt = func(prompt string, options ConfirmOptions) (bool, error) {
|
||||
t.Fatalf("did not expect messaging prompt during desktop launch: %s", prompt)
|
||||
return false, nil
|
||||
}
|
||||
|
||||
if err := (&HermesDesktop{}).Run("", nil, tt.args); err != nil {
|
||||
t.Fatalf("Run returned error: %v", err)
|
||||
}
|
||||
if got := readHermesDesktopInvocations(t, tmpDir); got != tt.want {
|
||||
t.Fatalf("expected %q, got %q", tt.want, got)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func writeHermesVersionedTestBinary(t *testing.T, dir, version string) {
|
||||
t.Helper()
|
||||
script := "#!/bin/sh\n" +
|
||||
"case \"$1\" in\n" +
|
||||
" --version)\n" +
|
||||
" printf 'Hermes Agent " + version + " (test)\\n'\n" +
|
||||
" ;;\n" +
|
||||
" update)\n" +
|
||||
" printf 'update\\n' >> \"$HOME/hermes-update.log\"\n" +
|
||||
" ;;\n" +
|
||||
" *)\n" +
|
||||
" printf '[%s]\\n' \"$*\" >> \"$HOME/hermes-invocations.log\"\n" +
|
||||
" ;;\n" +
|
||||
"esac\n"
|
||||
bin := filepath.Join(dir, "hermes")
|
||||
if err := os.WriteFile(bin, []byte(script), 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseHermesVersion(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
input string
|
||||
want string
|
||||
}{
|
||||
{"standard release", "Hermes Agent v0.16.0 (2026.6.5)", "v0.16.0"},
|
||||
{"newer release", "Hermes Agent v0.17.0 (2026.6.19)", "v0.17.0"},
|
||||
{"older release", "Hermes Agent v0.15.1 (2026.5.29)", "v0.15.1"},
|
||||
{"prerelease", "Hermes Agent v0.16.0-rc1 (2026.6.5)", "v0.16.0-rc1"},
|
||||
{"no version token", "Hermes Agent", ""},
|
||||
{"empty", "", ""},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if got := parseHermesVersion(tt.input); got != tt.want {
|
||||
t.Fatalf("parseHermesVersion(%q) = %q, want %q", tt.input, got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestHermesDesktopRun_UpdatesCliOlderThanMinVersion(t *testing.T) {
|
||||
if runtime.GOOS == "windows" {
|
||||
t.Skip("uses a POSIX shell test binary")
|
||||
}
|
||||
|
||||
tmpDir := t.TempDir()
|
||||
setTestHome(t, tmpDir)
|
||||
withLauncherHooks(t)
|
||||
withInteractiveSession(t, true)
|
||||
withHermesPlatform(t, runtime.GOOS)
|
||||
clearHermesMessagingEnvVars(t)
|
||||
clearHermesDesktopPackageEnvVars(t)
|
||||
t.Setenv("PATH", tmpDir+string(os.PathListSeparator)+os.Getenv("PATH"))
|
||||
writeHermesVersionedTestBinary(t, tmpDir, "v0.15.1")
|
||||
|
||||
DefaultConfirmPrompt = func(prompt string, options ConfirmOptions) (bool, error) {
|
||||
t.Fatalf("did not expect messaging prompt during desktop launch: %s", prompt)
|
||||
return false, nil
|
||||
}
|
||||
|
||||
if err := (&HermesDesktop{}).Run("", nil, []string{"--foreground"}); err != nil {
|
||||
t.Fatalf("Run returned error: %v", err)
|
||||
}
|
||||
|
||||
updateLog, err := os.ReadFile(filepath.Join(tmpDir, "hermes-update.log"))
|
||||
if err != nil {
|
||||
t.Fatalf("expected hermes update to run for an older CLI: %v", err)
|
||||
}
|
||||
if strings.TrimSpace(string(updateLog)) != "update" {
|
||||
t.Fatalf("expected update log 'update', got %q", updateLog)
|
||||
}
|
||||
if got := readHermesDesktopInvocations(t, tmpDir); got != "[desktop --foreground]" {
|
||||
t.Fatalf("expected desktop launch after update, got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestHermesDesktopRun_SkipsMinVersionCheckOnWindows(t *testing.T) {
|
||||
if runtime.GOOS == "windows" {
|
||||
t.Skip("uses a POSIX shell test binary")
|
||||
}
|
||||
|
||||
tmpDir := t.TempDir()
|
||||
setTestHome(t, tmpDir)
|
||||
withLauncherHooks(t)
|
||||
withInteractiveSession(t, true)
|
||||
withHermesPlatform(t, "windows")
|
||||
clearHermesMessagingEnvVars(t)
|
||||
clearHermesDesktopPackageEnvVars(t)
|
||||
t.Setenv("PATH", tmpDir+string(os.PathListSeparator)+os.Getenv("PATH"))
|
||||
writeHermesVersionedTestBinary(t, tmpDir, "v0.15.1")
|
||||
|
||||
DefaultConfirmPrompt = func(prompt string, options ConfirmOptions) (bool, error) {
|
||||
t.Fatalf("did not expect messaging prompt during desktop launch: %s", prompt)
|
||||
return false, nil
|
||||
}
|
||||
|
||||
if err := (&HermesDesktop{}).Run("", nil, []string{"--foreground"}); err != nil {
|
||||
t.Fatalf("Run returned error: %v", err)
|
||||
}
|
||||
|
||||
if _, err := os.Stat(filepath.Join(tmpDir, "hermes-update.log")); err == nil {
|
||||
t.Fatal("expected hermes update NOT to run on Windows, but hermes-update.log exists")
|
||||
}
|
||||
if got := readHermesDesktopInvocations(t, tmpDir); got != "[desktop --foreground]" {
|
||||
t.Fatalf("expected desktop launch without update, got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestHermesDesktopRun_DoesNotUpdateCliAtOrAboveMinVersion(t *testing.T) {
|
||||
if runtime.GOOS == "windows" {
|
||||
t.Skip("uses a POSIX shell test binary")
|
||||
}
|
||||
|
||||
tmpDir := t.TempDir()
|
||||
setTestHome(t, tmpDir)
|
||||
withLauncherHooks(t)
|
||||
withInteractiveSession(t, true)
|
||||
withHermesPlatform(t, runtime.GOOS)
|
||||
clearHermesMessagingEnvVars(t)
|
||||
clearHermesDesktopPackageEnvVars(t)
|
||||
t.Setenv("PATH", tmpDir+string(os.PathListSeparator)+os.Getenv("PATH"))
|
||||
writeHermesVersionedTestBinary(t, tmpDir, "v0.17.0")
|
||||
|
||||
DefaultConfirmPrompt = func(prompt string, options ConfirmOptions) (bool, error) {
|
||||
t.Fatalf("did not expect messaging prompt during desktop launch: %s", prompt)
|
||||
return false, nil
|
||||
}
|
||||
|
||||
if err := (&HermesDesktop{}).Run("", nil, []string{"--foreground"}); err != nil {
|
||||
t.Fatalf("Run returned error: %v", err)
|
||||
}
|
||||
|
||||
if _, err := os.Stat(filepath.Join(tmpDir, "hermes-update.log")); err == nil {
|
||||
t.Fatal("expected hermes update NOT to run for a current CLI, but hermes-update.log exists")
|
||||
}
|
||||
if got := readHermesDesktopInvocations(t, tmpDir); got != "[desktop --foreground]" {
|
||||
t.Fatalf("expected desktop launch without update, got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestHermesDesktopRunUsesWindowsLocalAppDataPackage(t *testing.T) {
|
||||
tmpDir := t.TempDir()
|
||||
setTestHome(t, tmpDir)
|
||||
withHermesPlatform(t, "windows")
|
||||
t.Setenv("LOCALAPPDATA", filepath.Join(tmpDir, "LocalAppData"))
|
||||
|
||||
writeHermesDesktopExecutable(t,
|
||||
filepath.Join(tmpDir, "LocalAppData", "hermes", "hermes-agent", "apps", "desktop", "release"),
|
||||
hermesDesktopTestExecutableRelativePath("windows"),
|
||||
)
|
||||
|
||||
got := (&HermesDesktop{}).launchArgs([]string{"--cwd", `C:\Users\me\project`})
|
||||
want := []string{"desktop", "--skip-build", "--cwd", `C:\Users\me\project`}
|
||||
if diff := compareStrings(got, want); diff != "" {
|
||||
t.Fatalf("Hermes Desktop launch args mismatch: %s", diff)
|
||||
}
|
||||
}
|
||||
|
||||
func TestHermesDesktopReleaseRootsIncludeLinuxRootInstall(t *testing.T) {
|
||||
tmpDir := t.TempDir()
|
||||
setTestHome(t, tmpDir)
|
||||
withHermesPlatform(t, "linux")
|
||||
|
||||
got := hermesDesktopReleaseRoots()
|
||||
want := filepath.Join(string(filepath.Separator), "usr", "local", "lib", "hermes-agent", "apps", "desktop", "release")
|
||||
if !slices.Contains(got, want) {
|
||||
t.Fatalf("expected Linux root install release path %q in %v", want, got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestHermesRun_PromptsForMessagingSetupBeforeDefaultLaunch(t *testing.T) {
|
||||
if runtime.GOOS == "windows" {
|
||||
t.Skip("uses a POSIX shell test binary")
|
||||
@@ -943,26 +1283,59 @@ func TestHermesMessagingConfiguredRecognizesSupportedGatewayVars(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestHermesEnsureInstalledWindowsShowsWSLGuidance(t *testing.T) {
|
||||
func TestHermesEnsureInstalledWindowsRunsPowerShellInstaller(t *testing.T) {
|
||||
if runtime.GOOS == "windows" {
|
||||
t.Skip("uses a POSIX shell test binary")
|
||||
}
|
||||
|
||||
tmpDir := t.TempDir()
|
||||
setTestHome(t, tmpDir)
|
||||
withLauncherHooks(t)
|
||||
withHermesPlatform(t, "windows")
|
||||
t.Setenv("PATH", tmpDir)
|
||||
t.Setenv("LOCALAPPDATA", filepath.Join(tmpDir, "AppData", "Local"))
|
||||
|
||||
powershell := filepath.Join(tmpDir, "powershell.exe")
|
||||
script := fmt.Sprintf(`#!/bin/sh
|
||||
printf '%%s\n' "$*" >> %q
|
||||
/bin/mkdir -p %q
|
||||
/bin/cat > %q <<'EOS'
|
||||
#!/bin/sh
|
||||
exit 0
|
||||
EOS
|
||||
/bin/chmod +x %q
|
||||
exit 0
|
||||
`,
|
||||
filepath.Join(tmpDir, "powershell.log"),
|
||||
filepath.Dir(filepath.Join(tmpDir, "AppData", "Local", "hermes", "hermes-agent", "venv", "Scripts", "hermes.exe")),
|
||||
filepath.Join(tmpDir, "AppData", "Local", "hermes", "hermes-agent", "venv", "Scripts", "hermes.exe"),
|
||||
filepath.Join(tmpDir, "AppData", "Local", "hermes", "hermes-agent", "venv", "Scripts", "hermes.exe"),
|
||||
)
|
||||
if err := os.WriteFile(powershell, []byte(script), 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
DefaultConfirmPrompt = func(prompt string, options ConfirmOptions) (bool, error) {
|
||||
if prompt != "Hermes is not installed. Install now?" {
|
||||
t.Fatalf("unexpected install prompt %q", prompt)
|
||||
}
|
||||
return true, nil
|
||||
}
|
||||
|
||||
h := &Hermes{}
|
||||
err := h.ensureInstalled()
|
||||
if err == nil {
|
||||
t.Fatal("expected WSL guidance error")
|
||||
if err := h.ensureInstalled(); err != nil {
|
||||
t.Fatalf("ensureInstalled returned error: %v", err)
|
||||
}
|
||||
msg := err.Error()
|
||||
if !strings.Contains(msg, "wsl --install") {
|
||||
t.Fatalf("expected install command in guidance, got %v", err)
|
||||
|
||||
data, err := os.ReadFile(filepath.Join(tmpDir, "powershell.log"))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !strings.Contains(msg, "hermes-agent.nousresearch.com") {
|
||||
t.Fatalf("expected docs link in guidance, got %v", err)
|
||||
}
|
||||
if strings.Contains(msg, "hermes is not installed") {
|
||||
t.Fatalf("guidance should not lead with 'hermes is not installed', got %v", err)
|
||||
logs := string(data)
|
||||
for _, want := range []string{"-NoProfile", "-ExecutionPolicy", "Bypass", "-Command", hermesWindowsInstallURL, "-SkipSetup"} {
|
||||
if !strings.Contains(logs, want) {
|
||||
t.Fatalf("expected PowerShell installer args to contain %q, got logs:\n%s", want, logs)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -58,12 +58,15 @@ func TestIntegrationLookup(t *testing.T) {
|
||||
{"claude desktop", "claude-desktop", true, "Claude Desktop"},
|
||||
{"claude desktop alias", "claude-app", true, "Claude Desktop"},
|
||||
{"codex", "codex", true, "Codex"},
|
||||
{"codex app", "codex-app", true, "Codex App"},
|
||||
{"codex app desktop alias", "codex-desktop", true, "Codex App"},
|
||||
{"codex app gui alias", "codex-gui", true, "Codex App"},
|
||||
{"chatgpt", "chatgpt", true, "ChatGPT"},
|
||||
{"codex app legacy alias", "codex-app", true, "ChatGPT"},
|
||||
{"codex app desktop alias", "codex-desktop", true, "ChatGPT"},
|
||||
{"codex app gui alias", "codex-gui", true, "ChatGPT"},
|
||||
{"hermes desktop", "hermes-desktop", true, "Hermes Desktop"},
|
||||
{"kimi", "kimi", true, "Kimi Code CLI"},
|
||||
{"droid", "droid", true, "Droid"},
|
||||
{"opencode", "opencode", true, "OpenCode"},
|
||||
{"omp", "omp", true, "OMP"},
|
||||
{"pool", "pool", true, "Pool"},
|
||||
{"unknown integration", "unknown", false, ""},
|
||||
{"empty string", "", false, ""},
|
||||
@@ -83,7 +86,7 @@ func TestIntegrationLookup(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestIntegrationRegistry(t *testing.T) {
|
||||
expectedIntegrations := []string{"claude", "claude-desktop", "cline", "codex", "codex-app", "kimi", "droid", "opencode", "hermes", "pool"}
|
||||
expectedIntegrations := []string{"claude", "claude-desktop", "cline", "codex", "chatgpt", "kimi", "droid", "opencode", "omp", "hermes", "hermes-desktop", "pool", "qwen"}
|
||||
for _, name := range expectedIntegrations {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
r, ok := integrations[name]
|
||||
@@ -97,6 +100,30 @@ func TestIntegrationRegistry(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestChatGPTMigratesLegacyCodexAppLaunchConfig(t *testing.T) {
|
||||
setTestHome(t, t.TempDir())
|
||||
if err := config.SaveIntegration(codexAppIntegrationName, []string{"qwen3.5"}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := config.MarkIntegrationOnboarded(codexAppIntegrationName); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
got, err := loadStoredIntegrationConfig(chatGPTIntegrationName)
|
||||
if err != nil {
|
||||
t.Fatalf("loadStoredIntegrationConfig returned error: %v", err)
|
||||
}
|
||||
if diff := compareStrings(got.Models, []string{"qwen3.5"}); diff != "" {
|
||||
t.Fatalf("migrated models mismatch: %s", diff)
|
||||
}
|
||||
if !got.Onboarded {
|
||||
t.Fatal("migrated integration should remain onboarded")
|
||||
}
|
||||
if _, err := config.LoadIntegration(chatGPTIntegrationName); err != nil {
|
||||
t.Fatalf("canonical ChatGPT config was not written: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestHiddenIntegrationsExcludedFromVisibleLists(t *testing.T) {
|
||||
for _, info := range ListIntegrationInfos() {
|
||||
switch info.Name {
|
||||
@@ -1079,6 +1106,51 @@ func TestShowOrPullWithPolicy_CloudModelNotFound_FailsEarlyForAllPolicies(t *tes
|
||||
}
|
||||
}
|
||||
|
||||
func TestShowOrPullWithPolicy_CloudModelShowUnavailableAllowsSelection(t *testing.T) {
|
||||
oldHook := DefaultConfirmPrompt
|
||||
DefaultConfirmPrompt = func(prompt string, options ConfirmOptions) (bool, error) {
|
||||
t.Fatal("confirm prompt should not be called for explicit cloud models")
|
||||
return false, nil
|
||||
}
|
||||
defer func() { DefaultConfirmPrompt = oldHook }()
|
||||
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
switch r.URL.Path {
|
||||
case "/api/show":
|
||||
w.WriteHeader(http.StatusInternalServerError)
|
||||
fmt.Fprintf(w, `{"error":"temporary failure"}`)
|
||||
case "/api/status":
|
||||
w.WriteHeader(http.StatusInternalServerError)
|
||||
fmt.Fprintf(w, `{"error":"temporary failure"}`)
|
||||
case "/api/pull":
|
||||
t.Fatal("pull should not be called for explicit cloud models")
|
||||
default:
|
||||
w.WriteHeader(http.StatusNotFound)
|
||||
}
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
u, _ := url.Parse(srv.URL)
|
||||
client := api.NewClient(u, srv.Client())
|
||||
|
||||
if err := showOrPullWithPolicy(context.Background(), client, "glm-5.1:cloud", missingModelFail, true); err != nil {
|
||||
t.Fatalf("showOrPullWithPolicy returned error: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestShowOrPullWithPolicy_CloudModelShowUnreachableAllowsSelection(t *testing.T) {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
t.Fatalf("unexpected request after server close: %s %s", r.Method, r.URL.Path)
|
||||
}))
|
||||
u, _ := url.Parse(srv.URL)
|
||||
client := api.NewClient(u, srv.Client())
|
||||
srv.Close()
|
||||
|
||||
if err := showOrPullWithPolicy(context.Background(), client, "glm-5.1:cloud", missingModelFail, true); err != nil {
|
||||
t.Fatalf("showOrPullWithPolicy returned error: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestShowOrPullWithPolicy_CloudModelDisabled_FailsWithCloudDisabledError(t *testing.T) {
|
||||
oldHook := DefaultConfirmPrompt
|
||||
DefaultConfirmPrompt = func(prompt string, options ConfirmOptions) (bool, error) {
|
||||
@@ -1741,9 +1813,9 @@ func TestIntegration_InstallHint(t *testing.T) {
|
||||
wantURL: "https://developers.openai.com/codex/cli/",
|
||||
},
|
||||
{
|
||||
name: "codex app has hint",
|
||||
input: "codex-app",
|
||||
wantURL: "https://developers.openai.com/codex/quickstart",
|
||||
name: "chatgpt has hint",
|
||||
input: "chatgpt",
|
||||
wantURL: "https://chatgpt.com/download",
|
||||
},
|
||||
{
|
||||
name: "openclaw has hint",
|
||||
@@ -1829,7 +1901,7 @@ func TestListIntegrationInfos(t *testing.T) {
|
||||
if codexAppSupported() != nil {
|
||||
filtered := make([]string, 0, len(want))
|
||||
for _, name := range want {
|
||||
if name != "codex-app" {
|
||||
if name != "chatgpt" {
|
||||
filtered = append(filtered, name)
|
||||
}
|
||||
}
|
||||
@@ -1846,9 +1918,9 @@ func TestListIntegrationInfos(t *testing.T) {
|
||||
for _, info := range infos {
|
||||
got = append(got, info.Name)
|
||||
}
|
||||
wantPrefix := []string{"claude", "codex-app", "hermes", "openclaw"}
|
||||
wantPrefix := []string{"claude", "chatgpt", "hermes", "openclaw", "opencode", "hermes-desktop", "codex", "copilot", "omp"}
|
||||
if codexAppSupported() != nil {
|
||||
wantPrefix = []string{"claude", "hermes", "openclaw", "opencode"}
|
||||
wantPrefix = []string{"claude", "hermes", "openclaw", "opencode", "hermes-desktop", "codex", "copilot", "omp"}
|
||||
}
|
||||
if len(got) < len(wantPrefix) {
|
||||
t.Fatalf("expected at least %d integrations, got %v", len(wantPrefix), got)
|
||||
@@ -1870,9 +1942,9 @@ func TestListIntegrationInfos(t *testing.T) {
|
||||
})
|
||||
|
||||
t.Run("includes known integrations", func(t *testing.T) {
|
||||
known := map[string]bool{"claude": false, "cline": false, "codex": false, "opencode": false}
|
||||
known := map[string]bool{"claude": false, "cline": false, "codex": false, "opencode": false, "omp": false}
|
||||
if codexAppSupported() == nil {
|
||||
known["codex-app"] = false
|
||||
known["chatgpt"] = false
|
||||
}
|
||||
if poolsideGOOS != "windows" {
|
||||
known["pool"] = false
|
||||
@@ -1898,6 +1970,15 @@ func TestListIntegrationInfos(t *testing.T) {
|
||||
t.Fatal("expected hermes to be included in ListIntegrationInfos")
|
||||
})
|
||||
|
||||
t.Run("includes hermes desktop", func(t *testing.T) {
|
||||
for _, info := range infos {
|
||||
if info.Name == "hermes-desktop" {
|
||||
return
|
||||
}
|
||||
}
|
||||
t.Fatal("expected hermes-desktop to be included in ListIntegrationInfos")
|
||||
})
|
||||
|
||||
t.Run("hermes still resolves explicitly", func(t *testing.T) {
|
||||
name, runner, err := LookupIntegration("hermes")
|
||||
if err != nil {
|
||||
@@ -1996,6 +2077,7 @@ func TestIntegration_Editor(t *testing.T) {
|
||||
{"claude", false},
|
||||
{"claude-desktop", false},
|
||||
{"codex", false},
|
||||
{"omp", false},
|
||||
{"nonexistent", false},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
@@ -2020,12 +2102,14 @@ func TestIntegration_AutoInstallable(t *testing.T) {
|
||||
{"openclaw", true},
|
||||
{"pi", true},
|
||||
{"hermes", true},
|
||||
{"hermes-desktop", true},
|
||||
{"cline", true},
|
||||
{"qwen", true},
|
||||
{"claude", false},
|
||||
{"claude", true},
|
||||
{"claude-desktop", false},
|
||||
{"codex", false},
|
||||
{"opencode", false},
|
||||
{"opencode", true},
|
||||
{"omp", false},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
|
||||
+99
-22
@@ -287,12 +287,14 @@ Flags and extra arguments require an integration name.
|
||||
|
||||
Supported integrations:
|
||||
claude Claude Code
|
||||
codex-app Codex App (aliases: codex-desktop, codex-gui)
|
||||
chatgpt ChatGPT (aliases: codex-app, codex-desktop, codex-gui)
|
||||
hermes Hermes Agent
|
||||
openclaw OpenClaw (aliases: clawdbot, moltbot)
|
||||
opencode OpenCode
|
||||
codex Codex
|
||||
hermes-desktop Hermes Desktop
|
||||
copilot Copilot CLI (aliases: copilot-cli)
|
||||
omp OMP
|
||||
droid Droid
|
||||
kimi Kimi Code CLI
|
||||
pi Pi
|
||||
@@ -305,9 +307,10 @@ Examples:
|
||||
ollama launch
|
||||
ollama launch claude
|
||||
ollama launch claude --model <model>
|
||||
ollama launch codex-app
|
||||
ollama launch codex-app --restore
|
||||
ollama launch chatgpt
|
||||
ollama launch chatgpt --restore
|
||||
ollama launch hermes
|
||||
ollama launch hermes-desktop
|
||||
ollama launch droid --config (does not auto-launch)
|
||||
ollama launch codex --restore
|
||||
ollama launch codex -- --sandbox workspace-write`,
|
||||
@@ -704,9 +707,12 @@ func (c *launcherClient) resolveRunModel(ctx context.Context, req RunModelReques
|
||||
}
|
||||
if usable {
|
||||
if err := c.ensureModelsReady(ctx, []string{current}); err != nil {
|
||||
return "", err
|
||||
if !errors.Is(err, errDeprecatedLaunchModelDeclined) {
|
||||
return "", err
|
||||
}
|
||||
} else {
|
||||
return current, nil
|
||||
}
|
||||
return current, nil
|
||||
}
|
||||
}
|
||||
|
||||
@@ -723,7 +729,7 @@ func (c *launcherClient) resolveRunModel(ctx context.Context, req RunModelReques
|
||||
}
|
||||
|
||||
func (c *launcherClient) launchSingleIntegration(ctx context.Context, name string, runner Runner, saved *config.IntegrationConfig, req IntegrationLaunchRequest) error {
|
||||
target, _, err := c.resolveSingleIntegrationTarget(ctx, runner, primaryModelFromConfig(saved), req)
|
||||
target, _, err := c.resolveSingleIntegrationTarget(ctx, name, runner, primaryModelFromConfig(saved), req)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -745,14 +751,22 @@ func (c *launcherClient) launchEditorIntegration(ctx context.Context, name strin
|
||||
models, needsConfigure := c.resolveEditorLaunchModels(ctx, saved, req)
|
||||
|
||||
if needsConfigure {
|
||||
selected, err := c.selectMultiModelsForIntegration(ctx, runner, models)
|
||||
selected, err := c.selectMultiModelsForIntegration(ctx, name, runner, models)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
models = selected
|
||||
} else if len(models) > 0 {
|
||||
if err := c.ensureModelsReady(ctx, models[:1]); err != nil {
|
||||
return err
|
||||
if err := c.ensureModelsReadyFor(ctx, models[:1], runner.String(), name); err != nil {
|
||||
if !errors.Is(err, errDeprecatedLaunchModelDeclined) || req.ModelOverride != "" {
|
||||
return err
|
||||
}
|
||||
selected, err := c.selectMultiModelsForIntegration(ctx, name, runner, models)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
models = selected
|
||||
needsConfigure = true
|
||||
}
|
||||
}
|
||||
|
||||
@@ -761,7 +775,8 @@ func (c *launcherClient) launchEditorIntegration(ctx context.Context, name strin
|
||||
}
|
||||
|
||||
var launchModels []LaunchModel
|
||||
if (needsConfigure || req.ModelOverride != "") && !savedMatchesModels(saved, models) {
|
||||
liveConfigMatches := slices.Equal(editor.Models(), models)
|
||||
if needsConfigure || req.ModelOverride != "" || !savedMatchesModels(saved, models) || !liveConfigMatches {
|
||||
launchModels = c.modelInventory().Resolve(ctx, models)
|
||||
if err := prepareEditorIntegration(name, editor, launchModels); err != nil {
|
||||
return err
|
||||
@@ -780,7 +795,7 @@ func (c *launcherClient) launchManagedSingleIntegration(ctx context.Context, nam
|
||||
selectionCurrent = primaryModelFromConfig(saved)
|
||||
}
|
||||
|
||||
target, needsConfigure, err := c.resolveSingleIntegrationTarget(ctx, runner, selectionCurrent, req)
|
||||
target, needsConfigure, err := c.resolveSingleIntegrationTarget(ctx, name, runner, selectionCurrent, req)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -950,7 +965,7 @@ func (c *launcherClient) managedSingleConfigureModels(ctx context.Context, manag
|
||||
return dedupeModelList(models), nil
|
||||
}
|
||||
|
||||
func (c *launcherClient) resolveSingleIntegrationTarget(ctx context.Context, runner Runner, current string, req IntegrationLaunchRequest) (string, bool, error) {
|
||||
func (c *launcherClient) resolveSingleIntegrationTarget(ctx context.Context, name string, runner Runner, current string, req IntegrationLaunchRequest) (string, bool, error) {
|
||||
target := req.ModelOverride
|
||||
needsConfigure := req.ForceConfigure
|
||||
skipReadiness := false
|
||||
@@ -974,14 +989,24 @@ func (c *launcherClient) resolveSingleIntegrationTarget(ctx context.Context, run
|
||||
}
|
||||
|
||||
if needsConfigure && req.ModelOverride == "" {
|
||||
selected, err := c.selectSingleModelWithSelectorReady(ctx, fmt.Sprintf("Select model for %s:", runner), target, DefaultSingleSelector, !skipReadiness)
|
||||
selected, err := c.selectSingleModelWithSelectorReady(ctx, fmt.Sprintf("Select model for %s:", runner), target, DefaultSingleSelector, !skipReadiness, runner.String(), name)
|
||||
if err != nil {
|
||||
return "", false, err
|
||||
}
|
||||
target = selected
|
||||
} else if !skipReadiness {
|
||||
if err := c.ensureModelsReady(ctx, []string{target}); err != nil {
|
||||
return "", false, err
|
||||
if err := c.ensureModelsReadyFor(ctx, []string{target}, runner.String(), name); err != nil {
|
||||
if !errors.Is(err, errDeprecatedLaunchModelDeclined) {
|
||||
return "", false, err
|
||||
}
|
||||
// "Pick another model" is an interactive recovery path, including
|
||||
// when --model supplied the initial target.
|
||||
selected, err := c.selectSingleModelWithSelectorReady(ctx, fmt.Sprintf("Select model for %s:", runner), target, DefaultSingleSelector, true, runner.String(), name)
|
||||
if err != nil {
|
||||
return "", false, err
|
||||
}
|
||||
target = selected
|
||||
needsConfigure = true
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1019,7 +1044,7 @@ func managedRequiresInteractiveOnboarding(managed any) bool {
|
||||
}
|
||||
|
||||
func (c *launcherClient) selectSingleModelWithSelector(ctx context.Context, title, current string, selector SingleSelector) (string, error) {
|
||||
return c.selectSingleModelWithSelectorReady(ctx, title, current, selector, true)
|
||||
return c.selectSingleModelWithSelectorReady(ctx, title, current, selector, true, "ollama launch", "")
|
||||
}
|
||||
|
||||
func (c *launcherClient) latestAccountState() *AccountState {
|
||||
@@ -1029,7 +1054,7 @@ func (c *launcherClient) latestAccountState() *AccountState {
|
||||
return c.accountState
|
||||
}
|
||||
|
||||
func (c *launcherClient) selectSingleModelWithSelectorReady(ctx context.Context, title, current string, selector SingleSelector, ensureReady bool) (string, error) {
|
||||
func (c *launcherClient) selectSingleModelWithSelectorReady(ctx context.Context, title, current string, selector SingleSelector, ensureReady bool, label, commandName string) (string, error) {
|
||||
if selector == nil && DefaultSingleSelectorWithUpdates == nil {
|
||||
return "", fmt.Errorf("no selector configured")
|
||||
}
|
||||
@@ -1054,11 +1079,15 @@ func (c *launcherClient) selectSingleModelWithSelectorReady(ctx context.Context,
|
||||
return "", ErrCancelled
|
||||
}
|
||||
if ensureReady {
|
||||
if err := c.ensureModelsReady(ctx, []string{selected}); err != nil {
|
||||
if err := c.ensureModelsReadyFor(ctx, []string{selected}, label, commandName); err != nil {
|
||||
if errors.Is(err, errUpgradeCancelled) {
|
||||
current = selected
|
||||
continue
|
||||
}
|
||||
if errors.Is(err, errDeprecatedLaunchModelDeclined) {
|
||||
current = selected
|
||||
continue
|
||||
}
|
||||
return "", err
|
||||
}
|
||||
}
|
||||
@@ -1066,7 +1095,7 @@ func (c *launcherClient) selectSingleModelWithSelectorReady(ctx context.Context,
|
||||
}
|
||||
}
|
||||
|
||||
func (c *launcherClient) selectMultiModelsForIntegration(ctx context.Context, runner Runner, preChecked []string) ([]string, error) {
|
||||
func (c *launcherClient) selectMultiModelsForIntegration(ctx context.Context, name string, runner Runner, preChecked []string) ([]string, error) {
|
||||
if DefaultMultiSelector == nil && DefaultMultiSelectorWithUpdates == nil {
|
||||
return nil, fmt.Errorf("no selector configured")
|
||||
}
|
||||
@@ -1088,12 +1117,16 @@ func (c *launcherClient) selectMultiModelsForIntegration(ctx context.Context, ru
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
accepted, skipped, err := c.selectReadyModelsForSave(ctx, selected)
|
||||
accepted, skipped, err := c.selectReadyModelsForSave(ctx, selected, runner.String(), name)
|
||||
if err != nil {
|
||||
if errors.Is(err, errUpgradeCancelled) {
|
||||
orderedChecked = append([]string(nil), selected...)
|
||||
continue
|
||||
}
|
||||
if errors.Is(err, errDeprecatedLaunchModelDeclined) {
|
||||
orderedChecked = append([]string(nil), selected...)
|
||||
continue
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
for _, skip := range skipped {
|
||||
@@ -1132,6 +1165,8 @@ func (c *launcherClient) loadSelectableModels(ctx context.Context, preChecked []
|
||||
|
||||
cloudDisabled, _ := cloudStatusDisabled(ctx, c.apiClient)
|
||||
items, orderedChecked, _, _ := buildModelListWithRecommendations(inventory, recommendations, preChecked, current)
|
||||
items = filterDeprecatedLaunchModelItems(items)
|
||||
orderedChecked = filterDeprecatedLaunchModelNames(orderedChecked)
|
||||
if cloudDisabled {
|
||||
items = filterCloudItems(items)
|
||||
orderedChecked = c.filterDisabledCloudModels(ctx, orderedChecked)
|
||||
@@ -1208,13 +1243,31 @@ func (c *launcherClient) requestRecommendations(ctx context.Context) ([]ModelIte
|
||||
}
|
||||
|
||||
func (c *launcherClient) ensureModelsReady(ctx context.Context, models []string) error {
|
||||
return c.ensureModelsReadyFor(ctx, models, "ollama launch", "")
|
||||
}
|
||||
|
||||
func (c *launcherClient) ensureModelsReadyFor(ctx context.Context, models []string, label, commandName string) error {
|
||||
models = dedupeModelList(models)
|
||||
if len(models) == 0 {
|
||||
return nil
|
||||
}
|
||||
cloudRec, localRec := c.agentCapableRecommendations(ctx)
|
||||
|
||||
cloudModels := make(map[string]bool, len(models))
|
||||
for _, model := range models {
|
||||
if prompt := deprecatedLaunchModelPrompt(model, label, commandName, cloudRec, localRec); prompt != "" {
|
||||
ok, err := ConfirmPromptWithOptions(prompt, ConfirmOptions{
|
||||
YesLabel: "Launch anyway",
|
||||
NoLabel: "Pick another model",
|
||||
Default: ConfirmDefaultNo,
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if !ok {
|
||||
return errDeprecatedLaunchModelDeclined
|
||||
}
|
||||
}
|
||||
isCloudModel := isCloudModelName(model)
|
||||
if isCloudModel {
|
||||
cloudModels[model] = true
|
||||
@@ -1229,6 +1282,27 @@ func (c *launcherClient) ensureModelsReady(ctx context.Context, models []string)
|
||||
return ensureAuth(ctx, c.apiClient, cloudModels, models)
|
||||
}
|
||||
|
||||
func (c *launcherClient) agentCapableRecommendations(ctx context.Context) (cloud, local string) {
|
||||
recs := c.recommendations(ctx)
|
||||
cloudDisabled, known := cloudStatusDisabled(ctx, c.apiClient)
|
||||
for _, rec := range recs {
|
||||
if rec.Name == "" || isDeprecatedLaunchModel(rec.Name) {
|
||||
continue
|
||||
}
|
||||
if isCloudModelName(rec.Name) {
|
||||
if cloud == "" && !(known && cloudDisabled) {
|
||||
cloud = rec.Name
|
||||
}
|
||||
} else if local == "" {
|
||||
local = rec.Name
|
||||
}
|
||||
if cloud != "" && local != "" {
|
||||
break
|
||||
}
|
||||
}
|
||||
return cloud, local
|
||||
}
|
||||
|
||||
func dedupeModelList(models []string) []string {
|
||||
deduped := make([]string, 0, len(models))
|
||||
seen := make(map[string]bool, len(models))
|
||||
@@ -1247,16 +1321,19 @@ type skippedModel struct {
|
||||
reason string
|
||||
}
|
||||
|
||||
func (c *launcherClient) selectReadyModelsForSave(ctx context.Context, selected []string) ([]string, []skippedModel, error) {
|
||||
func (c *launcherClient) selectReadyModelsForSave(ctx context.Context, selected []string, label, commandName string) ([]string, []skippedModel, error) {
|
||||
selected = dedupeModelList(selected)
|
||||
accepted := make([]string, 0, len(selected))
|
||||
skipped := make([]skippedModel, 0, len(selected))
|
||||
|
||||
for _, model := range selected {
|
||||
if err := c.ensureModelsReady(ctx, []string{model}); err != nil {
|
||||
if err := c.ensureModelsReadyFor(ctx, []string{model}, label, commandName); err != nil {
|
||||
if errors.Is(err, errUpgradeCancelled) {
|
||||
return nil, nil, err
|
||||
}
|
||||
if errors.Is(err, errDeprecatedLaunchModelDeclined) {
|
||||
return nil, nil, err
|
||||
}
|
||||
if errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) {
|
||||
return nil, nil, err
|
||||
}
|
||||
|
||||
+608
-91
File diff suppressed because it is too large.
Load diff
+15
-11
@@ -194,10 +194,10 @@ func ensureCloudAuth(ctx context.Context, client *api.Client, modelList string)
|
||||
}
|
||||
|
||||
var aErr api.AuthorizationError
|
||||
if !errors.As(err, &aErr) || aErr.SigninURL == "" {
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err != nil && !errors.As(err, &aErr) {
|
||||
return nil
|
||||
}
|
||||
if err == nil || aErr.SigninURL == "" {
|
||||
return fmt.Errorf("%s requires sign in", modelList)
|
||||
}
|
||||
|
||||
@@ -258,19 +258,23 @@ func showOrPullWithPolicy(ctx context.Context, client *api.Client, model string,
|
||||
if _, err := client.Show(ctx, &api.ShowRequest{Model: model}); err == nil {
|
||||
return nil
|
||||
} else {
|
||||
if isCloudModel {
|
||||
if disabled, known := cloudStatusDisabled(ctx, client); known && disabled {
|
||||
return errors.New(internalcloud.DisabledError("remote inference is unavailable"))
|
||||
}
|
||||
var statusErr api.StatusError
|
||||
if errors.As(err, &statusErr) && statusErr.StatusCode == http.StatusNotFound {
|
||||
return fmt.Errorf("model %q not found", model)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
var statusErr api.StatusError
|
||||
if !errors.As(err, &statusErr) || statusErr.StatusCode != http.StatusNotFound {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
if isCloudModel {
|
||||
if disabled, known := cloudStatusDisabled(ctx, client); known && disabled {
|
||||
return errors.New(internalcloud.DisabledError("remote inference is unavailable"))
|
||||
}
|
||||
return fmt.Errorf("model %q not found", model)
|
||||
}
|
||||
|
||||
switch policy {
|
||||
case missingModelAutoPull:
|
||||
return pullMissingModel(ctx, client, model)
|
||||
|
||||
@@ -0,0 +1,454 @@
|
||||
package launch
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"os/exec"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"slices"
|
||||
"strings"
|
||||
|
||||
"github.com/ollama/ollama/cmd/config"
|
||||
"github.com/ollama/ollama/cmd/internal/fileutil"
|
||||
"github.com/ollama/ollama/envconfig"
|
||||
"github.com/ollama/ollama/types/model"
|
||||
"gopkg.in/yaml.v3"
|
||||
)
|
||||
|
||||
const (
|
||||
ompIntegrationName = "omp"
|
||||
ompProviderName = "ollama"
|
||||
ompSetupVersion = 1
|
||||
ompWebSearchPlugin = "@ollama/pi-web-search"
|
||||
)
|
||||
|
||||
// OMP implements Runner for the OMP coding-agent integration.
|
||||
type OMP struct{}
|
||||
|
||||
func (o *OMP) String() string { return "OMP" }
|
||||
|
||||
func (o *OMP) Paths() []string {
|
||||
var paths []string
|
||||
for _, pathFn := range []func() (string, error){ompModelsPath, ompConfigPath} {
|
||||
path, err := pathFn()
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
if _, err := os.Stat(path); err == nil {
|
||||
paths = append(paths, path)
|
||||
}
|
||||
}
|
||||
return paths
|
||||
}
|
||||
|
||||
func (o *OMP) Configure(model string) error {
|
||||
return o.ConfigureWithModels(model, []LaunchModel{fallbackLaunchModel(model)})
|
||||
}
|
||||
|
||||
func (o *OMP) ConfigureWithModels(primary string, models []LaunchModel) error {
|
||||
if primary == "" {
|
||||
return nil
|
||||
}
|
||||
if len(models) == 0 {
|
||||
models = []LaunchModel{fallbackLaunchModel(primary)}
|
||||
}
|
||||
if err := writeOMPModelsConfig(primary, models); err != nil {
|
||||
return err
|
||||
}
|
||||
return writeOMPAgentConfig()
|
||||
}
|
||||
|
||||
func (o *OMP) CurrentModel() string {
|
||||
cfg, err := readOMPModelsConfig()
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
provider, ok := ompProvider(cfg)
|
||||
if !ok {
|
||||
return ""
|
||||
}
|
||||
if !ompProviderHealthy(provider) {
|
||||
return ""
|
||||
}
|
||||
models, _ := provider["models"].([]any)
|
||||
for _, raw := range models {
|
||||
entry, ok := raw.(map[string]any)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
if id, _ := entry["id"].(string); id != "" {
|
||||
return id
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (o *OMP) Onboard() error {
|
||||
return config.MarkIntegrationOnboarded(ompIntegrationName)
|
||||
}
|
||||
|
||||
func (o *OMP) RequiresInteractiveOnboarding() bool { return false }
|
||||
|
||||
func (o *OMP) args(model string, extra []string) []string {
|
||||
var args []string
|
||||
if model != "" {
|
||||
args = append(args, "--model", ompModelName(model))
|
||||
}
|
||||
args = append(args, extra...)
|
||||
return args
|
||||
}
|
||||
|
||||
func ompModelName(model string) string {
|
||||
if strings.HasPrefix(model, "ollama/") {
|
||||
return model
|
||||
}
|
||||
return "ollama/" + model
|
||||
}
|
||||
|
||||
func (o *OMP) findPath() (string, error) {
|
||||
if p, err := exec.LookPath("omp"); err == nil {
|
||||
return p, nil
|
||||
}
|
||||
home, err := os.UserHomeDir()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
for _, dir := range []string{
|
||||
filepath.Join(home, ".local", "bin"),
|
||||
filepath.Join(home, ".bun", "bin"),
|
||||
} {
|
||||
for _, name := range ompExecutableNames() {
|
||||
fallback := filepath.Join(dir, name)
|
||||
if _, err := os.Stat(fallback); err == nil {
|
||||
return fallback, nil
|
||||
}
|
||||
}
|
||||
}
|
||||
return "", exec.ErrNotFound
|
||||
}
|
||||
|
||||
func ompExecutableNames() []string {
|
||||
if runtime.GOOS == "windows" {
|
||||
return []string{"omp.exe", "omp.cmd", "omp.bat"}
|
||||
}
|
||||
return []string{"omp"}
|
||||
}
|
||||
|
||||
func (o *OMP) Run(model string, _ []LaunchModel, args []string) error {
|
||||
ompPath, err := o.findPath()
|
||||
if err != nil {
|
||||
return fmt.Errorf("omp is not installed, install from https://omp.sh")
|
||||
}
|
||||
|
||||
ensureOMPWebSearchPlugin(ompPath)
|
||||
|
||||
cmd := exec.Command(ompPath, o.args(model, args)...)
|
||||
cmd.Stdin = os.Stdin
|
||||
cmd.Stdout = os.Stdout
|
||||
cmd.Stderr = os.Stderr
|
||||
cmd.Env = os.Environ()
|
||||
return cmd.Run()
|
||||
}
|
||||
|
||||
func ensureOMPWebSearchPlugin(bin string) {
|
||||
if !shouldManageOllamaWebSearch() {
|
||||
fmt.Fprintf(os.Stderr, "%sCloud is disabled; skipping %s setup.%s\n", ansiGray, ompWebSearchPlugin, ansiReset)
|
||||
return
|
||||
}
|
||||
|
||||
fmt.Fprintf(os.Stderr, "%sChecking OMP web search plugin...%s\n", ansiGray, ansiReset)
|
||||
|
||||
installed, err := ompPluginInstalled(bin, ompWebSearchPlugin)
|
||||
if err != nil {
|
||||
fmt.Fprintf(os.Stderr, "%s Warning: could not check %s installation: %v%s\n", ansiYellow, ompWebSearchPlugin, err, ansiReset)
|
||||
return
|
||||
}
|
||||
|
||||
verb := "Installing"
|
||||
warnVerb := "install"
|
||||
doneVerb := "Installed"
|
||||
if installed {
|
||||
verb = "Updating"
|
||||
warnVerb = "update"
|
||||
doneVerb = "Updated"
|
||||
}
|
||||
|
||||
fmt.Fprintf(os.Stderr, "%s%s %s...%s\n", ansiGray, verb, ompWebSearchPlugin, ansiReset)
|
||||
cmd := exec.Command(bin, "plugin", "install", ompWebSearchPlugin)
|
||||
cmd.Stdout = os.Stdout
|
||||
cmd.Stderr = os.Stderr
|
||||
if err := cmd.Run(); err != nil {
|
||||
fmt.Fprintf(os.Stderr, "%s Warning: could not %s %s: %v%s\n", ansiYellow, warnVerb, ompWebSearchPlugin, err, ansiReset)
|
||||
return
|
||||
}
|
||||
|
||||
fmt.Fprintf(os.Stderr, "%s ✓ %s %s%s\n", ansiGreen, doneVerb, ompWebSearchPlugin, ansiReset)
|
||||
}
|
||||
|
||||
func ompPluginInstalled(bin, plugin string) (bool, error) {
|
||||
cmd := exec.Command(bin, "plugin", "list")
|
||||
out, err := cmd.CombinedOutput()
|
||||
if err != nil {
|
||||
msg := strings.TrimSpace(string(out))
|
||||
if msg == "" {
|
||||
return false, err
|
||||
}
|
||||
return false, fmt.Errorf("%w: %s", err, msg)
|
||||
}
|
||||
|
||||
versioned := plugin + "@"
|
||||
for _, line := range strings.Split(string(out), "\n") {
|
||||
trimmed := strings.TrimSpace(line)
|
||||
if strings.Contains(trimmed, versioned) || trimmed == plugin {
|
||||
return true, nil
|
||||
}
|
||||
}
|
||||
return false, nil
|
||||
}
|
||||
|
||||
func ompModelsPath() (string, error) {
|
||||
dir, err := ompAgentDir()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return filepath.Join(dir, "models.yml"), nil
|
||||
}
|
||||
|
||||
func ompConfigPath() (string, error) {
|
||||
dir, err := ompAgentDir()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return filepath.Join(dir, "config.yml"), nil
|
||||
}
|
||||
|
||||
func ompAgentDir() (string, error) {
|
||||
if dir := strings.TrimSpace(os.Getenv("PI_CODING_AGENT_DIR")); dir != "" {
|
||||
return dir, nil
|
||||
}
|
||||
|
||||
home, err := os.UserHomeDir()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
configDir := strings.TrimSpace(os.Getenv("PI_CONFIG_DIR"))
|
||||
if configDir == "" {
|
||||
configDir = ".omp"
|
||||
}
|
||||
if filepath.IsAbs(configDir) {
|
||||
return filepath.Join(configDir, "agent"), nil
|
||||
}
|
||||
return filepath.Join(home, configDir, "agent"), nil
|
||||
}
|
||||
|
||||
func readOMPModelsConfig() (map[string]any, error) {
|
||||
path, err := ompModelsPath()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
data, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var cfg map[string]any
|
||||
if err := yaml.Unmarshal(data, &cfg); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if cfg == nil {
|
||||
cfg = make(map[string]any)
|
||||
}
|
||||
return cfg, nil
|
||||
}
|
||||
|
||||
func writeOMPModelsConfig(primary string, models []LaunchModel) error {
|
||||
path, err := ompModelsPath()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
cfg := make(map[string]any)
|
||||
if existing, err := readOMPModelsConfig(); err == nil {
|
||||
cfg = existing
|
||||
}
|
||||
|
||||
provider := ensureOMPProvider(cfg)
|
||||
existingByID := ompModelEntriesByID(provider)
|
||||
ordered := append([]LaunchModel(nil), models...)
|
||||
if model, ok := findLaunchModel(ordered, primary); ok {
|
||||
ordered = append([]LaunchModel{model}, removeLaunchModel(ordered, primary)...)
|
||||
} else {
|
||||
ordered = append([]LaunchModel{fallbackLaunchModel(primary)}, ordered...)
|
||||
}
|
||||
|
||||
var merged []any
|
||||
seen := make(map[string]bool, len(ordered))
|
||||
for _, model := range ordered {
|
||||
if model.Name == "" || seen[model.Name] {
|
||||
continue
|
||||
}
|
||||
seen[model.Name] = true
|
||||
entry := ompModelConfig(model)
|
||||
if existing, ok := existingByID[model.Name]; ok {
|
||||
for key, value := range existing {
|
||||
if _, overridden := entry[key]; !overridden {
|
||||
entry[key] = value
|
||||
}
|
||||
}
|
||||
}
|
||||
merged = append(merged, entry)
|
||||
}
|
||||
|
||||
for _, raw := range ompProviderModels(provider) {
|
||||
entry, ok := raw.(map[string]any)
|
||||
if !ok {
|
||||
merged = append(merged, raw)
|
||||
continue
|
||||
}
|
||||
id, _ := entry["id"].(string)
|
||||
if id == "" || seen[id] {
|
||||
continue
|
||||
}
|
||||
merged = append(merged, entry)
|
||||
}
|
||||
provider["models"] = merged
|
||||
|
||||
data, err := yaml.Marshal(cfg)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return fileutil.WriteWithBackup(path, data, ompIntegrationName)
|
||||
}
|
||||
|
||||
func writeOMPAgentConfig() error {
|
||||
path, err := ompConfigPath()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
cfg := make(map[string]any)
|
||||
if data, err := os.ReadFile(path); err == nil {
|
||||
if err := yaml.Unmarshal(data, &cfg); err != nil {
|
||||
return err
|
||||
}
|
||||
if cfg == nil {
|
||||
cfg = make(map[string]any)
|
||||
}
|
||||
}
|
||||
cfg["setupVersion"] = ompSetupVersion
|
||||
|
||||
data, err := yaml.Marshal(cfg)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return fileutil.WriteWithBackup(path, data, ompIntegrationName)
|
||||
}
|
||||
|
||||
func ensureOMPProvider(cfg map[string]any) map[string]any {
|
||||
providers, _ := cfg["providers"].(map[string]any)
|
||||
if providers == nil {
|
||||
providers = make(map[string]any)
|
||||
cfg["providers"] = providers
|
||||
}
|
||||
provider, _ := providers[ompProviderName].(map[string]any)
|
||||
if provider == nil {
|
||||
provider = make(map[string]any)
|
||||
providers[ompProviderName] = provider
|
||||
}
|
||||
|
||||
provider["baseUrl"] = ompBaseURL()
|
||||
provider["api"] = "openai-responses"
|
||||
provider["auth"] = "none"
|
||||
provider["discovery"] = map[string]any{"type": "ollama"}
|
||||
return provider
|
||||
}
|
||||
|
||||
func ompBaseURL() string {
|
||||
return strings.TrimRight(envconfig.ConnectableHost().String(), "/") + "/v1"
|
||||
}
|
||||
|
||||
func ompProviderHealthy(provider map[string]any) bool {
|
||||
baseURL, _ := provider["baseUrl"].(string)
|
||||
if strings.TrimRight(baseURL, "/") != strings.TrimRight(ompBaseURL(), "/") {
|
||||
return false
|
||||
}
|
||||
api, _ := provider["api"].(string)
|
||||
if api != "openai-responses" {
|
||||
return false
|
||||
}
|
||||
auth, _ := provider["auth"].(string)
|
||||
if auth != "none" {
|
||||
return false
|
||||
}
|
||||
discovery, _ := provider["discovery"].(map[string]any)
|
||||
if discovery == nil {
|
||||
return false
|
||||
}
|
||||
discoveryType, _ := discovery["type"].(string)
|
||||
return discoveryType == "ollama"
|
||||
}
|
||||
|
||||
func ompProvider(cfg map[string]any) (map[string]any, bool) {
|
||||
providers, ok := cfg["providers"].(map[string]any)
|
||||
if !ok {
|
||||
return nil, false
|
||||
}
|
||||
provider, ok := providers[ompProviderName].(map[string]any)
|
||||
return provider, ok
|
||||
}
|
||||
|
||||
func ompProviderModels(provider map[string]any) []any {
|
||||
models, _ := provider["models"].([]any)
|
||||
return models
|
||||
}
|
||||
|
||||
func ompModelEntriesByID(provider map[string]any) map[string]map[string]any {
|
||||
out := make(map[string]map[string]any)
|
||||
for _, raw := range ompProviderModels(provider) {
|
||||
entry, ok := raw.(map[string]any)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
if id, _ := entry["id"].(string); id != "" {
|
||||
out[id] = entry
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func ompModelConfig(modelInfo LaunchModel) map[string]any {
|
||||
entry := map[string]any{
|
||||
"id": modelInfo.Name,
|
||||
"name": modelInfo.Name,
|
||||
}
|
||||
input := []string{"text"}
|
||||
if slices.Contains(modelInfo.Capabilities, model.CapabilityVision) {
|
||||
input = append(input, "image")
|
||||
}
|
||||
entry["input"] = input
|
||||
|
||||
if modelInfo.ContextLength > 0 {
|
||||
entry["contextWindow"] = modelInfo.ContextLength
|
||||
}
|
||||
if modelInfo.MaxOutputTokens > 0 {
|
||||
entry["maxTokens"] = modelInfo.MaxOutputTokens
|
||||
}
|
||||
return entry
|
||||
}
|
||||
|
||||
func removeLaunchModel(models []LaunchModel, name string) []LaunchModel {
|
||||
out := make([]LaunchModel, 0, len(models))
|
||||
for _, model := range models {
|
||||
if launchModelMatches(model.Name, name) || launchModelMatches(name, model.Name) {
|
||||
continue
|
||||
}
|
||||
out = append(out, model)
|
||||
}
|
||||
return out
|
||||
}
|
||||
@@ -0,0 +1,687 @@
|
||||
package launch
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"slices"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
modelpkg "github.com/ollama/ollama/types/model"
|
||||
"gopkg.in/yaml.v3"
|
||||
)
|
||||
|
||||
func TestMain(m *testing.M) {
|
||||
if os.Getenv("OLLAMA_LAUNCH_OMP_TEST_HELPER") == "1" {
|
||||
runOMPTestHelper()
|
||||
return
|
||||
}
|
||||
os.Exit(m.Run())
|
||||
}
|
||||
|
||||
func runOMPTestHelper() {
|
||||
logPath := os.Getenv("OLLAMA_LAUNCH_OMP_TEST_LOG")
|
||||
if logPath != "" {
|
||||
f, err := os.OpenFile(logPath, os.O_CREATE|os.O_APPEND|os.O_WRONLY, 0o644)
|
||||
if err == nil {
|
||||
_, _ = fmt.Fprintln(f, strings.Join(os.Args[1:], " "))
|
||||
_ = f.Close()
|
||||
}
|
||||
}
|
||||
|
||||
if len(os.Args) >= 3 && os.Args[1] == "plugin" && os.Args[2] == "list" {
|
||||
fmt.Print(os.Getenv("OLLAMA_LAUNCH_OMP_TEST_PLUGIN_LIST"))
|
||||
os.Exit(0)
|
||||
}
|
||||
if len(os.Args) >= 4 && os.Args[1] == "plugin" && os.Args[2] == "install" {
|
||||
if os.Getenv("OLLAMA_LAUNCH_OMP_TEST_FAIL_INSTALL") == "1" {
|
||||
_, _ = fmt.Fprintln(os.Stderr, "install failed")
|
||||
os.Exit(1)
|
||||
}
|
||||
os.Exit(0)
|
||||
}
|
||||
os.Exit(0)
|
||||
}
|
||||
|
||||
func setOMPTestHome(t *testing.T, dir string) {
|
||||
t.Helper()
|
||||
setTestHome(t, dir)
|
||||
t.Setenv("PI_CONFIG_DIR", "")
|
||||
t.Setenv("PI_CODING_AGENT_DIR", "")
|
||||
}
|
||||
|
||||
func TestOMPIntegration(t *testing.T) {
|
||||
o := &OMP{}
|
||||
|
||||
t.Run("String", func(t *testing.T) {
|
||||
if got := o.String(); got != "OMP" {
|
||||
t.Errorf("String() = %q, want %q", got, "OMP")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("implements Runner", func(t *testing.T) {
|
||||
var _ Runner = o
|
||||
})
|
||||
|
||||
t.Run("implements ManagedSingleModel", func(t *testing.T) {
|
||||
var _ ManagedSingleModel = o
|
||||
})
|
||||
|
||||
t.Run("implements ManagedModelListConfigurer", func(t *testing.T) {
|
||||
var _ ManagedModelListConfigurer = o
|
||||
})
|
||||
|
||||
t.Run("does not require interactive onboarding", func(t *testing.T) {
|
||||
var _ ManagedInteractiveOnboarding = o
|
||||
if o.RequiresInteractiveOnboarding() {
|
||||
t.Fatal("OMP onboarding should not require an interactive terminal")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestOMPArgs(t *testing.T) {
|
||||
o := &OMP{}
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
model string
|
||||
args []string
|
||||
want []string
|
||||
}{
|
||||
{"with model", "gemma4", nil, []string{"--model", "ollama/gemma4"}},
|
||||
{"with cloud model", "kimi-k2.6:cloud", nil, []string{"--model", "ollama/kimi-k2.6:cloud"}},
|
||||
{"empty model", "", nil, nil},
|
||||
{"with model and extra", "gemma4", []string{"--help"}, []string{"--model", "ollama/gemma4", "--help"}},
|
||||
{"already qualified", "ollama/gemma4", nil, []string{"--model", "ollama/gemma4"}},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
got := o.args(tt.model, tt.args)
|
||||
if !slices.Equal(got, tt.want) {
|
||||
t.Errorf("args(%q, %v) = %v, want %v", tt.model, tt.args, got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestOMPRun_WebSearchPluginLifecycle(t *testing.T) {
|
||||
seedOMPHelperBinary := func(t *testing.T, dir string) {
|
||||
t.Helper()
|
||||
src, err := os.Executable()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
data, err := os.ReadFile(src)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
dst := filepath.Join(dir, ompExecutableNames()[0])
|
||||
if err := os.WriteFile(dst, data, 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
setCloudStatus := func(t *testing.T, disabled bool) {
|
||||
t.Helper()
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.URL.Path == "/api/status" {
|
||||
fmt.Fprintf(w, `{"cloud":{"disabled":%t,"source":"config"}}`, disabled)
|
||||
return
|
||||
}
|
||||
http.NotFound(w, r)
|
||||
}))
|
||||
t.Cleanup(srv.Close)
|
||||
t.Setenv("OLLAMA_HOST", srv.URL)
|
||||
}
|
||||
|
||||
setup := func(t *testing.T, pluginList string, cloudDisabled bool) (string, *OMP) {
|
||||
t.Helper()
|
||||
tmpDir := t.TempDir()
|
||||
setOMPTestHome(t, tmpDir)
|
||||
t.Setenv("PATH", tmpDir)
|
||||
t.Setenv("OLLAMA_LAUNCH_OMP_TEST_HELPER", "1")
|
||||
t.Setenv("OLLAMA_LAUNCH_OMP_TEST_PLUGIN_LIST", pluginList)
|
||||
logPath := filepath.Join(tmpDir, "omp.log")
|
||||
t.Setenv("OLLAMA_LAUNCH_OMP_TEST_LOG", logPath)
|
||||
setCloudStatus(t, cloudDisabled)
|
||||
seedOMPHelperBinary(t, tmpDir)
|
||||
return logPath, &OMP{}
|
||||
}
|
||||
|
||||
t.Run("web search missing installs before launch", func(t *testing.T) {
|
||||
logPath, o := setup(t, "No plugins installed\n", false)
|
||||
|
||||
if err := o.Run("kimi-k2.6:cloud", nil, []string{"session"}); err != nil {
|
||||
t.Fatalf("Run() error = %v", err)
|
||||
}
|
||||
|
||||
calls, err := os.ReadFile(logPath)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
got := string(calls)
|
||||
if !strings.Contains(got, "plugin list\n") {
|
||||
t.Fatalf("expected plugin list call, got:\n%s", got)
|
||||
}
|
||||
if !strings.Contains(got, "plugin install "+ompWebSearchPlugin+"\n") {
|
||||
t.Fatalf("expected plugin install call, got:\n%s", got)
|
||||
}
|
||||
if !strings.Contains(got, "--model ollama/kimi-k2.6:cloud session\n") {
|
||||
t.Fatalf("expected final omp launch call, got:\n%s", got)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("web search present refreshes before launch", func(t *testing.T) {
|
||||
logPath, o := setup(t, "npm Plugins:\n\n● "+ompWebSearchPlugin+"@0.0.5\n", false)
|
||||
|
||||
if err := o.Run("gemma4", nil, []string{"chat"}); err != nil {
|
||||
t.Fatalf("Run() error = %v", err)
|
||||
}
|
||||
|
||||
calls, err := os.ReadFile(logPath)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
got := string(calls)
|
||||
if !strings.Contains(got, "plugin install "+ompWebSearchPlugin+"\n") {
|
||||
t.Fatalf("expected plugin refresh install call, got:\n%s", got)
|
||||
}
|
||||
if !strings.Contains(got, "--model ollama/gemma4 chat\n") {
|
||||
t.Fatalf("expected final omp launch call, got:\n%s", got)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("web search install failure warns and continues", func(t *testing.T) {
|
||||
logPath, o := setup(t, "No plugins installed\n", false)
|
||||
t.Setenv("OLLAMA_LAUNCH_OMP_TEST_FAIL_INSTALL", "1")
|
||||
|
||||
stderr := captureStderr(t, func() {
|
||||
if err := o.Run("gemma4", nil, []string{"chat"}); err != nil {
|
||||
t.Fatalf("Run() should continue after plugin install failure, got %v", err)
|
||||
}
|
||||
})
|
||||
if !strings.Contains(stderr, "Warning: could not install "+ompWebSearchPlugin) {
|
||||
t.Fatalf("expected install warning, got:\n%s", stderr)
|
||||
}
|
||||
|
||||
calls, err := os.ReadFile(logPath)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !strings.Contains(string(calls), "--model ollama/gemma4 chat\n") {
|
||||
t.Fatalf("expected final omp launch call, got:\n%s", calls)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("cloud disabled skips web search plugin management", func(t *testing.T) {
|
||||
logPath, o := setup(t, "No plugins installed\n", true)
|
||||
|
||||
stderr := captureStderr(t, func() {
|
||||
if err := o.Run("gemma4", nil, []string{"chat"}); err != nil {
|
||||
t.Fatalf("Run() error = %v", err)
|
||||
}
|
||||
})
|
||||
if !strings.Contains(stderr, "Cloud is disabled; skipping "+ompWebSearchPlugin+" setup.") {
|
||||
t.Fatalf("expected cloud-disabled skip message, got:\n%s", stderr)
|
||||
}
|
||||
|
||||
calls, err := os.ReadFile(logPath)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
got := string(calls)
|
||||
if strings.Contains(got, "plugin list\n") || strings.Contains(got, "plugin install "+ompWebSearchPlugin+"\n") {
|
||||
t.Fatalf("did not expect plugin management calls, got:\n%s", got)
|
||||
}
|
||||
if !strings.Contains(got, "--model ollama/gemma4 chat\n") {
|
||||
t.Fatalf("expected final omp launch call, got:\n%s", got)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestOMPFindPath(t *testing.T) {
|
||||
o := &OMP{}
|
||||
|
||||
t.Run("finds omp in PATH", func(t *testing.T) {
|
||||
tmpDir := t.TempDir()
|
||||
name := "omp"
|
||||
if runtime.GOOS == "windows" {
|
||||
name = "omp.exe"
|
||||
}
|
||||
fakeBin := filepath.Join(tmpDir, name)
|
||||
os.WriteFile(fakeBin, []byte("#!/bin/sh\n"), 0o755)
|
||||
t.Setenv("PATH", tmpDir)
|
||||
|
||||
got, err := o.findPath()
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if got != fakeBin {
|
||||
t.Errorf("findPath() = %q, want %q", got, fakeBin)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("falls back to ~/.local/bin/omp", func(t *testing.T) {
|
||||
home := t.TempDir()
|
||||
setOMPTestHome(t, home)
|
||||
t.Setenv("PATH", t.TempDir())
|
||||
|
||||
fallback := filepath.Join(home, ".local", "bin", ompExecutableNames()[0])
|
||||
os.MkdirAll(filepath.Dir(fallback), 0o755)
|
||||
os.WriteFile(fallback, []byte("#!/bin/sh\n"), 0o755)
|
||||
|
||||
got, err := o.findPath()
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if got != fallback {
|
||||
t.Errorf("findPath() = %q, want %q", got, fallback)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("falls back to ~/.bun/bin/omp", func(t *testing.T) {
|
||||
home := t.TempDir()
|
||||
setOMPTestHome(t, home)
|
||||
t.Setenv("PATH", t.TempDir())
|
||||
|
||||
fallback := filepath.Join(home, ".bun", "bin", ompExecutableNames()[0])
|
||||
os.MkdirAll(filepath.Dir(fallback), 0o755)
|
||||
os.WriteFile(fallback, []byte("#!/bin/sh\n"), 0o755)
|
||||
|
||||
got, err := o.findPath()
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if got != fallback {
|
||||
t.Errorf("findPath() = %q, want %q", got, fallback)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("returns error when not found", func(t *testing.T) {
|
||||
home := t.TempDir()
|
||||
setOMPTestHome(t, home)
|
||||
t.Setenv("PATH", t.TempDir())
|
||||
|
||||
if _, err := o.findPath(); err == nil {
|
||||
t.Fatal("expected error, got nil")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestOMPConfigureWithModelsWritesModelsYML(t *testing.T) {
|
||||
home := t.TempDir()
|
||||
setOMPTestHome(t, home)
|
||||
t.Setenv("OLLAMA_HOST", "http://0.0.0.0:11434")
|
||||
|
||||
o := &OMP{}
|
||||
models := []LaunchModel{
|
||||
{
|
||||
Name: "glm-5.1:cloud",
|
||||
ContextLength: 202_752,
|
||||
MaxOutputTokens: 131_072,
|
||||
},
|
||||
{
|
||||
Name: "qwen3.6",
|
||||
Capabilities: []modelpkg.Capability{modelpkg.CapabilityVision},
|
||||
},
|
||||
}
|
||||
if err := o.ConfigureWithModels("glm-5.1:cloud", models); err != nil {
|
||||
t.Fatalf("ConfigureWithModels returned error: %v", err)
|
||||
}
|
||||
|
||||
path := filepath.Join(home, ".omp", "agent", "models.yml")
|
||||
data, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to read models.yml: %v", err)
|
||||
}
|
||||
|
||||
cfg := parseOMPConfigYAML(t, data)
|
||||
provider := ompProviderFromYAML(t, cfg)
|
||||
if provider["baseUrl"] != "http://127.0.0.1:11434/v1" {
|
||||
t.Fatalf("baseUrl = %v, want connectable OpenAI-compatible host", provider["baseUrl"])
|
||||
}
|
||||
if provider["api"] != "openai-responses" {
|
||||
t.Fatalf("api = %v, want openai-responses", provider["api"])
|
||||
}
|
||||
if provider["auth"] != "none" {
|
||||
t.Fatalf("auth = %v, want none", provider["auth"])
|
||||
}
|
||||
discovery, _ := provider["discovery"].(map[string]any)
|
||||
if discovery["type"] != "ollama" {
|
||||
t.Fatalf("discovery = %v, want type ollama", discovery)
|
||||
}
|
||||
|
||||
entries := ompModelEntriesFromYAML(t, provider)
|
||||
if len(entries) != 2 {
|
||||
t.Fatalf("models length = %d, want 2", len(entries))
|
||||
}
|
||||
if entries[0]["id"] != "glm-5.1:cloud" {
|
||||
t.Fatalf("first model id = %v, want primary first", entries[0]["id"])
|
||||
}
|
||||
if got := numericYAMLValue(entries[0]["contextWindow"]); got != 202_752 {
|
||||
t.Fatalf("contextWindow = %d, want 202752", got)
|
||||
}
|
||||
if got := numericYAMLValue(entries[0]["maxTokens"]); got != 131_072 {
|
||||
t.Fatalf("maxTokens = %d, want 131072", got)
|
||||
}
|
||||
if input := stringSliceYAMLValue(entries[1]["input"]); !slices.Equal(input, []string{"text", "image"}) {
|
||||
t.Fatalf("vision input = %v, want [text image]", input)
|
||||
}
|
||||
if got := o.CurrentModel(); got != "glm-5.1:cloud" {
|
||||
t.Fatalf("CurrentModel = %q, want glm-5.1:cloud", got)
|
||||
}
|
||||
|
||||
configPath := filepath.Join(home, ".omp", "agent", "config.yml")
|
||||
configData, err := os.ReadFile(configPath)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to read config.yml: %v", err)
|
||||
}
|
||||
config := parseOMPConfigYAML(t, configData)
|
||||
if got := numericYAMLValue(config["setupVersion"]); got != ompSetupVersion {
|
||||
t.Fatalf("setupVersion = %d, want %d", got, ompSetupVersion)
|
||||
}
|
||||
if paths := o.Paths(); !slices.Equal(paths, []string{path, configPath}) {
|
||||
t.Fatalf("Paths = %v, want [%s %s]", paths, path, configPath)
|
||||
}
|
||||
}
|
||||
|
||||
func TestOMPConfigureWithModelsPreservesExistingConfig(t *testing.T) {
|
||||
home := t.TempDir()
|
||||
setOMPTestHome(t, home)
|
||||
|
||||
modelsPath := filepath.Join(home, ".omp", "agent", "models.yml")
|
||||
if err := os.MkdirAll(filepath.Dir(modelsPath), 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
existing := []byte(`
|
||||
providers:
|
||||
anthropic:
|
||||
baseUrl: https://example.com/anthropic
|
||||
ollama:
|
||||
baseUrl: http://old-host:11434
|
||||
api: openai-responses
|
||||
auth: none
|
||||
models:
|
||||
- id: old-model
|
||||
name: Old Model
|
||||
customField: keep-me
|
||||
`)
|
||||
if err := os.WriteFile(modelsPath, existing, 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
configPath := filepath.Join(home, ".omp", "agent", "config.yml")
|
||||
existingConfig := []byte(`
|
||||
lastChangelogVersion: 15.7.6
|
||||
setupVersion: 0
|
||||
theme: monochrome
|
||||
`)
|
||||
if err := os.WriteFile(configPath, existingConfig, 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
o := &OMP{}
|
||||
if err := o.ConfigureWithModels("new-model", []LaunchModel{{Name: "new-model"}, {Name: "old-model"}}); err != nil {
|
||||
t.Fatalf("ConfigureWithModels returned error: %v", err)
|
||||
}
|
||||
|
||||
data, err := os.ReadFile(modelsPath)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
cfg := parseOMPConfigYAML(t, data)
|
||||
providers, _ := cfg["providers"].(map[string]any)
|
||||
if _, ok := providers["anthropic"]; !ok {
|
||||
t.Fatalf("expected non-Ollama provider to be preserved: %v", providers)
|
||||
}
|
||||
|
||||
provider := ompProviderFromYAML(t, cfg)
|
||||
if provider["baseUrl"] != "http://127.0.0.1:11434/v1" {
|
||||
t.Fatalf("baseUrl = %v, want repaired OpenAI-compatible host", provider["baseUrl"])
|
||||
}
|
||||
|
||||
entries := ompModelEntriesFromYAML(t, provider)
|
||||
if len(entries) != 2 {
|
||||
t.Fatalf("models length = %d, want 2", len(entries))
|
||||
}
|
||||
if entries[0]["id"] != "new-model" {
|
||||
t.Fatalf("first model id = %v, want new-model", entries[0]["id"])
|
||||
}
|
||||
if entries[1]["id"] != "old-model" {
|
||||
t.Fatalf("second model id = %v, want old-model", entries[1]["id"])
|
||||
}
|
||||
if entries[1]["customField"] != "keep-me" {
|
||||
t.Fatalf("custom field was not preserved: %v", entries[1])
|
||||
}
|
||||
|
||||
configData, err := os.ReadFile(configPath)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
config := parseOMPConfigYAML(t, configData)
|
||||
if got := numericYAMLValue(config["setupVersion"]); got != ompSetupVersion {
|
||||
t.Fatalf("setupVersion = %d, want %d", got, ompSetupVersion)
|
||||
}
|
||||
if config["theme"] != "monochrome" {
|
||||
t.Fatalf("theme was not preserved: %v", config)
|
||||
}
|
||||
if config["lastChangelogVersion"] != "15.7.6" {
|
||||
t.Fatalf("lastChangelogVersion was not preserved: %v", config)
|
||||
}
|
||||
}
|
||||
|
||||
func TestOMPConfigureWithModelsAlwaysMarksSetupComplete(t *testing.T) {
|
||||
home := t.TempDir()
|
||||
setOMPTestHome(t, home)
|
||||
|
||||
configPath := filepath.Join(home, ".omp", "agent", "config.yml")
|
||||
if err := os.MkdirAll(filepath.Dir(configPath), 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := os.WriteFile(configPath, []byte("setupVersion: 2\n"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
o := &OMP{}
|
||||
if err := o.ConfigureWithModels("new-model", []LaunchModel{{Name: "new-model"}}); err != nil {
|
||||
t.Fatalf("ConfigureWithModels returned error: %v", err)
|
||||
}
|
||||
|
||||
configData, err := os.ReadFile(configPath)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
config := parseOMPConfigYAML(t, configData)
|
||||
if got := numericYAMLValue(config["setupVersion"]); got != ompSetupVersion {
|
||||
t.Fatalf("setupVersion = %d, want %d", got, ompSetupVersion)
|
||||
}
|
||||
}
|
||||
|
||||
func TestOMPConfigureWithModelsRespectsPiConfigDir(t *testing.T) {
|
||||
home := t.TempDir()
|
||||
setOMPTestHome(t, home)
|
||||
t.Setenv("PI_CONFIG_DIR", ".custom-omp")
|
||||
|
||||
o := &OMP{}
|
||||
if err := o.ConfigureWithModels("new-model", []LaunchModel{{Name: "new-model"}}); err != nil {
|
||||
t.Fatalf("ConfigureWithModels returned error: %v", err)
|
||||
}
|
||||
|
||||
modelsPath := filepath.Join(home, ".custom-omp", "agent", "models.yml")
|
||||
configPath := filepath.Join(home, ".custom-omp", "agent", "config.yml")
|
||||
for _, path := range []string{modelsPath, configPath} {
|
||||
if _, err := os.Stat(path); err != nil {
|
||||
t.Fatalf("expected %s to be written: %v", path, err)
|
||||
}
|
||||
}
|
||||
if _, err := os.Stat(filepath.Join(home, ".omp", "agent", "models.yml")); !os.IsNotExist(err) {
|
||||
t.Fatalf("expected default OMP models path to be untouched, got err %v", err)
|
||||
}
|
||||
if paths := o.Paths(); !slices.Equal(paths, []string{modelsPath, configPath}) {
|
||||
t.Fatalf("Paths = %v, want [%s %s]", paths, modelsPath, configPath)
|
||||
}
|
||||
}
|
||||
|
||||
func TestOMPConfigureWithModelsRespectsPiCodingAgentDir(t *testing.T) {
|
||||
home := t.TempDir()
|
||||
setOMPTestHome(t, home)
|
||||
agentDir := filepath.Join(home, "agent-override")
|
||||
t.Setenv("PI_CONFIG_DIR", ".ignored-omp")
|
||||
t.Setenv("PI_CODING_AGENT_DIR", agentDir)
|
||||
|
||||
o := &OMP{}
|
||||
if err := o.ConfigureWithModels("new-model", []LaunchModel{{Name: "new-model"}}); err != nil {
|
||||
t.Fatalf("ConfigureWithModels returned error: %v", err)
|
||||
}
|
||||
|
||||
modelsPath := filepath.Join(agentDir, "models.yml")
|
||||
configPath := filepath.Join(agentDir, "config.yml")
|
||||
for _, path := range []string{modelsPath, configPath} {
|
||||
if _, err := os.Stat(path); err != nil {
|
||||
t.Fatalf("expected %s to be written: %v", path, err)
|
||||
}
|
||||
}
|
||||
if _, err := os.Stat(filepath.Join(home, ".ignored-omp", "agent", "models.yml")); !os.IsNotExist(err) {
|
||||
t.Fatalf("expected PI_CONFIG_DIR path to be ignored when PI_CODING_AGENT_DIR is set, got err %v", err)
|
||||
}
|
||||
if got := o.CurrentModel(); got != "new-model" {
|
||||
t.Fatalf("CurrentModel = %q, want new-model", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestOMPCurrentModelRequiresHealthyProvider(t *testing.T) {
|
||||
home := t.TempDir()
|
||||
setOMPTestHome(t, home)
|
||||
t.Setenv("OLLAMA_HOST", "http://127.0.0.1:11434")
|
||||
|
||||
modelsPath := filepath.Join(home, ".omp", "agent", "models.yml")
|
||||
if err := os.MkdirAll(filepath.Dir(modelsPath), 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
provider string
|
||||
}{
|
||||
{
|
||||
name: "wrong base url",
|
||||
provider: "" +
|
||||
" baseUrl: http://127.0.0.1:9999/v1\n" +
|
||||
" api: openai-responses\n" +
|
||||
" auth: none\n" +
|
||||
" discovery:\n" +
|
||||
" type: ollama\n",
|
||||
},
|
||||
{
|
||||
name: "wrong api",
|
||||
provider: "" +
|
||||
" baseUrl: http://127.0.0.1:11434/v1\n" +
|
||||
" api: openai-chat\n" +
|
||||
" auth: none\n" +
|
||||
" discovery:\n" +
|
||||
" type: ollama\n",
|
||||
},
|
||||
{
|
||||
name: "wrong auth",
|
||||
provider: "" +
|
||||
" baseUrl: http://127.0.0.1:11434/v1\n" +
|
||||
" api: openai-responses\n" +
|
||||
" auth: api-key\n" +
|
||||
" discovery:\n" +
|
||||
" type: ollama\n",
|
||||
},
|
||||
{
|
||||
name: "wrong discovery",
|
||||
provider: "" +
|
||||
" baseUrl: http://127.0.0.1:11434/v1\n" +
|
||||
" api: openai-responses\n" +
|
||||
" auth: none\n" +
|
||||
" discovery:\n" +
|
||||
" type: static\n",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
cfg := "providers:\n" +
|
||||
" ollama:\n" +
|
||||
tt.provider +
|
||||
" models:\n" +
|
||||
" - id: gemma4\n"
|
||||
if err := os.WriteFile(modelsPath, []byte(cfg), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if got := (&OMP{}).CurrentModel(); got != "" {
|
||||
t.Fatalf("expected stale config to return empty current model, got %q", got)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func parseOMPConfigYAML(t *testing.T, data []byte) map[string]any {
|
||||
t.Helper()
|
||||
var cfg map[string]any
|
||||
if err := yaml.Unmarshal(data, &cfg); err != nil {
|
||||
t.Fatalf("generated YAML did not parse: %v\n%s", err, data)
|
||||
}
|
||||
return cfg
|
||||
}
|
||||
|
||||
func ompProviderFromYAML(t *testing.T, cfg map[string]any) map[string]any {
|
||||
t.Helper()
|
||||
providers, ok := cfg["providers"].(map[string]any)
|
||||
if !ok {
|
||||
t.Fatalf("providers missing from config: %v", cfg)
|
||||
}
|
||||
provider, ok := providers["ollama"].(map[string]any)
|
||||
if !ok {
|
||||
t.Fatalf("ollama provider missing from config: %v", providers)
|
||||
}
|
||||
return provider
|
||||
}
|
||||
|
||||
func ompModelEntriesFromYAML(t *testing.T, provider map[string]any) []map[string]any {
|
||||
t.Helper()
|
||||
rawModels, ok := provider["models"].([]any)
|
||||
if !ok {
|
||||
t.Fatalf("provider models missing: %v", provider)
|
||||
}
|
||||
models := make([]map[string]any, 0, len(rawModels))
|
||||
for _, raw := range rawModels {
|
||||
entry, ok := raw.(map[string]any)
|
||||
if !ok {
|
||||
t.Fatalf("model entry has unexpected type %T: %v", raw, raw)
|
||||
}
|
||||
models = append(models, entry)
|
||||
}
|
||||
return models
|
||||
}
|
||||
|
||||
func numericYAMLValue(value any) int {
|
||||
switch v := value.(type) {
|
||||
case int:
|
||||
return v
|
||||
case int64:
|
||||
return int(v)
|
||||
case float64:
|
||||
return int(v)
|
||||
default:
|
||||
return 0
|
||||
}
|
||||
}
|
||||
|
||||
func stringSliceYAMLValue(value any) []string {
|
||||
raw, _ := value.([]any)
|
||||
out := make([]string, 0, len(raw))
|
||||
for _, item := range raw {
|
||||
if s, ok := item.(string); ok {
|
||||
out = append(out, s)
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
+117
-4
@@ -8,11 +8,16 @@ import (
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"slices"
|
||||
"strings"
|
||||
|
||||
"github.com/ollama/ollama/cmd/internal/fileutil"
|
||||
"github.com/ollama/ollama/envconfig"
|
||||
)
|
||||
|
||||
const openCodeInstallScript = "curl -fsSL https://opencode.ai/install | bash"
|
||||
|
||||
var openCodeGOOS = runtime.GOOS
|
||||
|
||||
// OpenCode implements Runner and Editor for OpenCode integration.
|
||||
// Config is passed via OPENCODE_CONFIG_CONTENT env var at launch time
|
||||
// instead of writing to opencode's config files.
|
||||
@@ -33,7 +38,7 @@ func findOpenCode() (string, bool) {
|
||||
return "", false
|
||||
}
|
||||
name := "opencode"
|
||||
if runtime.GOOS == "windows" {
|
||||
if openCodeGOOS == "windows" {
|
||||
name = "opencode.exe"
|
||||
}
|
||||
fallback := filepath.Join(home, ".opencode", "bin", name)
|
||||
@@ -44,9 +49,9 @@ func findOpenCode() (string, bool) {
|
||||
}
|
||||
|
||||
func (o *OpenCode) Run(model string, models []LaunchModel, args []string) error {
|
||||
opencodePath, ok := findOpenCode()
|
||||
if !ok {
|
||||
return fmt.Errorf("opencode is not installed, install from https://opencode.ai")
|
||||
opencodePath, err := ensureOpenCodeInstalled()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
cmd := exec.Command(opencodePath, args...)
|
||||
@@ -60,6 +65,78 @@ func (o *OpenCode) Run(model string, models []LaunchModel, args []string) error
|
||||
return cmd.Run()
|
||||
}
|
||||
|
||||
func ensureOpenCodeInstalled() (string, error) {
|
||||
if opencodePath, ok := findOpenCode(); ok {
|
||||
return opencodePath, nil
|
||||
}
|
||||
|
||||
if err := checkOpenCodeInstallerDependencies(); err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
ok, err := ConfirmPrompt("OpenCode is not installed. Install now?")
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if !ok {
|
||||
return "", fmt.Errorf("opencode installation cancelled")
|
||||
}
|
||||
|
||||
bin, args, err := openCodeInstallerCommand(openCodeGOOS)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
fmt.Fprintf(os.Stderr, "\nInstalling OpenCode...\n")
|
||||
cmd := exec.Command(bin, args...)
|
||||
cmd.Stdin = os.Stdin
|
||||
cmd.Stdout = os.Stdout
|
||||
cmd.Stderr = os.Stderr
|
||||
if err := cmd.Run(); err != nil {
|
||||
return "", fmt.Errorf("failed to install opencode: %w", err)
|
||||
}
|
||||
|
||||
opencodePath, ok := findOpenCode()
|
||||
if !ok {
|
||||
return "", fmt.Errorf("opencode was installed but the binary was not found on PATH\n\nYou may need to restart your shell")
|
||||
}
|
||||
|
||||
fmt.Fprintf(os.Stderr, "%sOpenCode installed successfully%s\n\n", ansiGreen, ansiReset)
|
||||
return opencodePath, nil
|
||||
}
|
||||
|
||||
func checkOpenCodeInstallerDependencies() error {
|
||||
switch openCodeGOOS {
|
||||
case "windows":
|
||||
if _, err := exec.LookPath("npm"); err != nil {
|
||||
return fmt.Errorf("opencode is not installed and required dependencies are missing\n\nInstall the following first:\n npm (Node.js): https://nodejs.org/\n\nThen re-run:\n ollama launch opencode")
|
||||
}
|
||||
default:
|
||||
var missing []string
|
||||
if _, err := exec.LookPath("curl"); err != nil {
|
||||
missing = append(missing, "curl: https://curl.se/")
|
||||
}
|
||||
if _, err := exec.LookPath("bash"); err != nil {
|
||||
missing = append(missing, "bash: https://www.gnu.org/software/bash/")
|
||||
}
|
||||
if len(missing) > 0 {
|
||||
return fmt.Errorf("opencode is not installed and required dependencies are missing\n\nInstall the following first:\n %s\n\nThen re-run:\n ollama launch opencode", strings.Join(missing, "\n "))
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func openCodeInstallerCommand(goos string) (string, []string, error) {
|
||||
switch goos {
|
||||
case "windows":
|
||||
return "npm", []string{"install", "-g", "opencode-ai@latest"}, nil
|
||||
case "darwin", "linux":
|
||||
return "bash", []string{"-c", "set -o pipefail; " + openCodeInstallScript}, nil
|
||||
default:
|
||||
return "", nil, fmt.Errorf("unsupported platform for opencode install: %s", goos)
|
||||
}
|
||||
}
|
||||
|
||||
// resolveContent returns the inline config to send via OPENCODE_CONFIG_CONTENT.
|
||||
// Returns content built by Edit if available, otherwise builds from model.json
|
||||
// with the requested model as primary (e.g. re-launch with saved config).
|
||||
@@ -278,6 +355,25 @@ func buildModelEntries(modelList []LaunchModel) map[string]any {
|
||||
"output": []string{"text"},
|
||||
}
|
||||
}
|
||||
if model.HasCapability("thinking") {
|
||||
entry["reasoning"] = true
|
||||
if openCodeModelSupportsThinkingLevels(model) {
|
||||
entry["options"] = map[string]any{"reasoningEffort": "medium"}
|
||||
entry["variants"] = map[string]any{
|
||||
"low": map[string]any{"reasoningEffort": "low"},
|
||||
"medium": map[string]any{"reasoningEffort": "medium"},
|
||||
"high": map[string]any{"reasoningEffort": "high"},
|
||||
"max": map[string]any{"reasoningEffort": "max"},
|
||||
}
|
||||
} else {
|
||||
entry["variants"] = map[string]any{
|
||||
"none": map[string]any{"reasoningEffort": "none"},
|
||||
"low": map[string]any{"disabled": true},
|
||||
"medium": map[string]any{"disabled": true},
|
||||
"high": map[string]any{"disabled": true},
|
||||
}
|
||||
}
|
||||
}
|
||||
if model.MaxOutputTokens > 0 {
|
||||
limit := make(map[string]any)
|
||||
if model.ContextLength > 0 {
|
||||
@@ -290,3 +386,20 @@ func buildModelEntries(modelList []LaunchModel) map[string]any {
|
||||
}
|
||||
return models
|
||||
}
|
||||
|
||||
func openCodeModelSupportsThinkingLevels(model LaunchModel) bool {
|
||||
for _, family := range append([]string{model.Details.Family}, model.Details.Families...) {
|
||||
if normalizeOpenCodeModelFamily(family) == "gptoss" {
|
||||
return true
|
||||
}
|
||||
}
|
||||
|
||||
return strings.Contains(normalizeOpenCodeModelFamily(model.Name), "gptoss")
|
||||
}
|
||||
|
||||
func normalizeOpenCodeModelFamily(s string) string {
|
||||
s = strings.ToLower(s)
|
||||
s = strings.ReplaceAll(s, "-", "")
|
||||
s = strings.ReplaceAll(s, "_", "")
|
||||
return s
|
||||
}
|
||||
@@ -6,8 +6,10 @@ import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/ollama/ollama/api"
|
||||
"github.com/ollama/ollama/types/model"
|
||||
)
|
||||
|
||||
@@ -174,6 +176,54 @@ func TestOpenCodeEdit(t *testing.T) {
|
||||
t.Fatalf("modalities.output = %v, want [text]", output)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("thinking model gets on off reasoning variants", func(t *testing.T) {
|
||||
models := buildModelEntries([]LaunchModel{{Name: "thinking-model", Capabilities: []model.Capability{model.CapabilityThinking}}})
|
||||
entry, _ := models["thinking-model"].(map[string]any)
|
||||
|
||||
if entry["reasoning"] != true {
|
||||
t.Fatalf("reasoning = %v, want true", entry["reasoning"])
|
||||
}
|
||||
variants, _ := entry["variants"].(map[string]any)
|
||||
none, _ := variants["none"].(map[string]any)
|
||||
if none["reasoningEffort"] != "none" {
|
||||
t.Fatalf("variants.none.reasoningEffort = %v, want none", none["reasoningEffort"])
|
||||
}
|
||||
for _, level := range []string{"low", "medium", "high"} {
|
||||
variant, _ := variants[level].(map[string]any)
|
||||
if variant["disabled"] != true {
|
||||
t.Fatalf("variants.%s.disabled = %v, want true", level, variant["disabled"])
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("gpt oss gets reasoning level variants", func(t *testing.T) {
|
||||
models := buildModelEntries([]LaunchModel{{Name: "gpt-oss:120b-cloud", Capabilities: []model.Capability{model.CapabilityThinking}}})
|
||||
entry, _ := models["gpt-oss:120b-cloud"].(map[string]any)
|
||||
options, _ := entry["options"].(map[string]any)
|
||||
|
||||
if options["reasoningEffort"] != "medium" {
|
||||
t.Fatalf("options.reasoningEffort = %v, want medium", options["reasoningEffort"])
|
||||
}
|
||||
variants, _ := entry["variants"].(map[string]any)
|
||||
for _, level := range []string{"low", "medium", "high", "max"} {
|
||||
variant, _ := variants[level].(map[string]any)
|
||||
if variant["reasoningEffort"] != level {
|
||||
t.Fatalf("variants.%s.reasoningEffort = %v, want %s", level, variant["reasoningEffort"], level)
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("gpt oss family gets reasoning level variants", func(t *testing.T) {
|
||||
models := buildModelEntries([]LaunchModel{{Name: "reasoning-model", Capabilities: []model.Capability{model.CapabilityThinking}, Details: api.ModelDetails{Families: []string{"gptoss"}}}})
|
||||
entry, _ := models["reasoning-model"].(map[string]any)
|
||||
variants, _ := entry["variants"].(map[string]any)
|
||||
max, _ := variants["max"].(map[string]any)
|
||||
|
||||
if max["reasoningEffort"] != "max" {
|
||||
t.Fatalf("variants.max.reasoningEffort = %v, want max", max["reasoningEffort"])
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestBuildModelEntries(t *testing.T) {
|
||||
@@ -289,12 +339,16 @@ func TestLookupCloudModelLimit(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestFindOpenCode(t *testing.T) {
|
||||
oldGOOS := openCodeGOOS
|
||||
t.Cleanup(func() { openCodeGOOS = oldGOOS })
|
||||
|
||||
t.Run("fallback to ~/.opencode/bin", func(t *testing.T) {
|
||||
tmpDir := t.TempDir()
|
||||
setTestHome(t, tmpDir)
|
||||
|
||||
// Ensure opencode is not on PATH
|
||||
t.Setenv("PATH", tmpDir)
|
||||
openCodeGOOS = runtime.GOOS
|
||||
|
||||
// Without the fallback binary, findOpenCode should fail
|
||||
if _, ok := findOpenCode(); ok {
|
||||
@@ -322,6 +376,259 @@ func TestFindOpenCode(t *testing.T) {
|
||||
})
|
||||
}
|
||||
|
||||
func TestEnsureOpenCodeInstalled(t *testing.T) {
|
||||
oldGOOS := openCodeGOOS
|
||||
t.Cleanup(func() { openCodeGOOS = oldGOOS })
|
||||
|
||||
withConfirm := func(t *testing.T, fn func(prompt string) (bool, error)) {
|
||||
t.Helper()
|
||||
oldConfirm := DefaultConfirmPrompt
|
||||
DefaultConfirmPrompt = func(prompt string, options ConfirmOptions) (bool, error) {
|
||||
return fn(prompt)
|
||||
}
|
||||
t.Cleanup(func() { DefaultConfirmPrompt = oldConfirm })
|
||||
}
|
||||
|
||||
t.Run("already installed", func(t *testing.T) {
|
||||
setTestHome(t, t.TempDir())
|
||||
tmpDir := t.TempDir()
|
||||
t.Setenv("PATH", tmpDir)
|
||||
openCodeGOOS = runtime.GOOS
|
||||
writeFakeBinary(t, tmpDir, "opencode")
|
||||
|
||||
withConfirm(t, func(prompt string) (bool, error) {
|
||||
t.Fatalf("did not expect prompt, got %q", prompt)
|
||||
return false, nil
|
||||
})
|
||||
|
||||
bin, err := ensureOpenCodeInstalled()
|
||||
if err != nil {
|
||||
t.Fatalf("ensureOpenCodeInstalled() error = %v", err)
|
||||
}
|
||||
if filepath.Base(bin) == "" {
|
||||
t.Fatalf("expected opencode binary path, got %q", bin)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("missing dependencies", func(t *testing.T) {
|
||||
setTestHome(t, t.TempDir())
|
||||
t.Setenv("PATH", t.TempDir())
|
||||
openCodeGOOS = "linux"
|
||||
|
||||
withConfirm(t, func(prompt string) (bool, error) {
|
||||
t.Fatalf("did not expect prompt, got %q", prompt)
|
||||
return false, nil
|
||||
})
|
||||
|
||||
_, err := ensureOpenCodeInstalled()
|
||||
if err == nil || !strings.Contains(err.Error(), "required dependencies are missing") {
|
||||
t.Fatalf("expected missing dependency error, got %v", err)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("missing and user declines install", func(t *testing.T) {
|
||||
setTestHome(t, t.TempDir())
|
||||
tmpDir := t.TempDir()
|
||||
t.Setenv("PATH", tmpDir)
|
||||
openCodeGOOS = "linux"
|
||||
writeFakeBinary(t, tmpDir, "curl")
|
||||
writeFakeBinary(t, tmpDir, "bash")
|
||||
|
||||
withConfirm(t, func(prompt string) (bool, error) {
|
||||
if !strings.Contains(prompt, "OpenCode is not installed.") {
|
||||
t.Fatalf("unexpected prompt: %q", prompt)
|
||||
}
|
||||
return false, nil
|
||||
})
|
||||
|
||||
_, err := ensureOpenCodeInstalled()
|
||||
if err == nil || !strings.Contains(err.Error(), "installation cancelled") {
|
||||
t.Fatalf("expected cancellation error, got %v", err)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("missing and user confirms unix install succeeds", func(t *testing.T) {
|
||||
if runtime.GOOS == "windows" {
|
||||
t.Skip("uses POSIX shell fake binaries")
|
||||
}
|
||||
|
||||
homeDir := t.TempDir()
|
||||
setTestHome(t, homeDir)
|
||||
tmpDir := t.TempDir()
|
||||
t.Setenv("PATH", tmpDir)
|
||||
openCodeGOOS = "linux"
|
||||
writeFakeBinary(t, tmpDir, "curl")
|
||||
|
||||
installLog := filepath.Join(tmpDir, "bash.log")
|
||||
opencodePath := filepath.Join(homeDir, ".opencode", "bin", "opencode")
|
||||
bashScript := fmt.Sprintf(`#!/bin/sh
|
||||
echo "$@" >> %q
|
||||
if [ "$1" = "-c" ]; then
|
||||
/bin/mkdir -p %q
|
||||
/bin/cat > %q <<'EOS'
|
||||
#!/bin/sh
|
||||
exit 0
|
||||
EOS
|
||||
/bin/chmod +x %q
|
||||
fi
|
||||
exit 0
|
||||
`, installLog, filepath.Dir(opencodePath), opencodePath, opencodePath)
|
||||
if err := os.WriteFile(filepath.Join(tmpDir, "bash"), []byte(bashScript), 0o755); err != nil {
|
||||
t.Fatalf("failed to write fake bash: %v", err)
|
||||
}
|
||||
|
||||
withConfirm(t, func(prompt string) (bool, error) {
|
||||
return true, nil
|
||||
})
|
||||
|
||||
bin, err := ensureOpenCodeInstalled()
|
||||
if err != nil {
|
||||
t.Fatalf("ensureOpenCodeInstalled() error = %v", err)
|
||||
}
|
||||
if bin != opencodePath {
|
||||
t.Fatalf("bin = %q, want %q", bin, opencodePath)
|
||||
}
|
||||
|
||||
logData, err := os.ReadFile(installLog)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to read install log: %v", err)
|
||||
}
|
||||
if !strings.Contains(string(logData), openCodeInstallScript) {
|
||||
t.Fatalf("expected opencode install script in log, got:\n%s", string(logData))
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("missing and user confirms windows install succeeds", func(t *testing.T) {
|
||||
if runtime.GOOS == "windows" {
|
||||
t.Skip("uses POSIX shell fake binaries")
|
||||
}
|
||||
|
||||
homeDir := t.TempDir()
|
||||
setTestHome(t, homeDir)
|
||||
tmpDir := t.TempDir()
|
||||
t.Setenv("PATH", tmpDir)
|
||||
openCodeGOOS = "windows"
|
||||
|
||||
installLog := filepath.Join(tmpDir, "npm.log")
|
||||
opencodePath := filepath.Join(homeDir, ".opencode", "bin", "opencode.exe")
|
||||
npmScript := fmt.Sprintf(`#!/bin/sh
|
||||
echo "$@" >> %q
|
||||
/bin/mkdir -p %q
|
||||
/bin/cat > %q <<'EOS'
|
||||
@echo off
|
||||
exit /b 0
|
||||
EOS
|
||||
/bin/chmod +x %q
|
||||
exit 0
|
||||
`, installLog, filepath.Dir(opencodePath), opencodePath, opencodePath)
|
||||
if err := os.WriteFile(filepath.Join(tmpDir, "npm"), []byte(npmScript), 0o755); err != nil {
|
||||
t.Fatalf("failed to write fake npm: %v", err)
|
||||
}
|
||||
|
||||
withConfirm(t, func(prompt string) (bool, error) {
|
||||
return true, nil
|
||||
})
|
||||
|
||||
bin, err := ensureOpenCodeInstalled()
|
||||
if err != nil {
|
||||
t.Fatalf("ensureOpenCodeInstalled() error = %v", err)
|
||||
}
|
||||
if bin != opencodePath {
|
||||
t.Fatalf("bin = %q, want %q", bin, opencodePath)
|
||||
}
|
||||
|
||||
logData, err := os.ReadFile(installLog)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to read install log: %v", err)
|
||||
}
|
||||
if !strings.Contains(string(logData), "install -g opencode-ai@latest") {
|
||||
t.Fatalf("expected npm install command in log, got:\n%s", string(logData))
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("install command fails", func(t *testing.T) {
|
||||
if runtime.GOOS == "windows" {
|
||||
t.Skip("uses POSIX shell fake binaries")
|
||||
}
|
||||
|
||||
setTestHome(t, t.TempDir())
|
||||
tmpDir := t.TempDir()
|
||||
t.Setenv("PATH", tmpDir)
|
||||
openCodeGOOS = "linux"
|
||||
writeFakeBinary(t, tmpDir, "curl")
|
||||
if err := os.WriteFile(filepath.Join(tmpDir, "bash"), []byte("#!/bin/sh\nexit 1\n"), 0o755); err != nil {
|
||||
t.Fatalf("failed to write fake bash: %v", err)
|
||||
}
|
||||
|
||||
withConfirm(t, func(prompt string) (bool, error) {
|
||||
return true, nil
|
||||
})
|
||||
|
||||
_, err := ensureOpenCodeInstalled()
|
||||
if err == nil || !strings.Contains(err.Error(), "failed to install opencode") {
|
||||
t.Fatalf("expected install failure error, got %v", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestOpenCodeInstallerCommand(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
goos string
|
||||
wantBin string
|
||||
wantParts []string
|
||||
wantErr bool
|
||||
}{
|
||||
{
|
||||
name: "linux",
|
||||
goos: "linux",
|
||||
wantBin: "bash",
|
||||
wantParts: []string{"-c", "set -o pipefail", "https://opencode.ai/install"},
|
||||
},
|
||||
{
|
||||
name: "darwin",
|
||||
goos: "darwin",
|
||||
wantBin: "bash",
|
||||
wantParts: []string{"-c", "set -o pipefail", "https://opencode.ai/install"},
|
||||
},
|
||||
{
|
||||
name: "windows",
|
||||
goos: "windows",
|
||||
wantBin: "npm",
|
||||
wantParts: []string{"install", "-g", "opencode-ai@latest"},
|
||||
},
|
||||
{
|
||||
name: "unsupported",
|
||||
goos: "plan9",
|
||||
wantErr: true,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
bin, args, err := openCodeInstallerCommand(tt.goos)
|
||||
if tt.wantErr {
|
||||
if err == nil {
|
||||
t.Fatal("expected error")
|
||||
}
|
||||
return
|
||||
}
|
||||
if err != nil {
|
||||
t.Fatalf("openCodeInstallerCommand() error = %v", err)
|
||||
}
|
||||
if bin != tt.wantBin {
|
||||
t.Fatalf("bin = %q, want %q", bin, tt.wantBin)
|
||||
}
|
||||
joined := strings.Join(args, " ")
|
||||
for _, want := range tt.wantParts {
|
||||
if !strings.Contains(joined, want) {
|
||||
t.Fatalf("args %q missing %q", joined, want)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// Verify that the BackfillsCloudModelLimitOnExistingEntry test from the old
|
||||
// file-based approach is covered by the new inline config approach.
|
||||
func TestOpenCodeEdit_CloudModelLimitStructure(t *testing.T) {
|
||||
|
||||
+2
-2
@@ -351,7 +351,7 @@ func npmArgs(prefix string, args ...string) []string {
|
||||
}
|
||||
|
||||
func ensurePiWebSearchPackage(bin string) {
|
||||
if !shouldManagePiWebSearch() {
|
||||
if !shouldManageOllamaWebSearch() {
|
||||
fmt.Fprintf(os.Stderr, "%sCloud is disabled; skipping %s setup.%s\n", ansiGray, piWebSearchPkg, ansiReset)
|
||||
return
|
||||
}
|
||||
@@ -395,7 +395,7 @@ func ensurePiWebSearchPackage(bin string) {
|
||||
fmt.Fprintf(os.Stderr, "%s ✓ Updated %s%s\n", ansiGreen, piWebSearchPkg, ansiReset)
|
||||
}
|
||||
|
||||
func shouldManagePiWebSearch() bool {
|
||||
func shouldManageOllamaWebSearch() bool {
|
||||
client, err := api.ClientFromEnvironment()
|
||||
if err != nil {
|
||||
return true
|
||||
|
||||
+39
-5
@@ -33,7 +33,7 @@ type IntegrationInfo struct {
|
||||
Description string
|
||||
}
|
||||
|
||||
var launcherIntegrationOrder = []string{"claude", "codex-app", "hermes", "openclaw", "opencode", "codex", "copilot", "cline", "droid", "pi", "pool", "qwen"}
|
||||
var launcherIntegrationOrder = []string{"claude", "chatgpt", "hermes", "openclaw", "opencode", "hermes-desktop", "codex", "copilot", "omp", "cline", "droid", "pi", "pool", "qwen"}
|
||||
|
||||
var integrationSpecs = []*IntegrationSpec{
|
||||
{
|
||||
@@ -45,6 +45,10 @@ var integrationSpecs = []*IntegrationSpec{
|
||||
_, err := (&Claude{}).findPath()
|
||||
return err == nil
|
||||
},
|
||||
EnsureInstalled: func() error {
|
||||
_, err := ensureClaudeInstalled()
|
||||
return err
|
||||
},
|
||||
URL: "https://code.claude.com/docs/en/quickstart",
|
||||
},
|
||||
},
|
||||
@@ -91,15 +95,15 @@ var integrationSpecs = []*IntegrationSpec{
|
||||
},
|
||||
},
|
||||
{
|
||||
Name: "codex-app",
|
||||
Name: chatGPTIntegrationName,
|
||||
Runner: &CodexApp{},
|
||||
Aliases: []string{"codex-desktop", "codex-gui"},
|
||||
Description: "An AI agent you can delegate real work to, by OpenAI",
|
||||
Aliases: []string{codexAppIntegrationName, "codex-desktop", "codex-gui"},
|
||||
Description: "Complete work with ChatGPT",
|
||||
Install: IntegrationInstallSpec{
|
||||
CheckInstalled: func() bool {
|
||||
return codexAppInstalled()
|
||||
},
|
||||
URL: "https://developers.openai.com/codex/quickstart",
|
||||
URL: "https://chatgpt.com/download",
|
||||
},
|
||||
},
|
||||
{
|
||||
@@ -153,9 +157,25 @@ var integrationSpecs = []*IntegrationSpec{
|
||||
_, ok := findOpenCode()
|
||||
return ok
|
||||
},
|
||||
EnsureInstalled: func() error {
|
||||
_, err := ensureOpenCodeInstalled()
|
||||
return err
|
||||
},
|
||||
URL: "https://opencode.ai",
|
||||
},
|
||||
},
|
||||
{
|
||||
Name: "omp",
|
||||
Runner: &OMP{},
|
||||
Description: "AI coding agent with IDE integration",
|
||||
Install: IntegrationInstallSpec{
|
||||
CheckInstalled: func() bool {
|
||||
_, err := (&OMP{}).findPath()
|
||||
return err == nil
|
||||
},
|
||||
URL: "https://omp.sh",
|
||||
},
|
||||
},
|
||||
{
|
||||
Name: "openclaw",
|
||||
Runner: &Openclaw{},
|
||||
@@ -220,6 +240,20 @@ var integrationSpecs = []*IntegrationSpec{
|
||||
URL: "https://hermes-agent.nousresearch.com/docs/getting-started/installation/",
|
||||
},
|
||||
},
|
||||
{
|
||||
Name: "hermes-desktop",
|
||||
Runner: &HermesDesktop{},
|
||||
Description: "Desktop app for Hermes Agent by Nous Research",
|
||||
Install: IntegrationInstallSpec{
|
||||
CheckInstalled: func() bool {
|
||||
return (&Hermes{}).installed()
|
||||
},
|
||||
EnsureInstalled: func() error {
|
||||
return (&Hermes{}).ensureInstalledFor("hermes-desktop")
|
||||
},
|
||||
URL: "https://hermes-agent.nousresearch.com/docs/getting-started/installation/",
|
||||
},
|
||||
},
|
||||
{
|
||||
Name: "vscode",
|
||||
Runner: &VSCode{},
|
||||
|
||||
@@ -61,6 +61,14 @@ func TestEditorRunsDoNotRewriteConfig(t *testing.T) {
|
||||
return filepath.Join(home, ".kimi", "config.toml")
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "omp",
|
||||
binary: "omp",
|
||||
runner: &OMP{},
|
||||
checkPath: func(home string) string {
|
||||
return filepath.Join(home, ".omp", "agent", "models.yml")
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
|
||||
@@ -27,10 +27,18 @@ var errCancelled = ErrCancelled
|
||||
// When set, ConfirmPrompt delegates to it instead of using raw terminal I/O.
|
||||
var DefaultConfirmPrompt func(prompt string, options ConfirmOptions) (bool, error)
|
||||
|
||||
type ConfirmDefault int
|
||||
|
||||
const (
|
||||
ConfirmDefaultYes ConfirmDefault = iota
|
||||
ConfirmDefaultNo
|
||||
)
|
||||
|
||||
// ConfirmOptions customizes labels for confirmation prompts.
|
||||
type ConfirmOptions struct {
|
||||
YesLabel string
|
||||
NoLabel string
|
||||
Default ConfirmDefault
|
||||
}
|
||||
|
||||
// SingleSelector is a function type for single item selection.
|
||||
@@ -111,7 +119,12 @@ func ConfirmPromptWithOptions(prompt string, options ConfirmOptions) (bool, erro
|
||||
}
|
||||
defer term.Restore(fd, oldState)
|
||||
|
||||
fmt.Fprintf(os.Stderr, "%s (\033[1my\033[0m/n) ", prompt)
|
||||
defaultNo := options.Default == ConfirmDefaultNo
|
||||
if defaultNo {
|
||||
fmt.Fprintf(os.Stderr, "%s (y/\033[1mN\033[0m) ", prompt)
|
||||
} else {
|
||||
fmt.Fprintf(os.Stderr, "%s (\033[1my\033[0m/n) ", prompt)
|
||||
}
|
||||
|
||||
buf := make([]byte, 1)
|
||||
for {
|
||||
@@ -120,7 +133,14 @@ func ConfirmPromptWithOptions(prompt string, options ConfirmOptions) (bool, erro
|
||||
}
|
||||
|
||||
switch buf[0] {
|
||||
case 'Y', 'y', 13:
|
||||
case 'Y', 'y':
|
||||
fmt.Fprintf(os.Stderr, "yes\r\n")
|
||||
return true, nil
|
||||
case 13:
|
||||
if defaultNo {
|
||||
fmt.Fprintf(os.Stderr, "no\r\n")
|
||||
return false, nil
|
||||
}
|
||||
fmt.Fprintf(os.Stderr, "yes\r\n")
|
||||
return true, nil
|
||||
case 'N', 'n', 27, 3:
|
||||
|
||||
@@ -0,0 +1,109 @@
|
||||
package launch
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
// SpinnerFrames are the braille spinner frames used by the bubbletea TUIs in
|
||||
// this codebase (sign-in, upgrade). StartSpinner uses the same frames for its
|
||||
// fallback so the restart spinner matches the look of those flows.
|
||||
var SpinnerFrames = []string{"⠋", "⠙", "⠹", "⠸", "⠼", "⠴", "⠦", "⠧", "⠇", "⠏"}
|
||||
|
||||
// DefaultSpinner, when set, starts an animated spinner displaying message and
|
||||
// returns a *Spinner. cmd/cmd.go registers a bubbletea implementation from
|
||||
// cmd/tui; when unset (or when it returns nil, e.g. no TTY) StartSpinner falls
|
||||
// back to a simple ANSI spinner using SpinnerFrames.
|
||||
var DefaultSpinner func(message string) *Spinner
|
||||
|
||||
// Spinner is a handle on a running animated spinner. Stop halts the spinner
|
||||
// and clears its line (it blocks until the spinner has fully stopped and is
|
||||
// safe to call multiple times). Cancelled returns a channel that is closed if
|
||||
// the user interrupts the spinner (e.g. with Ctrl+C); wait loops can select on
|
||||
// it to abort early. For the non-interactive ANSI fallback the channel is never
|
||||
// closed because Ctrl+C raises SIGINT and terminates the process directly.
|
||||
type Spinner struct {
|
||||
stop func()
|
||||
cancelled chan struct{}
|
||||
}
|
||||
|
||||
// NewSpinner builds a Spinner from a stop function and a cancellation channel.
|
||||
// It is intended for implementations of DefaultSpinner (e.g. the bubbletea
|
||||
// spinner in cmd/tui). stop must be safe to call multiple times; cancelled is
|
||||
// closed by the implementation when the user interrupts the spinner, or left
|
||||
// open when interruption is handled another way (e.g. SIGINT).
|
||||
func NewSpinner(stop func(), cancelled chan struct{}) *Spinner {
|
||||
return &Spinner{stop: stop, cancelled: cancelled}
|
||||
}
|
||||
|
||||
// Stop halts the spinner and clears its line. It is a no-op when the spinner
|
||||
// already stopped (for example after the user cancelled it).
|
||||
func (s *Spinner) Stop() {
|
||||
if s != nil && s.stop != nil {
|
||||
s.stop()
|
||||
}
|
||||
}
|
||||
|
||||
// Cancelled returns a channel that is closed when the user interrupts the
|
||||
// spinner. Callers may select on it to abort a blocking wait.
|
||||
func (s *Spinner) Cancelled() <-chan struct{} {
|
||||
if s == nil {
|
||||
return nil
|
||||
}
|
||||
return s.cancelled
|
||||
}
|
||||
|
||||
// StartSpinner begins an animated spinner displaying message and returns a
|
||||
// *Spinner handle. It uses DefaultSpinner when available, otherwise a simple
|
||||
// ANSI fallback that renders SpinnerFrames to stderr without requiring a TTY.
|
||||
func StartSpinner(message string) *Spinner {
|
||||
if DefaultSpinner != nil {
|
||||
if s := DefaultSpinner(message); s != nil {
|
||||
return s
|
||||
}
|
||||
}
|
||||
return defaultSpinner(message)
|
||||
}
|
||||
|
||||
// defaultSpinner renders SpinnerFrames to stderr without requiring a TTY. It
|
||||
// runs in its own goroutine so it can animate while a caller polls; Stop
|
||||
// signals the goroutine to exit, waits for it, and clears the spinner line.
|
||||
func defaultSpinner(message string) *Spinner {
|
||||
frames := SpinnerFrames
|
||||
frame := 0
|
||||
fmt.Fprintf(os.Stderr, "\r\033[90m%s %s\033[0m", message, frames[0])
|
||||
|
||||
done := make(chan struct{})
|
||||
exited := make(chan struct{})
|
||||
var once sync.Once
|
||||
|
||||
go func() {
|
||||
ticker := time.NewTicker(100 * time.Millisecond)
|
||||
defer ticker.Stop()
|
||||
for {
|
||||
select {
|
||||
case <-done:
|
||||
close(exited)
|
||||
return
|
||||
case <-ticker.C:
|
||||
frame++
|
||||
fmt.Fprintf(os.Stderr, "\r\033[90m%s %s\033[0m", message, frames[frame%len(frames)])
|
||||
}
|
||||
}
|
||||
}()
|
||||
|
||||
stop := func() {
|
||||
once.Do(func() {
|
||||
close(done)
|
||||
<-exited
|
||||
fmt.Fprintf(os.Stderr, "\r\033[K")
|
||||
})
|
||||
}
|
||||
|
||||
// Ctrl+C in non-raw mode raises SIGINT and terminates the process by
|
||||
// default (the launch flow installs no SIGINT handler), so this cancelled
|
||||
// channel is intentionally never closed.
|
||||
return &Spinner{stop: stop, cancelled: make(chan struct{})}
|
||||
}
|
||||
@@ -0,0 +1,406 @@
|
||||
package chat
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"slices"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
tea "github.com/charmbracelet/bubbletea"
|
||||
|
||||
coreagent "github.com/ollama/ollama/agent"
|
||||
)
|
||||
|
||||
type chatApprovalChoice struct {
|
||||
label string
|
||||
key string
|
||||
allow bool
|
||||
allowTools bool
|
||||
allowAll bool
|
||||
reason string
|
||||
}
|
||||
|
||||
var chatApprovalChoices = []chatApprovalChoice{
|
||||
{label: "Approve once", key: "1", allow: true},
|
||||
{label: "Always allow tool", key: "2", allow: true, allowTools: true},
|
||||
{label: "Deny", key: "3", reason: "Tool execution denied."},
|
||||
}
|
||||
|
||||
type chatApprovalPrompt struct {
|
||||
request coreagent.ApprovalRequest
|
||||
reply chan<- coreagent.Approval
|
||||
cursor int
|
||||
}
|
||||
|
||||
func (m chatModel) approvalPrompterForRun(controller *chatApprovalController) coreagent.ApprovalPrompter {
|
||||
if m.opts.ApprovalPrompter != nil {
|
||||
return m.opts.ApprovalPrompter
|
||||
}
|
||||
return controller
|
||||
}
|
||||
|
||||
func (m *chatModel) ensureApprovalState() *coreagent.ApprovalState {
|
||||
if m.approvalState == nil {
|
||||
m.approvalState = &coreagent.ApprovalState{}
|
||||
m.approvalState.Set(m.defaultAllowAll, nil)
|
||||
}
|
||||
return m.approvalState
|
||||
}
|
||||
|
||||
func (m *chatModel) resetApprovalState() {
|
||||
m.approvalState = &coreagent.ApprovalState{}
|
||||
m.approvalState.Set(m.defaultAllowAll, nil)
|
||||
}
|
||||
|
||||
func (m chatModel) allowAllToolsEnabled() bool {
|
||||
if m.approvalState == nil {
|
||||
return m.defaultAllowAll
|
||||
}
|
||||
return m.approvalState.AllGranted()
|
||||
}
|
||||
|
||||
func (m *chatModel) setAllowAllTools(allowAll bool) {
|
||||
if allowAll {
|
||||
m.ensureApprovalState().GrantAll()
|
||||
} else {
|
||||
m.ensureApprovalState().Set(false, nil)
|
||||
}
|
||||
m.opts.AllowAllTools = allowAll
|
||||
}
|
||||
|
||||
func (m *chatModel) openApprovalPrompt(msg chatApprovalPromptMsg) {
|
||||
m.approvalPrompt = &chatApprovalPrompt{request: msg.request, reply: msg.reply}
|
||||
m.status = "approval required"
|
||||
m.thinking = false
|
||||
m.thinkingTokens = 0
|
||||
m.upsertApprovalToolEntries(msg.request)
|
||||
}
|
||||
|
||||
func (m *chatModel) togglePermissionMode() (tea.Model, tea.Cmd) {
|
||||
m.setAllowAllTools(!m.allowAllToolsEnabled())
|
||||
if m.allowAllToolsEnabled() {
|
||||
m.permissionNotice = "full access enabled"
|
||||
m.status = "full access enabled"
|
||||
if m.approvalPrompt != nil {
|
||||
updated, cmd := m.resolveApprovalPrompt(chatApprovalChoice{allow: true, allowAll: true})
|
||||
if model, ok := updated.(chatModel); ok {
|
||||
model.permissionNotice = "full access enabled"
|
||||
model.status = "full access enabled"
|
||||
return model, cmd
|
||||
}
|
||||
return updated, cmd
|
||||
}
|
||||
return *m, nil
|
||||
}
|
||||
m.permissionNotice = "review mode enabled"
|
||||
m.status = "review mode enabled"
|
||||
return *m, nil
|
||||
}
|
||||
|
||||
func (m *chatModel) upsertApprovalToolEntries(request coreagent.ApprovalRequest) {
|
||||
for _, call := range request.Calls {
|
||||
idx := m.findToolEntry(call.ToolCallID)
|
||||
if idx < 0 {
|
||||
m.groupCompletedToolHistory()
|
||||
m.entries = append(m.entries, newChatEntry(chatEntry{role: "tool"}))
|
||||
idx = len(m.entries) - 1
|
||||
}
|
||||
m.entries[idx].detail = call.ToolName
|
||||
m.entries[idx].label = toolInvocationLabel(call.ToolName, call.Args)
|
||||
m.entries[idx].status = "approval"
|
||||
m.entries[idx].toolID = call.ToolCallID
|
||||
m.entries[idx].args = call.Args
|
||||
m.entries[idx].startedAt = time.Now()
|
||||
m.applyToolOutputModeTo(idx)
|
||||
m.markEntryDirty(idx)
|
||||
}
|
||||
}
|
||||
|
||||
func (m chatModel) updateApprovalPrompt(msg tea.KeyMsg) (tea.Model, tea.Cmd) {
|
||||
switch msg.Type {
|
||||
case tea.KeyLeft, tea.KeyUp:
|
||||
m.moveApprovalChoice(-1)
|
||||
case tea.KeyRight, tea.KeyDown, tea.KeyTab:
|
||||
m.moveApprovalChoice(1)
|
||||
case tea.KeyRunes:
|
||||
switch string(msg.Runes) {
|
||||
case "1", "2", "3":
|
||||
choice := chatApprovalChoices[int(msg.Runes[0]-'1')]
|
||||
return m.resolveApprovalPrompt(choice)
|
||||
}
|
||||
case tea.KeyEnter:
|
||||
choice := chatApprovalChoices[clamp(m.approvalPrompt.cursor, 0, len(chatApprovalChoices)-1)]
|
||||
return m.resolveApprovalPrompt(choice)
|
||||
case tea.KeyEsc, tea.KeyCtrlC:
|
||||
return m.resolveApprovalPrompt(chatApprovalChoice{reason: "Tool execution denied."})
|
||||
}
|
||||
return m, nil
|
||||
}
|
||||
|
||||
func (m *chatModel) moveApprovalChoice(delta int) {
|
||||
if m.approvalPrompt == nil {
|
||||
return
|
||||
}
|
||||
m.approvalPrompt.cursor = (m.approvalPrompt.cursor + delta) % len(chatApprovalChoices)
|
||||
if m.approvalPrompt.cursor < 0 {
|
||||
m.approvalPrompt.cursor += len(chatApprovalChoices)
|
||||
}
|
||||
m.markApprovalPromptEntryDirty()
|
||||
}
|
||||
|
||||
func (m *chatModel) markApprovalPromptEntryDirty() {
|
||||
if m.approvalPrompt == nil {
|
||||
return
|
||||
}
|
||||
for _, call := range m.approvalPrompt.request.Calls {
|
||||
if idx := m.findToolEntry(call.ToolCallID); idx >= 0 {
|
||||
m.markEntryDirty(idx)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (m chatModel) resolveApprovalPrompt(choice chatApprovalChoice) (tea.Model, tea.Cmd) {
|
||||
if m.approvalPrompt == nil {
|
||||
return m, nil
|
||||
}
|
||||
printedLines := m.flowPrintedLines
|
||||
var printedTranscript []string
|
||||
if printedLines > 0 {
|
||||
printedTranscript = slices.Clone(m.transcriptLines(m.viewWidth()))
|
||||
}
|
||||
prompt := m.approvalPrompt
|
||||
m.approvalPrompt = nil
|
||||
m.status = "running"
|
||||
if !choice.allow {
|
||||
m.status = "denied"
|
||||
}
|
||||
if choice.allowAll {
|
||||
m.setAllowAllTools(true)
|
||||
}
|
||||
allowScopes := approvalScopes(prompt.request)
|
||||
if choice.allowTools {
|
||||
m.ensureApprovalState().GrantScopes(allowScopes)
|
||||
}
|
||||
for _, call := range prompt.request.Calls {
|
||||
if idx := m.findToolEntry(call.ToolCallID); idx >= 0 && m.entries[idx].status == "approval" {
|
||||
if !choice.allow {
|
||||
m.entries[idx].status = "error"
|
||||
m.entries[idx].err = choice.reason
|
||||
if m.entries[idx].err == "" {
|
||||
m.entries[idx].err = "Tool execution denied."
|
||||
}
|
||||
} else {
|
||||
m.entries[idx].status = "queued"
|
||||
}
|
||||
m.markEntryDirty(idx)
|
||||
}
|
||||
}
|
||||
result := coreagent.Approval{Allow: choice.allow, AllowAll: choice.allowAll, Reason: choice.reason}
|
||||
if choice.allowTools {
|
||||
result.AllowScopes = allowScopes
|
||||
}
|
||||
prompt.reply <- result
|
||||
return m.withFlowTranscriptRefreshAfter(printedTranscript, printedLines, waitForChatMsg(m.events))
|
||||
}
|
||||
|
||||
func (m chatModel) renderApprovalPromptLines(width int) []string {
|
||||
prompt := m.approvalPrompt
|
||||
if prompt == nil {
|
||||
return nil
|
||||
}
|
||||
if width <= 0 {
|
||||
width = 80
|
||||
}
|
||||
bodyWidth := max(20, width-2)
|
||||
|
||||
var lines []string
|
||||
if len(prompt.request.Calls) <= 1 {
|
||||
detail := approvalRequestDetail(prompt.request, bodyWidth)
|
||||
if detail == "" {
|
||||
label := "Tool request"
|
||||
if len(prompt.request.Calls) == 1 {
|
||||
label = toolDisplayName(prompt.request.Calls[0].ToolName)
|
||||
}
|
||||
lines = append(lines, wrapChatText(fmt.Sprintf("%s wants to run", label), width)...)
|
||||
} else {
|
||||
lines = append(lines, indentLines(splitRenderedBody(detail), " ")...)
|
||||
}
|
||||
lines = append(lines, "")
|
||||
}
|
||||
|
||||
lines = append(lines, indentLines(renderApprovalChoices(prompt.request, prompt.cursor, bodyWidth), " ")...)
|
||||
return lines
|
||||
}
|
||||
|
||||
func approvalRequestDetail(request coreagent.ApprovalRequest, width int) string {
|
||||
if len(request.Calls) == 0 {
|
||||
return ""
|
||||
}
|
||||
if len(request.Calls) == 1 {
|
||||
return approvalToolCallDetail(request.Calls[0], width)
|
||||
}
|
||||
lines := make([]string, 0, len(request.Calls))
|
||||
for _, call := range request.Calls {
|
||||
lines = append(lines, toolInvocationLabel(call.ToolName, call.Args))
|
||||
}
|
||||
return chatMetaStyle.Render(strings.Join(lines, "\n"))
|
||||
}
|
||||
|
||||
func approvalToolCallDetail(call coreagent.ApprovalToolCall, width int) string {
|
||||
if isShellToolName(call.ToolName) {
|
||||
command, ok := rawStringArg(call.Args, "command")
|
||||
if !ok {
|
||||
return ""
|
||||
}
|
||||
return strings.Join(wrapChatText(shellPromptPrefix(call.ToolName)+command, width), "\n")
|
||||
}
|
||||
switch call.ToolName {
|
||||
case "edit":
|
||||
path, ok := rawStringArg(call.Args, "path")
|
||||
if !ok {
|
||||
return ""
|
||||
}
|
||||
var lines []string
|
||||
lines = append(lines, "path: "+path)
|
||||
if oldText, ok := rawStringArg(call.Args, "old_text"); ok {
|
||||
lines = append(lines, fmt.Sprintf("old_text: %d chars", len([]rune(oldText))))
|
||||
}
|
||||
if newText, ok := rawStringArg(call.Args, "new_text"); ok {
|
||||
lines = append(lines, fmt.Sprintf("new_text: %d chars", len([]rune(newText))))
|
||||
}
|
||||
return chatMetaStyle.Render(strings.Join(lines, "\n"))
|
||||
default:
|
||||
if len(call.Args) == 0 {
|
||||
return ""
|
||||
}
|
||||
return strings.Join(renderToolCallArgs(call.Args, width), "\n")
|
||||
}
|
||||
}
|
||||
|
||||
func renderApprovalChoices(request coreagent.ApprovalRequest, cursor int, width int) []string {
|
||||
var lines []string
|
||||
for i, choice := range chatApprovalChoices {
|
||||
label := choice.key + ". " + approvalChoiceLabel(choice, request)
|
||||
wrapped := wrapChatText(label, max(20, width-2))
|
||||
if i == clamp(cursor, 0, len(chatApprovalChoices)-1) {
|
||||
for j, line := range wrapped {
|
||||
if j == 0 {
|
||||
lines = append(lines, chatPickerSelectedStyle.Render("> "+line))
|
||||
} else {
|
||||
lines = append(lines, chatPickerSelectedStyle.Render(" "+line))
|
||||
}
|
||||
}
|
||||
} else {
|
||||
for _, line := range wrapped {
|
||||
lines = append(lines, chatPickerTextStyle.Render(" "+line))
|
||||
}
|
||||
}
|
||||
}
|
||||
return lines
|
||||
}
|
||||
|
||||
func approvalChoiceLabel(choice chatApprovalChoice, request coreagent.ApprovalRequest) string {
|
||||
if !choice.allowTools {
|
||||
return choice.label
|
||||
}
|
||||
scopes := approvalScopes(request)
|
||||
if len(scopes) == 1 {
|
||||
call := approvalCallForScope(request, scopes[0])
|
||||
if isShellToolName(call.ToolName) {
|
||||
if command, ok := rawStringArg(call.Args, "command"); ok && strings.TrimSpace(command) != "" {
|
||||
return "Always allow this command"
|
||||
}
|
||||
}
|
||||
return "Always allow " + toolDisplayName(call.ToolName)
|
||||
}
|
||||
return "Always allow these requests"
|
||||
}
|
||||
|
||||
func approvalScopes(request coreagent.ApprovalRequest) []string {
|
||||
seen := make(map[string]bool, len(request.Calls))
|
||||
var scopes []string
|
||||
for _, call := range request.Calls {
|
||||
scope := approvalScope(call)
|
||||
if scope == "" || seen[scope] {
|
||||
continue
|
||||
}
|
||||
seen[scope] = true
|
||||
scopes = append(scopes, scope)
|
||||
}
|
||||
return scopes
|
||||
}
|
||||
|
||||
func approvalCallForScope(request coreagent.ApprovalRequest, scope string) coreagent.ApprovalToolCall {
|
||||
for _, call := range request.Calls {
|
||||
if approvalScope(call) == scope {
|
||||
return call
|
||||
}
|
||||
}
|
||||
return coreagent.ApprovalToolCall{}
|
||||
}
|
||||
|
||||
func approvalScope(call coreagent.ApprovalToolCall) string {
|
||||
if scope := strings.TrimSpace(call.ApprovalScope); scope != "" {
|
||||
return scope
|
||||
}
|
||||
return strings.TrimSpace(call.ToolName)
|
||||
}
|
||||
|
||||
type chatApprovalPrompter struct {
|
||||
ch chan<- tea.Msg
|
||||
}
|
||||
|
||||
func (p chatApprovalPrompter) PromptApproval(ctx context.Context, request coreagent.ApprovalRequest) (coreagent.Approval, error) {
|
||||
reply := make(chan coreagent.Approval, 1)
|
||||
select {
|
||||
case p.ch <- chatApprovalPromptMsg{request: request, reply: reply}:
|
||||
case <-ctx.Done():
|
||||
return coreagent.Approval{Reason: "Tool approval canceled."}, nil
|
||||
}
|
||||
|
||||
select {
|
||||
case result := <-reply:
|
||||
return result, nil
|
||||
case <-ctx.Done():
|
||||
return coreagent.Approval{Reason: "Tool approval canceled."}, nil
|
||||
}
|
||||
}
|
||||
|
||||
type chatApprovalController struct {
|
||||
ch chan<- tea.Msg
|
||||
state *coreagent.ApprovalState
|
||||
}
|
||||
|
||||
func newChatApprovalController(ch chan<- tea.Msg, state *coreagent.ApprovalState) *chatApprovalController {
|
||||
return &chatApprovalController{
|
||||
ch: ch,
|
||||
state: state,
|
||||
}
|
||||
}
|
||||
|
||||
func (c *chatApprovalController) PromptApproval(ctx context.Context, request coreagent.ApprovalRequest) (coreagent.Approval, error) {
|
||||
if result, ok := c.preapproved(request); ok {
|
||||
return result, nil
|
||||
}
|
||||
return chatApprovalPrompter{ch: c.ch}.PromptApproval(ctx, request)
|
||||
}
|
||||
|
||||
func (c *chatApprovalController) preapproved(request coreagent.ApprovalRequest) (coreagent.Approval, bool) {
|
||||
if c == nil {
|
||||
return coreagent.Approval{}, false
|
||||
}
|
||||
if c.state.AllGranted() {
|
||||
return coreagent.Approval{Allow: true, AllowAll: true}, true
|
||||
}
|
||||
scopes := approvalScopes(request)
|
||||
if len(scopes) == 0 {
|
||||
return coreagent.Approval{}, false
|
||||
}
|
||||
for _, scope := range scopes {
|
||||
if !c.state.Allows(scope) {
|
||||
return coreagent.Approval{}, false
|
||||
}
|
||||
}
|
||||
return coreagent.Approval{Allow: true, AllowScopes: scopes}, true
|
||||
}
|
||||
@@ -0,0 +1,495 @@
|
||||
package chat
|
||||
|
||||
import (
|
||||
"context"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
tea "github.com/charmbracelet/bubbletea"
|
||||
|
||||
coreagent "github.com/ollama/ollama/agent"
|
||||
)
|
||||
|
||||
func testApprovalRequest() coreagent.ApprovalRequest {
|
||||
return coreagent.ApprovalRequest{
|
||||
WorkingDir: "/repo",
|
||||
Calls: []coreagent.ApprovalToolCall{{
|
||||
ToolCallID: "call-1",
|
||||
ToolName: "edit",
|
||||
Args: map[string]any{"path": "note.txt"},
|
||||
ApprovalScope: "edit",
|
||||
}},
|
||||
}
|
||||
}
|
||||
|
||||
func testApprovalState(allowAll bool, scopes map[string]bool) *coreagent.ApprovalState {
|
||||
state := &coreagent.ApprovalState{}
|
||||
state.Set(allowAll, scopes)
|
||||
return state
|
||||
}
|
||||
|
||||
func TestChatApprovalApprovesOnce(t *testing.T) {
|
||||
reply := make(chan coreagent.Approval, 1)
|
||||
m := chatModel{
|
||||
approvalPrompt: &chatApprovalPrompt{
|
||||
request: testApprovalRequest(),
|
||||
reply: reply,
|
||||
},
|
||||
events: make(chan tea.Msg),
|
||||
}
|
||||
|
||||
updated, cmd := m.updateApprovalPrompt(tea.KeyMsg{Type: tea.KeyEnter})
|
||||
if cmd == nil {
|
||||
t.Fatal("approval should resume waiting for agent events")
|
||||
}
|
||||
fm := updated.(chatModel)
|
||||
if fm.approvalPrompt != nil {
|
||||
t.Fatal("approval prompt should close")
|
||||
}
|
||||
result := <-reply
|
||||
if !result.Allow || result.AllowAll {
|
||||
t.Fatalf("approval = %#v, want allow once", result)
|
||||
}
|
||||
}
|
||||
|
||||
func TestChatApprovalAllowsTool(t *testing.T) {
|
||||
reply := make(chan coreagent.Approval, 1)
|
||||
m := chatModel{
|
||||
approvalPrompt: &chatApprovalPrompt{
|
||||
request: testApprovalRequest(),
|
||||
reply: reply,
|
||||
cursor: 1,
|
||||
},
|
||||
events: make(chan tea.Msg),
|
||||
}
|
||||
|
||||
updated, _ := m.updateApprovalPrompt(tea.KeyMsg{Type: tea.KeyEnter})
|
||||
fm := updated.(chatModel)
|
||||
if fm.allowAllToolsEnabled() {
|
||||
t.Fatal("allowing a tool should not enable full access")
|
||||
}
|
||||
if !fm.approvalState.Allows("edit") {
|
||||
t.Fatal("edit scope was not saved")
|
||||
}
|
||||
result := <-reply
|
||||
if !result.Allow || result.AllowAll || len(result.AllowScopes) != 1 || result.AllowScopes[0] != "edit" {
|
||||
t.Fatalf("approval = %#v, want per-tool approval", result)
|
||||
}
|
||||
}
|
||||
|
||||
func TestChatApprovalLabelsSecondChoiceAsPerTool(t *testing.T) {
|
||||
lines := stripANSI(strings.Join(renderApprovalChoices(testApprovalRequest(), 1, 80), "\n"))
|
||||
if !strings.Contains(lines, "2. Always allow Edit") {
|
||||
t.Fatalf("approval choices = %q, want per-tool option", lines)
|
||||
}
|
||||
if strings.Contains(lines, "Approve all") {
|
||||
t.Fatalf("approval choices = %q, should not offer approve all as option 2", lines)
|
||||
}
|
||||
}
|
||||
|
||||
func TestChatApprovalLabelsShellChoiceAsCommandScoped(t *testing.T) {
|
||||
request := coreagent.ApprovalRequest{
|
||||
WorkingDir: "/repo",
|
||||
Calls: []coreagent.ApprovalToolCall{{
|
||||
ToolCallID: "call-1",
|
||||
ToolName: "bash",
|
||||
Args: map[string]any{"command": "pwd"},
|
||||
ApprovalScope: "bash\x00pwd",
|
||||
}},
|
||||
}
|
||||
lines := stripANSI(strings.Join(renderApprovalChoices(request, 1, 80), "\n"))
|
||||
if !strings.Contains(lines, "2. Always allow this command") {
|
||||
t.Fatalf("approval choices = %q, want command-scoped option", lines)
|
||||
}
|
||||
if strings.Contains(lines, "Always allow Bash") {
|
||||
t.Fatalf("approval choices = %q, should not offer top-level Bash approval", lines)
|
||||
}
|
||||
}
|
||||
|
||||
func TestChatApprovalUsesShellNameForPermissionPrompt(t *testing.T) {
|
||||
request := coreagent.ApprovalRequest{
|
||||
WorkingDir: "/repo",
|
||||
Calls: []coreagent.ApprovalToolCall{{
|
||||
ToolCallID: "call-1",
|
||||
ToolName: "bash",
|
||||
Args: map[string]any{"command": "pwd"},
|
||||
ApprovalScope: "bash\x00pwd",
|
||||
}},
|
||||
}
|
||||
|
||||
detail := stripANSI(approvalRequestDetail(request, 80))
|
||||
if !strings.Contains(detail, "$ pwd") {
|
||||
t.Fatalf("approval detail should show command prompt, got %q", detail)
|
||||
}
|
||||
|
||||
m := chatModel{}
|
||||
m.upsertApprovalToolEntries(request)
|
||||
if len(m.entries) != 1 {
|
||||
t.Fatalf("entries = %#v", m.entries)
|
||||
}
|
||||
line := stripANSI(toolStatusLine(m.entries[0]))
|
||||
if !strings.Contains(line, `Bash("pwd")`) || !strings.Contains(line, "needs approval") {
|
||||
t.Fatalf("approval status line = %q", line)
|
||||
}
|
||||
}
|
||||
|
||||
func TestChatApprovalPromptOmitsDuplicateBatchDetails(t *testing.T) {
|
||||
request := coreagent.ApprovalRequest{
|
||||
WorkingDir: "/repo",
|
||||
Calls: []coreagent.ApprovalToolCall{
|
||||
{
|
||||
ToolCallID: "call-1",
|
||||
ToolName: "bash",
|
||||
Args: map[string]any{"command": "git rev-parse --abbrev-ref HEAD"},
|
||||
ApprovalScope: "bash\x00git rev-parse --abbrev-ref HEAD",
|
||||
},
|
||||
{
|
||||
ToolCallID: "call-2",
|
||||
ToolName: "bash",
|
||||
Args: map[string]any{"command": "git branch -a"},
|
||||
ApprovalScope: "bash\x00git branch -a",
|
||||
},
|
||||
},
|
||||
}
|
||||
m := chatModel{
|
||||
approvalPrompt: &chatApprovalPrompt{request: request},
|
||||
}
|
||||
|
||||
lines := stripANSI(strings.Join(m.renderApprovalPromptLines(120), "\n"))
|
||||
if strings.Contains(lines, `Bash("git rev-parse --abbrev-ref HEAD")`) || strings.Contains(lines, `Bash("git branch -a")`) {
|
||||
t.Fatalf("batched approval prompt should not duplicate visible tool rows:\n%s", lines)
|
||||
}
|
||||
for _, want := range []string{"1. Approve once", "2. Always allow these requests", "3. Deny"} {
|
||||
if !strings.Contains(lines, want) {
|
||||
t.Fatalf("batched approval prompt missing %q:\n%s", want, lines)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestChatApprovalKeepsQueuedBatchCallsVisible(t *testing.T) {
|
||||
reply := make(chan coreagent.Approval, 1)
|
||||
request := coreagent.ApprovalRequest{
|
||||
WorkingDir: "/repo",
|
||||
Calls: []coreagent.ApprovalToolCall{
|
||||
{
|
||||
ToolCallID: "call-1",
|
||||
ToolName: "bash",
|
||||
Args: map[string]any{"command": "git rev-parse --abbrev-ref HEAD"},
|
||||
ApprovalScope: "bash\x00git rev-parse --abbrev-ref HEAD",
|
||||
},
|
||||
{
|
||||
ToolCallID: "call-2",
|
||||
ToolName: "bash",
|
||||
Args: map[string]any{"command": "git branch -a"},
|
||||
ApprovalScope: "bash\x00git branch -a",
|
||||
},
|
||||
},
|
||||
}
|
||||
m := chatModel{
|
||||
running: true,
|
||||
events: make(chan tea.Msg),
|
||||
}
|
||||
m.openApprovalPrompt(chatApprovalPromptMsg{request: request, reply: reply})
|
||||
|
||||
updated, _ := m.resolveApprovalPrompt(chatApprovalChoice{allow: true})
|
||||
m = updated.(chatModel)
|
||||
m.applyAgentEvent(coreagent.Event{
|
||||
Type: coreagent.EventToolStarted,
|
||||
ToolCallID: "call-1",
|
||||
ToolName: "bash",
|
||||
Args: request.Calls[0].Args,
|
||||
})
|
||||
m.applyAgentEvent(coreagent.Event{
|
||||
Type: coreagent.EventToolFinished,
|
||||
ToolCallID: "call-1",
|
||||
ToolName: "bash",
|
||||
Args: request.Calls[0].Args,
|
||||
Content: "parth-agent-tui\n",
|
||||
})
|
||||
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventMessageDelta, Content: "I'm on branch parth-agent-tui."})
|
||||
|
||||
transcript := stripANSI(m.renderTranscript(180))
|
||||
for _, want := range []string{
|
||||
`Bash("git rev-parse --abbrev-ref HEAD")`,
|
||||
`Bash("git branch -a")`,
|
||||
"I'm on branch parth-agent-tui.",
|
||||
} {
|
||||
if !strings.Contains(transcript, want) {
|
||||
t.Fatalf("transcript missing %q:\n%s", want, transcript)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestChatApprovalPromptRepaintsFlowTranscript(t *testing.T) {
|
||||
reply := make(chan coreagent.Approval, 1)
|
||||
request := coreagent.ApprovalRequest{
|
||||
WorkingDir: "/repo",
|
||||
Calls: []coreagent.ApprovalToolCall{
|
||||
{
|
||||
ToolCallID: "call-1",
|
||||
ToolName: "web_fetch",
|
||||
Args: map[string]any{"url": "https://parthsareen.com/"},
|
||||
ApprovalScope: "web_fetch",
|
||||
},
|
||||
{
|
||||
ToolCallID: "call-2",
|
||||
ToolName: "web_fetch",
|
||||
Args: map[string]any{"url": "https://github.com/ParthSareen"},
|
||||
ApprovalScope: "web_fetch",
|
||||
},
|
||||
},
|
||||
}
|
||||
m := chatModel{
|
||||
running: true,
|
||||
width: 160,
|
||||
flowPrintedLines: 1,
|
||||
entries: []chatEntry{
|
||||
{role: "user", content: "research parth"},
|
||||
},
|
||||
}
|
||||
|
||||
updated, cmd := m.Update(chatApprovalPromptMsg{request: request, reply: reply})
|
||||
if cmd == nil {
|
||||
t.Fatal("opening approval should repaint flow transcript")
|
||||
}
|
||||
fm := updated.(chatModel)
|
||||
transcript := stripANSI(fm.renderTranscript(160))
|
||||
for _, want := range []string{
|
||||
`Web Fetch("https://parthsareen.com/") needs approval`,
|
||||
`Web Fetch("https://github.com/ParthSareen") needs approval`,
|
||||
} {
|
||||
if !strings.Contains(transcript, want) {
|
||||
t.Fatalf("transcript missing %q:\n%s", want, transcript)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestChatApprovalResolutionRepaintsFlowTranscript(t *testing.T) {
|
||||
reply := make(chan coreagent.Approval, 1)
|
||||
request := coreagent.ApprovalRequest{
|
||||
WorkingDir: "/repo",
|
||||
Calls: []coreagent.ApprovalToolCall{
|
||||
{
|
||||
ToolCallID: "call-1",
|
||||
ToolName: "web_fetch",
|
||||
Args: map[string]any{"url": "https://parthsareen.com/"},
|
||||
ApprovalScope: "web_fetch",
|
||||
},
|
||||
{
|
||||
ToolCallID: "call-2",
|
||||
ToolName: "web_fetch",
|
||||
Args: map[string]any{"url": "https://github.com/ParthSareen"},
|
||||
ApprovalScope: "web_fetch",
|
||||
},
|
||||
},
|
||||
}
|
||||
m := chatModel{
|
||||
running: true,
|
||||
width: 160,
|
||||
events: make(chan tea.Msg),
|
||||
entries: []chatEntry{
|
||||
{role: "user", content: "research parth"},
|
||||
},
|
||||
}
|
||||
m.openApprovalPrompt(chatApprovalPromptMsg{request: request, reply: reply})
|
||||
printed := len(m.transcriptLines(160))
|
||||
m.flowPrintedLines = printed
|
||||
|
||||
updated, cmd := m.resolveApprovalPrompt(chatApprovalChoice{allow: true})
|
||||
if cmd == nil {
|
||||
t.Fatal("approval resolution should keep waiting for agent events")
|
||||
}
|
||||
fm := updated.(chatModel)
|
||||
if fm.flowPrintedLines >= printed {
|
||||
t.Fatalf("approval resolution should repaint and hold queued rows, flowPrintedLines = %d, was %d", fm.flowPrintedLines, printed)
|
||||
}
|
||||
if result := <-reply; !result.Allow {
|
||||
t.Fatalf("approval = %#v, want allow", result)
|
||||
}
|
||||
}
|
||||
|
||||
func TestChatApprovalBatchCollapsesAtNextToolBoundary(t *testing.T) {
|
||||
reply := make(chan coreagent.Approval, 1)
|
||||
request := coreagent.ApprovalRequest{
|
||||
WorkingDir: "/repo",
|
||||
Calls: []coreagent.ApprovalToolCall{
|
||||
{
|
||||
ToolCallID: "call-1",
|
||||
ToolName: "bash",
|
||||
Args: map[string]any{"command": "git rev-parse --abbrev-ref HEAD"},
|
||||
ApprovalScope: "bash\x00git rev-parse --abbrev-ref HEAD",
|
||||
},
|
||||
{
|
||||
ToolCallID: "call-2",
|
||||
ToolName: "bash",
|
||||
Args: map[string]any{"command": "git branch -a"},
|
||||
ApprovalScope: "bash\x00git branch -a",
|
||||
},
|
||||
},
|
||||
}
|
||||
m := chatModel{
|
||||
running: true,
|
||||
events: make(chan tea.Msg),
|
||||
}
|
||||
m.openApprovalPrompt(chatApprovalPromptMsg{request: request, reply: reply})
|
||||
|
||||
updated, _ := m.resolveApprovalPrompt(chatApprovalChoice{allow: true})
|
||||
m = updated.(chatModel)
|
||||
for _, call := range request.Calls {
|
||||
m.applyAgentEvent(coreagent.Event{
|
||||
Type: coreagent.EventToolStarted,
|
||||
ToolCallID: call.ToolCallID,
|
||||
ToolName: call.ToolName,
|
||||
Args: call.Args,
|
||||
})
|
||||
m.applyAgentEvent(coreagent.Event{
|
||||
Type: coreagent.EventToolFinished,
|
||||
ToolCallID: call.ToolCallID,
|
||||
ToolName: call.ToolName,
|
||||
Args: call.Args,
|
||||
Content: "ok\n",
|
||||
})
|
||||
}
|
||||
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventMessageDelta, Content: "I'm on branch parth-agent-tui."})
|
||||
|
||||
transcript := stripANSI(m.renderTranscript(180))
|
||||
if strings.Contains(transcript, "Ran 2 commands") {
|
||||
t.Fatalf("completed batch should stay expanded until the next tool boundary:\n%s", transcript)
|
||||
}
|
||||
if !strings.Contains(transcript, `Bash("git branch -a")`) {
|
||||
t.Fatalf("completed batch should keep concrete command rows before the next boundary:\n%s", transcript)
|
||||
}
|
||||
|
||||
m.applyAgentEvent(coreagent.Event{
|
||||
Type: coreagent.EventToolStarted,
|
||||
ToolCallID: "call-3",
|
||||
ToolName: "bash",
|
||||
Args: map[string]any{"command": "git status --short"},
|
||||
})
|
||||
transcript = stripANSI(m.renderTranscript(180))
|
||||
if !strings.Contains(transcript, "Ran 2 commands") {
|
||||
t.Fatalf("completed batch should collapse when a new tool starts:\n%s", transcript)
|
||||
}
|
||||
if !strings.Contains(transcript, `Bash("git status --short")`) {
|
||||
t.Fatalf("new running command should remain concrete after previous batch collapses:\n%s", transcript)
|
||||
}
|
||||
}
|
||||
|
||||
func TestChatApprovalPrompterCancels(t *testing.T) {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
result, err := (chatApprovalPrompter{ch: make(chan tea.Msg)}).PromptApproval(ctx, testApprovalRequest())
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if result.Allow || result.Reason == "" {
|
||||
t.Fatalf("approval = %#v, want canceled denial", result)
|
||||
}
|
||||
}
|
||||
|
||||
func TestChatApprovalControllerAutoApprovesAfterFullAccessToggle(t *testing.T) {
|
||||
events := make(chan tea.Msg, 1)
|
||||
state := testApprovalState(false, nil)
|
||||
controller := newChatApprovalController(events, state)
|
||||
state.GrantAll()
|
||||
|
||||
result, err := controller.PromptApproval(context.Background(), testApprovalRequest())
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !result.Allow || !result.AllowAll {
|
||||
t.Fatalf("approval = %#v, want full-access approval", result)
|
||||
}
|
||||
select {
|
||||
case msg := <-events:
|
||||
t.Fatalf("approval UI event should not be sent after full access toggle: %#v", msg)
|
||||
default:
|
||||
}
|
||||
}
|
||||
|
||||
func TestChatPermissionToggleSyncsRunningApprovalController(t *testing.T) {
|
||||
events := make(chan tea.Msg, 1)
|
||||
state := testApprovalState(false, nil)
|
||||
m := chatModel{
|
||||
approvalState: state,
|
||||
approvalController: newChatApprovalController(events, state),
|
||||
}
|
||||
|
||||
updated, _ := m.togglePermissionMode()
|
||||
fm := updated.(chatModel)
|
||||
result, err := fm.approvalController.PromptApproval(context.Background(), testApprovalRequest())
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !result.Allow || !result.AllowAll {
|
||||
t.Fatalf("approval = %#v, want full-access approval", result)
|
||||
}
|
||||
}
|
||||
|
||||
func TestChatPermissionToggleFromFullAccessRequiresReviewInRunningController(t *testing.T) {
|
||||
events := make(chan tea.Msg, 1)
|
||||
state := testApprovalState(true, nil)
|
||||
m := chatModel{
|
||||
approvalState: state,
|
||||
approvalController: newChatApprovalController(events, state),
|
||||
}
|
||||
|
||||
updated, _ := m.togglePermissionMode()
|
||||
fm := updated.(chatModel)
|
||||
if fm.allowAllToolsEnabled() {
|
||||
t.Fatal("full access should be disabled")
|
||||
}
|
||||
|
||||
resultCh := make(chan coreagent.Approval, 1)
|
||||
go func() {
|
||||
result, err := fm.approvalController.PromptApproval(context.Background(), testApprovalRequest())
|
||||
if err != nil {
|
||||
resultCh <- coreagent.Approval{Reason: err.Error()}
|
||||
return
|
||||
}
|
||||
resultCh <- result
|
||||
}()
|
||||
|
||||
select {
|
||||
case msg := <-events:
|
||||
prompt, ok := msg.(chatApprovalPromptMsg)
|
||||
if !ok {
|
||||
t.Fatalf("event = %#v, want approval prompt", msg)
|
||||
}
|
||||
prompt.reply <- coreagent.Approval{Reason: "denied"}
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("expected approval prompt after toggling from full access to review")
|
||||
}
|
||||
|
||||
result := <-resultCh
|
||||
if result.Allow {
|
||||
t.Fatalf("approval = %#v, want review prompt result", result)
|
||||
}
|
||||
}
|
||||
|
||||
func TestChatApprovalPromptSkippedWhenFullAccessEnabledInFlight(t *testing.T) {
|
||||
reply := make(chan coreagent.Approval, 1)
|
||||
// Full access is on by the time the buffered approval request reaches the
|
||||
// UI (toggled after the agent sent the request but before Update ran).
|
||||
// The stale prompt must not surface; the request is auto-approved.
|
||||
m := chatModel{approvalState: testApprovalState(true, nil), running: true}
|
||||
|
||||
updated, _ := m.Update(chatApprovalPromptMsg{request: testApprovalRequest(), reply: reply})
|
||||
fm := updated.(chatModel)
|
||||
|
||||
if fm.approvalPrompt != nil {
|
||||
t.Fatalf("approval prompt = %#v, want nil (full access on)", fm.approvalPrompt)
|
||||
}
|
||||
if got := fm.status; got == "approval required" {
|
||||
t.Fatalf("status = %q, should not show approval required", got)
|
||||
}
|
||||
select {
|
||||
case result := <-reply:
|
||||
if !result.Allow || !result.AllowAll {
|
||||
t.Fatalf("approval = %#v, want full-access approval", result)
|
||||
}
|
||||
default:
|
||||
t.Fatal("expected auto-approval sent on the reply channel")
|
||||
}
|
||||
}
|
||||
File diff suppressed because it is too large.
Load diff
@@ -0,0 +1,66 @@
|
||||
package chat
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"os/exec"
|
||||
"runtime"
|
||||
"strings"
|
||||
|
||||
tea "github.com/charmbracelet/bubbletea"
|
||||
)
|
||||
|
||||
type chatClipboardErrorMsg struct {
|
||||
err error
|
||||
}
|
||||
|
||||
var writeClipboard = writeSystemClipboard
|
||||
|
||||
func copyTextCmd(ctx context.Context, text string) tea.Cmd {
|
||||
return func() tea.Msg {
|
||||
if err := writeClipboard(ctx, text); err != nil {
|
||||
return chatClipboardErrorMsg{err: err}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
func writeSystemClipboard(ctx context.Context, text string) error {
|
||||
if ctx == nil {
|
||||
ctx = context.Background()
|
||||
}
|
||||
switch runtime.GOOS {
|
||||
case "darwin":
|
||||
return runClipboardCommand(ctx, text, "pbcopy")
|
||||
case "windows":
|
||||
return runClipboardCommand(ctx, text, "clip")
|
||||
default:
|
||||
for _, candidate := range []struct {
|
||||
name string
|
||||
args []string
|
||||
}{
|
||||
{name: "wl-copy"},
|
||||
{name: "xclip", args: []string{"-selection", "clipboard"}},
|
||||
{name: "xsel", args: []string{"--clipboard", "--input"}},
|
||||
} {
|
||||
if _, err := exec.LookPath(candidate.name); err != nil {
|
||||
continue
|
||||
}
|
||||
return runClipboardCommand(ctx, text, candidate.name, candidate.args...)
|
||||
}
|
||||
return errors.New("no clipboard command found")
|
||||
}
|
||||
}
|
||||
|
||||
func runClipboardCommand(ctx context.Context, text, name string, args ...string) error {
|
||||
cmd := exec.CommandContext(ctx, name, args...)
|
||||
cmd.Stdin = strings.NewReader(text)
|
||||
if output, err := cmd.CombinedOutput(); err != nil {
|
||||
if len(output) > 0 {
|
||||
return fmt.Errorf("%s: %w: %s", name, err, strings.TrimSpace(string(output)))
|
||||
}
|
||||
return fmt.Errorf("%s: %w", name, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,434 @@
|
||||
package chat
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
tea "github.com/charmbracelet/bubbletea"
|
||||
"github.com/ollama/ollama/api"
|
||||
"github.com/ollama/ollama/cmd/launch"
|
||||
"github.com/ollama/ollama/internal/modelref"
|
||||
)
|
||||
|
||||
type cloudAuthKind string
|
||||
|
||||
const (
|
||||
cloudAuthSignIn cloudAuthKind = "signin"
|
||||
cloudAuthUpgrade cloudAuthKind = "upgrade"
|
||||
cloudAuthChecking cloudAuthKind = "checking"
|
||||
)
|
||||
|
||||
const cloudPlanVerificationUnavailable = "Could not verify Ollama plan. Try again in a moment or use a local model."
|
||||
|
||||
// Sign-in/upgrade verification polling bounds. While the check is healthy but
|
||||
// the user hasn't signed in yet, polling stays prompt so completion is detected
|
||||
// quickly. When the check itself fails, polling backs off so a down server
|
||||
// isn't hammered, and gives up after maxPollFailures consecutive errors (or
|
||||
// pollHardCap elapsed) so the user isn't stuck on a spinner with no recourse
|
||||
// beyond Esc.
|
||||
const (
|
||||
maxPollFailures = 6
|
||||
pollBackoffBase = 3 * time.Second
|
||||
pollBackoffCap = 30 * time.Second
|
||||
pollHardCap = 2 * time.Minute
|
||||
)
|
||||
|
||||
// cloudAuthPrompt is an inline modal that handles sign-in and plan-upgrade
|
||||
// flows when a user selects a cloud model from the picker.
|
||||
type cloudAuthPrompt struct {
|
||||
modelName string
|
||||
requiredPlan string
|
||||
signInURL string
|
||||
upgradeURL string
|
||||
kind cloudAuthKind
|
||||
spinner int
|
||||
openNow bool
|
||||
polling bool
|
||||
// pollStarted tracks when sign-in/upgrade verification polling began, for
|
||||
// the hard-cap timeout. Lazily set on the first poll response.
|
||||
pollStarted time.Time
|
||||
// pollFailures counts consecutive verification-check errors; once it
|
||||
// reaches maxPollFailures the modal gives up and surfaces an error.
|
||||
pollFailures int
|
||||
// pollErr holds the last verification error, rendered while retrying.
|
||||
pollErr string
|
||||
}
|
||||
|
||||
type cloudAuthCheckMsg struct {
|
||||
err error
|
||||
signInURL string
|
||||
}
|
||||
|
||||
type cloudModelPreflightMsg struct {
|
||||
model string
|
||||
err error
|
||||
signInURL string
|
||||
}
|
||||
|
||||
type cloudAuthTickMsg struct{}
|
||||
|
||||
type cloudAuthPollMsg struct {
|
||||
done bool
|
||||
err error
|
||||
}
|
||||
|
||||
func checkCloudModelCmd(ctx context.Context, check func(context.Context, string, string) error, model, requiredPlan string) tea.Cmd {
|
||||
if check == nil {
|
||||
return nil
|
||||
}
|
||||
return func() tea.Msg {
|
||||
if ctx == nil {
|
||||
ctx = context.Background()
|
||||
}
|
||||
err := check(ctx, model, requiredPlan)
|
||||
var signInURL string
|
||||
if err != nil {
|
||||
var authErr api.AuthorizationError
|
||||
if errors.As(err, &authErr) && authErr.SigninURL != "" {
|
||||
signInURL = authErr.SigninURL
|
||||
}
|
||||
}
|
||||
return cloudAuthCheckMsg{err: err, signInURL: signInURL}
|
||||
}
|
||||
}
|
||||
|
||||
func cloudModelPreflightCmd(ctx context.Context, opts Options, modelName, requiredPlan string) tea.Cmd {
|
||||
modelName = strings.TrimSpace(modelName)
|
||||
if opts.CheckCloudModel == nil || modelName == "" || !modelref.HasExplicitCloudSource(modelName) {
|
||||
return nil
|
||||
}
|
||||
return func() tea.Msg {
|
||||
if ctx == nil {
|
||||
ctx = context.Background()
|
||||
}
|
||||
plan := strings.TrimSpace(requiredPlan)
|
||||
if plan == "" && opts.ModelOptions != nil {
|
||||
models, err := opts.ModelOptions(ctx)
|
||||
if err == nil {
|
||||
for _, model := range models {
|
||||
if strings.EqualFold(strings.TrimSpace(model.Name), modelName) {
|
||||
plan = strings.TrimSpace(model.RequiredPlan)
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
err := opts.CheckCloudModel(ctx, modelName, plan)
|
||||
return cloudModelPreflightMsg{
|
||||
model: modelName,
|
||||
err: err,
|
||||
signInURL: cloudAuthSignInURL(err),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func cloudAuthSignInURL(err error) string {
|
||||
if err == nil {
|
||||
return ""
|
||||
}
|
||||
var authErr api.AuthorizationError
|
||||
if errors.As(err, &authErr) && (authErr.StatusCode == http.StatusUnauthorized || authErr.SigninURL != "") {
|
||||
return authErr.SigninURL
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func cloudAuthTickCmd() tea.Cmd {
|
||||
return tea.Tick(200*time.Millisecond, func(t time.Time) tea.Msg {
|
||||
return cloudAuthTickMsg{}
|
||||
})
|
||||
}
|
||||
|
||||
func (m chatModel) updateCloudModelPreflight(msg cloudModelPreflightMsg) (tea.Model, tea.Cmd) {
|
||||
if msg.model == "" || !strings.EqualFold(strings.TrimSpace(m.opts.Model), strings.TrimSpace(msg.model)) {
|
||||
return m, nil
|
||||
}
|
||||
if msg.err == nil {
|
||||
if m.status == cloudPlanVerificationUnavailable {
|
||||
m.status = "ready"
|
||||
}
|
||||
return m, nil
|
||||
}
|
||||
if msg.signInURL != "" {
|
||||
return m.startCloudAuthSignIn(msg.model, "", msg.signInURL)
|
||||
}
|
||||
m.status = cloudPlanVerificationUnavailable
|
||||
return m, nil
|
||||
}
|
||||
|
||||
func pollCloudAuthCmd(ctx context.Context, poll func(context.Context) (string, bool, error), delay time.Duration) tea.Cmd {
|
||||
if poll == nil {
|
||||
return nil
|
||||
}
|
||||
return func() tea.Msg {
|
||||
if ctx == nil {
|
||||
ctx = context.Background()
|
||||
}
|
||||
// Back off before the next check when the previous one failed. Honor
|
||||
// context cancellation so an abandoned modal doesn't block on the
|
||||
// full delay.
|
||||
if delay > 0 {
|
||||
timer := time.NewTimer(delay)
|
||||
defer timer.Stop()
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
case <-timer.C:
|
||||
}
|
||||
}
|
||||
pollCtx, cancel := context.WithTimeout(ctx, 3*time.Second)
|
||||
defer cancel()
|
||||
_, done, err := poll(pollCtx)
|
||||
return cloudAuthPollMsg{done: done, err: err}
|
||||
}
|
||||
}
|
||||
|
||||
func (m *chatModel) startCloudAuthSignIn(modelName, requiredPlan, signInURL string) (tea.Model, tea.Cmd) {
|
||||
// When no sign-in URL is available yet, show the "checking" state while
|
||||
// we verify the plan, rather than rendering a blank "Navigate to:" URL.
|
||||
kind := cloudAuthSignIn
|
||||
if signInURL == "" {
|
||||
kind = cloudAuthChecking
|
||||
}
|
||||
m.cloudAuthPrompt = &cloudAuthPrompt{
|
||||
modelName: modelName,
|
||||
requiredPlan: requiredPlan,
|
||||
kind: kind,
|
||||
signInURL: signInURL,
|
||||
polling: true,
|
||||
}
|
||||
m.status = "cloud-auth"
|
||||
m.modelPicker = nil
|
||||
m.modelPickerModels = nil
|
||||
if m.opts.OpenBrowser != nil && signInURL != "" {
|
||||
m.opts.OpenBrowser(signInURL)
|
||||
}
|
||||
if signInURL == "" {
|
||||
return m, checkCloudModelCmd(m.ctx, m.opts.CheckCloudModel, modelName, requiredPlan)
|
||||
}
|
||||
return m, tea.Batch(cloudAuthTickCmd(), pollCloudAuthCmd(m.ctx, m.opts.PollCloudAuth, 0))
|
||||
}
|
||||
|
||||
func (m *chatModel) startCloudAuthUpgrade(modelName, requiredPlan string) (tea.Model, tea.Cmd) {
|
||||
m.cloudAuthPrompt = &cloudAuthPrompt{
|
||||
modelName: modelName,
|
||||
requiredPlan: requiredPlan,
|
||||
kind: cloudAuthUpgrade,
|
||||
upgradeURL: launch.DefaultUpgradeURL,
|
||||
openNow: true,
|
||||
}
|
||||
m.status = "cloud-auth"
|
||||
m.modelPicker = nil
|
||||
m.modelPickerModels = nil
|
||||
return m, nil
|
||||
}
|
||||
|
||||
func (m chatModel) updateCloudAuthPrompt(msg tea.Msg) (tea.Model, tea.Cmd) {
|
||||
switch msg := msg.(type) {
|
||||
case cloudAuthCheckMsg:
|
||||
if msg.err == nil {
|
||||
// Auth passed — apply the pending model.
|
||||
return m.completeCloudAuth()
|
||||
}
|
||||
// Determine if sign-in or upgrade is needed.
|
||||
if msg.signInURL != "" {
|
||||
m.cloudAuthPrompt.kind = cloudAuthSignIn
|
||||
m.cloudAuthPrompt.signInURL = msg.signInURL
|
||||
m.cloudAuthPrompt.polling = true
|
||||
if m.opts.OpenBrowser != nil {
|
||||
m.opts.OpenBrowser(msg.signInURL)
|
||||
}
|
||||
return m, tea.Batch(cloudAuthTickCmd(), pollCloudAuthCmd(m.ctx, m.opts.PollCloudAuth, 0))
|
||||
}
|
||||
// Could be a plan upgrade error or unknown error.
|
||||
m.cloudAuthPrompt = nil
|
||||
m.openModelOnInit = false
|
||||
m.status = "ready"
|
||||
m.entries = append(m.entries, newChatEntry(chatEntry{role: "error", content: fmt.Sprintf("Could not switch model: %v", msg.err), err: msg.err.Error()}))
|
||||
return m, nil
|
||||
|
||||
case cloudAuthTickMsg:
|
||||
if m.cloudAuthPrompt == nil {
|
||||
return m, nil
|
||||
}
|
||||
m.cloudAuthPrompt.spinner++
|
||||
return m, cloudAuthTickCmd()
|
||||
|
||||
case cloudAuthPollMsg:
|
||||
if m.cloudAuthPrompt == nil {
|
||||
return m, nil
|
||||
}
|
||||
if msg.done {
|
||||
// Signed in — re-check auth to see if plan is satisfied.
|
||||
m.cloudAuthPrompt.polling = false
|
||||
m.cloudAuthPrompt.pollFailures = 0
|
||||
m.cloudAuthPrompt.pollErr = ""
|
||||
return m, checkCloudModelCmd(m.ctx, m.opts.CheckCloudModel, m.cloudAuthPrompt.modelName, m.cloudAuthPrompt.requiredPlan)
|
||||
}
|
||||
// Lazily mark the start of the polling window on the first response.
|
||||
if m.cloudAuthPrompt.pollStarted.IsZero() {
|
||||
m.cloudAuthPrompt.pollStarted = time.Now()
|
||||
}
|
||||
// Hard cap: give up if verification drags on too long for any reason.
|
||||
if time.Since(m.cloudAuthPrompt.pollStarted) > pollHardCap {
|
||||
return m.failCloudAuthPoll(errors.New("sign-in is taking longer than expected; check your connection and try again"))
|
||||
}
|
||||
if msg.err != nil {
|
||||
// The verification check itself failed (network down, server 5xx).
|
||||
// Back off and retry, but give up after a handful of consecutive
|
||||
// failures so the user isn't stuck on a spinner with no signal.
|
||||
m.cloudAuthPrompt.pollFailures++
|
||||
m.cloudAuthPrompt.pollErr = msg.err.Error()
|
||||
if m.cloudAuthPrompt.pollFailures >= maxPollFailures {
|
||||
return m.failCloudAuthPoll(fmt.Errorf("couldn't verify sign-in: %w", msg.err))
|
||||
}
|
||||
delay := pollBackoffCap
|
||||
if d := pollBackoffBase << (m.cloudAuthPrompt.pollFailures - 1); d < pollBackoffCap {
|
||||
delay = d
|
||||
}
|
||||
return m, pollCloudAuthCmd(m.ctx, m.opts.PollCloudAuth, delay)
|
||||
}
|
||||
// Healthy but not signed in yet — keep polling promptly so sign-in
|
||||
// completion is detected without added latency.
|
||||
m.cloudAuthPrompt.pollFailures = 0
|
||||
m.cloudAuthPrompt.pollErr = ""
|
||||
return m, pollCloudAuthCmd(m.ctx, m.opts.PollCloudAuth, 0)
|
||||
|
||||
case tea.KeyMsg:
|
||||
if msg.Type == tea.KeyEsc || msg.Type == tea.KeyCtrlC {
|
||||
m.cloudAuthPrompt = nil
|
||||
m.pendingModel = ""
|
||||
m.openModelOnInit = false
|
||||
m.status = "ready"
|
||||
return m, nil
|
||||
}
|
||||
if m.cloudAuthPrompt.kind == cloudAuthUpgrade && !m.cloudAuthPrompt.polling {
|
||||
switch msg.Type {
|
||||
case tea.KeyLeft, tea.KeyRight, tea.KeyTab:
|
||||
m.cloudAuthPrompt.openNow = !m.cloudAuthPrompt.openNow
|
||||
case tea.KeyEnter:
|
||||
if m.cloudAuthPrompt.openNow {
|
||||
m.cloudAuthPrompt.polling = true
|
||||
if m.opts.OpenBrowser != nil && m.cloudAuthPrompt.upgradeURL != "" {
|
||||
m.opts.OpenBrowser(m.cloudAuthPrompt.upgradeURL)
|
||||
}
|
||||
return m, tea.Batch(cloudAuthTickCmd(), pollCloudAuthCmd(m.ctx, m.opts.PollCloudAuth, 0))
|
||||
}
|
||||
m.cloudAuthPrompt = nil
|
||||
m.pendingModel = ""
|
||||
m.openModelOnInit = false
|
||||
m.status = "ready"
|
||||
return m, nil
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return m, nil
|
||||
}
|
||||
|
||||
// failCloudAuthPoll abandons the sign-in/upgrade verification modal, surfaces
|
||||
// an error entry to the user, and returns to the ready state so they can
|
||||
// re-pick a model and retry.
|
||||
func (m chatModel) failCloudAuthPoll(err error) (tea.Model, tea.Cmd) {
|
||||
m.cloudAuthPrompt = nil
|
||||
m.pendingModel = ""
|
||||
m.openModelOnInit = false
|
||||
m.status = "ready"
|
||||
m.entries = append(m.entries, newChatEntry(chatEntry{role: "error", content: fmt.Sprintf("Could not switch model: %v", err), err: err.Error()}))
|
||||
return m, nil
|
||||
}
|
||||
|
||||
func (m chatModel) completeCloudAuth() (tea.Model, tea.Cmd) {
|
||||
pending := m.cloudAuthPrompt.modelName
|
||||
m.cloudAuthPrompt = nil
|
||||
m.pendingModel = ""
|
||||
m.modelPicker = nil
|
||||
m.modelPickerModels = nil
|
||||
m.openModelOnInit = false
|
||||
m.status = "ready"
|
||||
if err := m.applyModelSelection(pending, true); err != nil {
|
||||
m.entries = append(m.entries, newChatEntry(chatEntry{role: "error", content: fmt.Sprintf("Could not switch model: %v", err), err: err.Error()}))
|
||||
m.status = "error"
|
||||
return m, nil
|
||||
}
|
||||
return m, m.startModelPreload(pending)
|
||||
}
|
||||
|
||||
func (m chatModel) renderCloudAuthPrompt(width int) string {
|
||||
if m.cloudAuthPrompt == nil {
|
||||
return ""
|
||||
}
|
||||
if width <= 0 {
|
||||
width = 80
|
||||
}
|
||||
|
||||
p := m.cloudAuthPrompt
|
||||
spinnerFrames := []string{"⠋", "⠙", "⠹", "⠸", "⠼", "⠴", "⠦", "⠧", "⠇", "⠏"}
|
||||
frame := spinnerFrames[p.spinner%len(spinnerFrames)]
|
||||
|
||||
var b strings.Builder
|
||||
|
||||
switch p.kind {
|
||||
case cloudAuthChecking:
|
||||
fmt.Fprintf(&b, "%s Checking %s...\n\n", frame, chatPickerSelectedStyle.Render(p.modelName))
|
||||
b.WriteString(chatPickerMetaStyle.Render("esc cancel"))
|
||||
case cloudAuthSignIn:
|
||||
fmt.Fprintf(&b, "To use %s, please sign in.\n\n", chatPickerSelectedStyle.Render(p.modelName))
|
||||
b.WriteString("Navigate to:\n")
|
||||
urlWrap := chatPickerTextStyle
|
||||
if width > 4 {
|
||||
urlWrap = chatPickerTextStyle.Width(width - 4)
|
||||
}
|
||||
b.WriteString(urlWrap.Render(p.signInURL))
|
||||
b.WriteString("\n\n")
|
||||
if p.pollErr != "" {
|
||||
b.WriteString(chatPickerMetaStyle.Render(frame + " Couldn't verify sign-in: " + p.pollErr + " — retrying..."))
|
||||
} else {
|
||||
b.WriteString(chatPickerMetaStyle.Render(frame + " Waiting for sign in to complete..."))
|
||||
}
|
||||
b.WriteString("\n\n")
|
||||
b.WriteString(chatPickerMetaStyle.Render("esc cancel"))
|
||||
case cloudAuthUpgrade:
|
||||
fmt.Fprintf(&b, "To use %s, upgrade your Ollama plan.\n\n", chatPickerSelectedStyle.Render(p.modelName))
|
||||
if !p.polling {
|
||||
var yesBtn, noBtn string
|
||||
if p.openNow {
|
||||
yesBtn = chatPickerSelectedStyle.Render("› Yes ")
|
||||
noBtn = chatPickerMetaStyle.Render(" No ")
|
||||
} else {
|
||||
yesBtn = chatPickerMetaStyle.Render(" Yes ")
|
||||
noBtn = chatPickerSelectedStyle.Render("› No ")
|
||||
}
|
||||
b.WriteString("Open upgrade page now?\n")
|
||||
b.WriteString(yesBtn + " " + noBtn)
|
||||
b.WriteString("\n\n")
|
||||
if !p.openNow {
|
||||
b.WriteString("Or navigate to:\n")
|
||||
urlWrap := chatPickerTextStyle
|
||||
if width > 4 {
|
||||
urlWrap = chatPickerTextStyle.Width(width - 4)
|
||||
}
|
||||
if u := p.upgradeURL; u != "" {
|
||||
b.WriteString(urlWrap.Render(u))
|
||||
} else {
|
||||
b.WriteString(urlWrap.Render(launch.DefaultUpgradeURL))
|
||||
}
|
||||
b.WriteString("\n\n")
|
||||
}
|
||||
b.WriteString(chatPickerMetaStyle.Render("←/→ navigate • enter confirm • esc cancel"))
|
||||
} else {
|
||||
if p.pollErr != "" {
|
||||
b.WriteString(chatPickerMetaStyle.Render(frame + " Couldn't verify upgrade: " + p.pollErr + " — retrying..."))
|
||||
} else {
|
||||
b.WriteString(chatPickerMetaStyle.Render(frame + " Waiting for upgrade to complete..."))
|
||||
}
|
||||
b.WriteString("\n\n")
|
||||
b.WriteString(chatPickerMetaStyle.Render("esc cancel"))
|
||||
}
|
||||
}
|
||||
|
||||
return b.String()
|
||||
}
|
||||
@@ -0,0 +1,247 @@
|
||||
package chat
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestCloudAuthTickDoesNotPoll(t *testing.T) {
|
||||
polls := 0
|
||||
m := chatModel{
|
||||
cloudAuthPrompt: &cloudAuthPrompt{polling: true},
|
||||
opts: Options{
|
||||
PollCloudAuth: func(context.Context) (string, bool, error) {
|
||||
polls++
|
||||
return "", false, nil
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
updated, cmd := m.updateCloudAuthPrompt(cloudAuthTickMsg{})
|
||||
m = updated.(chatModel)
|
||||
|
||||
if m.cloudAuthPrompt.spinner != 1 {
|
||||
t.Fatalf("spinner = %d, want 1", m.cloudAuthPrompt.spinner)
|
||||
}
|
||||
if polls != 0 {
|
||||
t.Fatalf("polls = %d, want 0 before running returned tick command", polls)
|
||||
}
|
||||
if cmd == nil {
|
||||
t.Fatal("tick should schedule the next tick")
|
||||
}
|
||||
if _, ok := cmd().(cloudAuthTickMsg); !ok {
|
||||
t.Fatal("tick should schedule another tick, not a poll")
|
||||
}
|
||||
if polls != 0 {
|
||||
t.Fatalf("polls = %d, want 0 after running returned tick command", polls)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCloudAuthPollSchedulesNextPoll(t *testing.T) {
|
||||
polls := 0
|
||||
m := chatModel{
|
||||
cloudAuthPrompt: &cloudAuthPrompt{polling: true},
|
||||
opts: Options{
|
||||
PollCloudAuth: func(context.Context) (string, bool, error) {
|
||||
polls++
|
||||
return "", false, nil
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
_, cmd := m.updateCloudAuthPrompt(cloudAuthPollMsg{})
|
||||
if cmd == nil {
|
||||
t.Fatal("poll should schedule the next poll")
|
||||
}
|
||||
msg, ok := cmd().(cloudAuthPollMsg)
|
||||
if !ok {
|
||||
t.Fatal("poll should schedule another poll, not a tick")
|
||||
}
|
||||
if msg.done {
|
||||
t.Fatal("poll should report not done")
|
||||
}
|
||||
if polls != 1 {
|
||||
t.Fatalf("polls = %d, want 1", polls)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCloudModelPreflightFailureShowsPlanVerificationNotice(t *testing.T) {
|
||||
m := chatModel{
|
||||
opts: Options{
|
||||
Model: "glm-5.2:cloud",
|
||||
},
|
||||
}
|
||||
|
||||
updated, cmd := m.updateCloudModelPreflight(cloudModelPreflightMsg{
|
||||
model: "glm-5.2:cloud",
|
||||
err: errors.New("temporary network failure"),
|
||||
})
|
||||
if cmd != nil {
|
||||
t.Fatal("transient preflight failure should not start an auth modal")
|
||||
}
|
||||
m = updated.(chatModel)
|
||||
|
||||
if got := m.status; got != cloudPlanVerificationUnavailable {
|
||||
t.Fatalf("status = %q", got)
|
||||
}
|
||||
if m.cloudAuthPrompt != nil {
|
||||
t.Fatalf("cloud auth prompt = %#v, want nil", m.cloudAuthPrompt)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCloudModelPreflightIgnoresStaleModel(t *testing.T) {
|
||||
m := chatModel{
|
||||
opts: Options{
|
||||
Model: "glm-5.2:cloud",
|
||||
},
|
||||
status: "ready",
|
||||
}
|
||||
|
||||
updated, _ := m.updateCloudModelPreflight(cloudModelPreflightMsg{
|
||||
model: "kimi-k2.7-code:cloud",
|
||||
err: errors.New("temporary network failure"),
|
||||
})
|
||||
m = updated.(chatModel)
|
||||
|
||||
if got := m.status; got != "ready" {
|
||||
t.Fatalf("status = %q, want unchanged", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCloudModelPreflightCommandChecksCloudModel(t *testing.T) {
|
||||
var checkedModel, checkedPlan string
|
||||
cmd := cloudModelPreflightCmd(context.Background(), Options{
|
||||
CheckCloudModel: func(_ context.Context, model, requiredPlan string) error {
|
||||
checkedModel = model
|
||||
checkedPlan = requiredPlan
|
||||
return errors.New("temporary network failure")
|
||||
},
|
||||
ModelOptions: func(context.Context) ([]ModelOption, error) {
|
||||
return []ModelOption{{Name: "glm-5.2:cloud", RequiredPlan: "pro", Cloud: true}}, nil
|
||||
},
|
||||
}, "glm-5.2:cloud", "")
|
||||
if cmd == nil {
|
||||
t.Fatal("cloud preflight command should be scheduled")
|
||||
}
|
||||
raw := cmd()
|
||||
msg, ok := raw.(cloudModelPreflightMsg)
|
||||
if !ok {
|
||||
t.Fatalf("message = %T, want cloudModelPreflightMsg", raw)
|
||||
}
|
||||
if checkedModel != "glm-5.2:cloud" || checkedPlan != "pro" {
|
||||
t.Fatalf("checked model/plan = %q/%q", checkedModel, checkedPlan)
|
||||
}
|
||||
if msg.model != "glm-5.2:cloud" || msg.err == nil || !strings.Contains(msg.err.Error(), "temporary") {
|
||||
t.Fatalf("message = %#v", msg)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCloudAuthPollGivesUpAfterConsecutiveFailures(t *testing.T) {
|
||||
pollErr := errors.New("whoami: connection refused")
|
||||
m := chatModel{
|
||||
cloudAuthPrompt: &cloudAuthPrompt{polling: true, kind: cloudAuthSignIn},
|
||||
opts: Options{
|
||||
PollCloudAuth: func(context.Context) (string, bool, error) {
|
||||
return "", false, pollErr
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
// The first maxPollFailures-1 failures should keep retrying.
|
||||
for i := 1; i < maxPollFailures; i++ {
|
||||
updated, _ := m.updateCloudAuthPrompt(cloudAuthPollMsg{done: false, err: pollErr})
|
||||
m = updated.(chatModel)
|
||||
if m.cloudAuthPrompt == nil {
|
||||
t.Fatalf("failure %d: prompt cleared early", i)
|
||||
}
|
||||
if got := m.cloudAuthPrompt.pollFailures; got != i {
|
||||
t.Fatalf("failure %d: pollFailures = %d, want %d", i, got, i)
|
||||
}
|
||||
if m.cloudAuthPrompt.pollErr != pollErr.Error() {
|
||||
t.Fatalf("failure %d: pollErr = %q, want %q", i, m.cloudAuthPrompt.pollErr, pollErr.Error())
|
||||
}
|
||||
}
|
||||
|
||||
// The threshold failure gives up: prompt cleared, back to ready, error entry.
|
||||
updated, cmd := m.updateCloudAuthPrompt(cloudAuthPollMsg{done: false, err: pollErr})
|
||||
m = updated.(chatModel)
|
||||
if cmd != nil {
|
||||
t.Fatalf("threshold failure should not reschedule, got cmd %T", cmd)
|
||||
}
|
||||
if m.cloudAuthPrompt != nil {
|
||||
t.Fatalf("prompt = %#v, want nil after give-up", m.cloudAuthPrompt)
|
||||
}
|
||||
if m.status != "ready" {
|
||||
t.Fatalf("status = %q, want ready", m.status)
|
||||
}
|
||||
if len(m.entries) == 0 {
|
||||
t.Fatal("expected an error entry after give-up")
|
||||
}
|
||||
last := m.entries[len(m.entries)-1]
|
||||
if last.role != "error" || !strings.Contains(last.content, "couldn't verify sign-in") {
|
||||
t.Fatalf("last entry = %+v, want error containing sign-in failure", last)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCloudAuthPollResetsFailuresOnHealthyResponse(t *testing.T) {
|
||||
pollErr := errors.New("whoami: timeout")
|
||||
m := chatModel{
|
||||
cloudAuthPrompt: &cloudAuthPrompt{polling: true, kind: cloudAuthSignIn},
|
||||
opts: Options{
|
||||
PollCloudAuth: func(context.Context) (string, bool, error) {
|
||||
return "", false, pollErr
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
// Accumulate some failures without hitting the threshold.
|
||||
for range maxPollFailures - 2 {
|
||||
updated, _ := m.updateCloudAuthPrompt(cloudAuthPollMsg{done: false, err: pollErr})
|
||||
m = updated.(chatModel)
|
||||
}
|
||||
if got := m.cloudAuthPrompt.pollFailures; got != maxPollFailures-2 {
|
||||
t.Fatalf("pollFailures = %d, want %d", got, maxPollFailures-2)
|
||||
}
|
||||
|
||||
// A healthy (no-error, not-done) response resets the streak so a later
|
||||
// transient blip isn't counted against a recovered connection.
|
||||
updated, _ := m.updateCloudAuthPrompt(cloudAuthPollMsg{done: false, err: nil})
|
||||
m = updated.(chatModel)
|
||||
if m.cloudAuthPrompt == nil {
|
||||
t.Fatal("healthy response should keep the prompt open")
|
||||
}
|
||||
if got := m.cloudAuthPrompt.pollFailures; got != 0 {
|
||||
t.Fatalf("pollFailures = %d, want 0 after healthy response", got)
|
||||
}
|
||||
if m.cloudAuthPrompt.pollErr != "" {
|
||||
t.Fatalf("pollErr = %q, want empty after healthy response", m.cloudAuthPrompt.pollErr)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCloudAuthPollCompletesAfterFailures(t *testing.T) {
|
||||
pollErr := errors.New("whoami: timeout")
|
||||
m := chatModel{
|
||||
cloudAuthPrompt: &cloudAuthPrompt{
|
||||
modelName: "glm-5.2:cloud",
|
||||
polling: true,
|
||||
kind: cloudAuthSignIn,
|
||||
pollFailures: maxPollFailures - 1,
|
||||
},
|
||||
opts: Options{
|
||||
CheckCloudModel: func(context.Context, string, string) error { return nil },
|
||||
PollCloudAuth: func(context.Context) (string, bool, error) { return "", false, pollErr },
|
||||
},
|
||||
}
|
||||
|
||||
// A successful sign-in mid-retry should clear the failure state and re-check.
|
||||
updated, _ := m.updateCloudAuthPrompt(cloudAuthPollMsg{done: true})
|
||||
m = updated.(chatModel)
|
||||
if m.cloudAuthPrompt.polling {
|
||||
t.Fatal("done should stop polling")
|
||||
}
|
||||
if m.cloudAuthPrompt.pollFailures != 0 || m.cloudAuthPrompt.pollErr != "" {
|
||||
t.Fatalf("failure state not reset: failures=%d err=%q", m.cloudAuthPrompt.pollFailures, m.cloudAuthPrompt.pollErr)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,101 @@
|
||||
package chat
|
||||
|
||||
import (
|
||||
"context"
|
||||
"slices"
|
||||
|
||||
tea "github.com/charmbracelet/bubbletea"
|
||||
|
||||
coreagent "github.com/ollama/ollama/agent"
|
||||
"github.com/ollama/ollama/api"
|
||||
)
|
||||
|
||||
func (m *chatModel) startManualCompaction() (tea.Model, tea.Cmd) {
|
||||
if m.running || m.compacting {
|
||||
m.entries = append(m.entries, newChatEntry(chatEntry{role: "system", content: "Wait for the current response to finish before compacting."}))
|
||||
return *m, nil
|
||||
}
|
||||
m.refreshContextWindowTokens(m.opts.Model)
|
||||
if m.opts.Compactor == nil {
|
||||
m.entries = append(m.entries, newChatEntry(chatEntry{role: "system", content: coreagent.CompactionSkippedMessage("compaction is unavailable")}))
|
||||
m.status = "compact skipped"
|
||||
return *m, nil
|
||||
}
|
||||
|
||||
ctx := m.ctx
|
||||
if ctx == nil {
|
||||
ctx = context.Background()
|
||||
}
|
||||
runCtx, cancel := context.WithCancel(ctx)
|
||||
compactor := m.opts.Compactor
|
||||
events := make(chan tea.Msg, 128)
|
||||
m.compacting = true
|
||||
m.compactingTokens = 0
|
||||
m.cancel = cancel
|
||||
m.compactEvents = events
|
||||
m.status = "compacting"
|
||||
messages := slices.Clone(m.messages)
|
||||
var tools api.Tools
|
||||
if m.opts.Tools != nil {
|
||||
tools = m.opts.Tools.Tools()
|
||||
}
|
||||
req := coreagent.CompactionRequest{
|
||||
ChatID: m.chatID,
|
||||
Model: m.opts.Model,
|
||||
SystemPrompt: m.systemPrompt(""),
|
||||
Messages: messages,
|
||||
Tools: tools,
|
||||
Format: m.opts.Format,
|
||||
Options: m.opts.Options,
|
||||
KeepAlive: m.opts.KeepAlive,
|
||||
Force: true,
|
||||
Progress: func(progress coreagent.CompactionProgress) {
|
||||
select {
|
||||
case events <- chatCompactProgressMsg{tokens: progress.Tokens}:
|
||||
case <-runCtx.Done():
|
||||
}
|
||||
},
|
||||
}
|
||||
go func() {
|
||||
defer close(events)
|
||||
result, err := compactor.MaybeCompact(runCtx, req)
|
||||
select {
|
||||
case events <- chatCompactDoneMsg{result: result, err: err}:
|
||||
case <-runCtx.Done():
|
||||
}
|
||||
}()
|
||||
tickCmd := m.scheduleTick()
|
||||
return *m, tea.Batch(waitForChatMsg(events), tickCmd)
|
||||
}
|
||||
|
||||
func (m chatModel) finishManualCompaction(msg chatCompactDoneMsg) (tea.Model, tea.Cmd) {
|
||||
wasCanceling := m.status == "canceling"
|
||||
m.compacting = false
|
||||
m.compactEvents = nil
|
||||
m.cancel = nil
|
||||
m.compactingTokens = 0
|
||||
if wasCanceling || isChatContextCanceledError(msg.err) {
|
||||
m.status = "compact canceled"
|
||||
return m.withFlowTranscriptFlush(nil)
|
||||
}
|
||||
if msg.err != nil {
|
||||
m.entries = append(m.entries, newChatEntry(chatEntry{role: "system", content: coreagent.CompactionSkippedMessage(msg.err.Error())}))
|
||||
m.status = "compact skipped"
|
||||
return m.withFlowTranscriptFlush(nil)
|
||||
}
|
||||
if !msg.result.Compacted {
|
||||
m.entries = append(m.entries, newChatEntry(chatEntry{role: "system", content: coreagent.CompactionSkippedMessage(msg.result.Reason)}))
|
||||
m.status = "compact skipped"
|
||||
return m.withFlowTranscriptFlush(nil)
|
||||
}
|
||||
|
||||
m.messages = msg.result.Messages
|
||||
m.liveMessages = nil
|
||||
m.entries = entriesFromMessages(m.messages)
|
||||
m.contextTokens = m.estimatePromptTokens(m.messages, "")
|
||||
m.contextEstimate = true
|
||||
m.scroll = 0
|
||||
m.flowPrintedLines = 0
|
||||
m.status = "compacted"
|
||||
return m.withFlowTranscriptFlush(nil)
|
||||
}
|
||||
@@ -0,0 +1,558 @@
|
||||
package chat
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"slices"
|
||||
"strings"
|
||||
|
||||
tea "github.com/charmbracelet/bubbletea"
|
||||
"github.com/charmbracelet/lipgloss"
|
||||
|
||||
coreagent "github.com/ollama/ollama/agent"
|
||||
"github.com/ollama/ollama/api"
|
||||
)
|
||||
|
||||
type chatPromptDebug struct {
|
||||
request api.ChatRequest
|
||||
tokens int
|
||||
scroll int
|
||||
}
|
||||
|
||||
const maxPromptDebugToolResultRunes = 400
|
||||
|
||||
func (m *chatModel) handleSaveCommand(args string) (tea.Model, tea.Cmd) {
|
||||
filename, err := saveRequestFilename(args)
|
||||
if err != nil {
|
||||
return m.addDebugError(err)
|
||||
}
|
||||
raw, err := m.rawRequestJSON()
|
||||
if err != nil {
|
||||
return m.addDebugError(err)
|
||||
}
|
||||
|
||||
dir, err := m.debugWorkingDir()
|
||||
if err != nil {
|
||||
return m.addDebugError(err)
|
||||
}
|
||||
path := filepath.Join(dir, filename)
|
||||
if err := os.WriteFile(path, []byte(raw+"\n"), 0o644); err != nil {
|
||||
return m.addDebugError(err)
|
||||
}
|
||||
m.entries = append(m.entries, newSlashEntry(fmt.Sprintf("saved as %s", filename)))
|
||||
m.status = "saved"
|
||||
return *m, nil
|
||||
}
|
||||
|
||||
func (m *chatModel) handlePromptCommand(args string) (tea.Model, tea.Cmd) {
|
||||
if strings.TrimSpace(args) != "" {
|
||||
return m.addDebugError(fmt.Errorf("usage: /prompt"))
|
||||
}
|
||||
req, tokens := m.requestPreview()
|
||||
m.promptDebug = &chatPromptDebug{
|
||||
request: req,
|
||||
tokens: tokens,
|
||||
}
|
||||
m.flowPrintedLines = 0
|
||||
m.selection = chatSelection{}
|
||||
m.status = "prompt"
|
||||
return *m, tea.Batch(tea.ClearScreen, tea.EnableMouseCellMotion)
|
||||
}
|
||||
|
||||
func (m chatModel) updatePromptDebug(msg tea.KeyMsg) (tea.Model, tea.Cmd) {
|
||||
if m.promptDebug == nil {
|
||||
return m, nil
|
||||
}
|
||||
switch msg.Type {
|
||||
case tea.KeyEsc, tea.KeyCtrlC, tea.KeyEnter:
|
||||
return m.closePromptDebug()
|
||||
case tea.KeyUp, tea.KeyCtrlP:
|
||||
m.promptDebug.scroll--
|
||||
case tea.KeyDown, tea.KeyCtrlN:
|
||||
m.promptDebug.scroll++
|
||||
case tea.KeyPgUp:
|
||||
m.promptDebug.scroll -= max(1, m.promptDebugPageSize())
|
||||
case tea.KeyPgDown:
|
||||
m.promptDebug.scroll += max(1, m.promptDebugPageSize())
|
||||
case tea.KeyHome, tea.KeyCtrlHome:
|
||||
m.promptDebug.scroll = 0
|
||||
case tea.KeyEnd, tea.KeyCtrlEnd:
|
||||
m.promptDebug.scroll = m.promptDebugMaxScroll()
|
||||
}
|
||||
if m.promptDebug != nil {
|
||||
m.promptDebug.scroll = clamp(m.promptDebug.scroll, 0, m.promptDebugMaxScroll())
|
||||
}
|
||||
return m, nil
|
||||
}
|
||||
|
||||
func (m chatModel) closePromptDebug() (tea.Model, tea.Cmd) {
|
||||
m.promptDebug = nil
|
||||
m.status = "ready"
|
||||
m.flowPrintedLines = 0
|
||||
next, printCmd := m.flowTranscriptFlushCmd()
|
||||
return next, tea.Sequence(tea.DisableMouse, tea.ClearScreen, printCmd)
|
||||
}
|
||||
|
||||
func (m chatModel) renderPromptDebug(width, height int) string {
|
||||
if width <= 0 {
|
||||
width = 80
|
||||
}
|
||||
if height <= 0 {
|
||||
height = 24
|
||||
}
|
||||
if m.promptDebug == nil {
|
||||
return renderFullFrame("", width, height)
|
||||
}
|
||||
|
||||
header := []string{
|
||||
chatPickerTitleStyle.Render("Prompt"),
|
||||
chatPickerMetaStyle.Render("full request preview • /save <filename> saved as <filename>.json"),
|
||||
"",
|
||||
}
|
||||
footer := chatPickerMetaStyle.Render("↑/↓ scroll • pgup/pgdn page • enter/esc close")
|
||||
bodyHeight := max(0, height-len(header)-1)
|
||||
body := m.promptDebugLines(width)
|
||||
maxScroll := max(0, len(body)-bodyHeight)
|
||||
scroll := clamp(m.promptDebug.scroll, 0, maxScroll)
|
||||
if bodyHeight < len(body) {
|
||||
body = body[scroll:min(len(body), scroll+bodyHeight)]
|
||||
}
|
||||
|
||||
lines := slices.Clone(header)
|
||||
lines = append(lines, body...)
|
||||
for len(lines) < height-1 {
|
||||
lines = append(lines, "")
|
||||
}
|
||||
lines = append(lines, footer)
|
||||
return renderFrameLines(lines, width, height)
|
||||
}
|
||||
|
||||
func (m chatModel) promptDebugPageSize() int {
|
||||
height := m.height
|
||||
if height <= 0 {
|
||||
height = 24
|
||||
}
|
||||
return max(1, height-5)
|
||||
}
|
||||
|
||||
func (m chatModel) promptDebugMaxScroll() int {
|
||||
if m.promptDebug == nil {
|
||||
return 0
|
||||
}
|
||||
width := m.viewWidth()
|
||||
height := m.height
|
||||
if height <= 0 {
|
||||
height = 24
|
||||
}
|
||||
bodyHeight := max(0, height-4)
|
||||
return max(0, len(m.promptDebugLines(width))-bodyHeight)
|
||||
}
|
||||
|
||||
func (m *chatModel) addDebugError(err error) (tea.Model, tea.Cmd) {
|
||||
m.entries = append(m.entries, newChatEntry(chatEntry{role: "error", content: err.Error(), err: err.Error()}))
|
||||
m.status = "error"
|
||||
return *m, nil
|
||||
}
|
||||
|
||||
func (m chatModel) rawRequestJSON() (string, error) {
|
||||
req, _ := m.requestPreview()
|
||||
data, err := json.MarshalIndent(req, "", " ")
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return string(data), nil
|
||||
}
|
||||
|
||||
func (m chatModel) requestPreview() (api.ChatRequest, int) {
|
||||
opts := m.previewRunOptions()
|
||||
messages := m.previewMessages()
|
||||
req := m.previewChatRequest(opts, messages)
|
||||
return req, m.estimatePromptTokens(messages, opts.SystemPrompt)
|
||||
}
|
||||
|
||||
func (m chatModel) previewRunOptions() coreagent.RunOptions {
|
||||
return coreagent.RunOptions{
|
||||
ChatID: m.chatID,
|
||||
Model: m.opts.Model,
|
||||
SystemPrompt: m.systemPrompt(""),
|
||||
Format: m.opts.Format,
|
||||
Options: m.opts.Options,
|
||||
Think: m.opts.Think,
|
||||
KeepAlive: m.opts.KeepAlive,
|
||||
}
|
||||
}
|
||||
|
||||
func (m chatModel) previewMessages() []api.Message {
|
||||
if len(m.liveMessages) > 0 {
|
||||
return slices.Clone(m.liveMessages)
|
||||
}
|
||||
return slices.Clone(m.messages)
|
||||
}
|
||||
|
||||
func (m chatModel) previewChatRequest(opts coreagent.RunOptions, messages []api.Message) api.ChatRequest {
|
||||
requestMessages := slices.Clone(messages)
|
||||
if strings.TrimSpace(opts.SystemPrompt) != "" {
|
||||
withSystem := make([]api.Message, 0, len(requestMessages)+1)
|
||||
withSystem = append(withSystem, api.Message{Role: "system", Content: opts.SystemPrompt})
|
||||
requestMessages = append(withSystem, requestMessages...)
|
||||
}
|
||||
|
||||
format := opts.Format
|
||||
if format == "json" {
|
||||
format = `"` + format + `"`
|
||||
}
|
||||
|
||||
req := api.ChatRequest{
|
||||
Model: opts.Model,
|
||||
Messages: requestMessages,
|
||||
Format: json.RawMessage(format),
|
||||
Options: opts.Options,
|
||||
Think: opts.Think,
|
||||
}
|
||||
if opts.KeepAlive != nil {
|
||||
req.KeepAlive = opts.KeepAlive
|
||||
}
|
||||
if m.opts.Tools != nil && !m.opts.ToolsDisabled {
|
||||
req.Tools = m.opts.Tools.Tools()
|
||||
}
|
||||
return req
|
||||
}
|
||||
|
||||
func (m chatModel) promptDebugLines(width int) []string {
|
||||
if m.promptDebug == nil {
|
||||
return nil
|
||||
}
|
||||
req := m.promptDebug.request
|
||||
innerWidth := max(20, width-2)
|
||||
lines := []string{
|
||||
chatHeaderStyle.Render("Request"),
|
||||
promptDebugFieldLine("model", req.Model, innerWidth),
|
||||
promptDebugFieldLine("estimated prompt", m.promptTokenText(m.promptDebug.tokens), innerWidth),
|
||||
promptDebugFieldLine("messages", fmt.Sprint(len(req.Messages)), innerWidth),
|
||||
promptDebugFieldLine("tools", fmt.Sprint(len(req.Tools)), innerWidth),
|
||||
}
|
||||
if len(req.Format) > 0 {
|
||||
lines = append(lines, promptDebugFieldLine("format", strings.TrimSpace(string(req.Format)), innerWidth))
|
||||
}
|
||||
if req.Options != nil {
|
||||
lines = append(lines, promptDebugMapLines("options", req.Options, innerWidth)...)
|
||||
}
|
||||
if req.Think != nil {
|
||||
lines = append(lines, promptDebugBlockLines("think", req.Think.String(), innerWidth, chatHistoryTextStyle)...)
|
||||
}
|
||||
if req.KeepAlive != nil {
|
||||
lines = append(lines, promptDebugFieldLine("keep_alive", req.KeepAlive.String(), innerWidth))
|
||||
}
|
||||
lines = append(lines, "", chatHeaderStyle.Render("Messages"))
|
||||
if len(req.Messages) == 0 {
|
||||
lines = append(lines, chatMetaStyle.Render("none"))
|
||||
} else {
|
||||
for i, msg := range req.Messages {
|
||||
if i > 0 {
|
||||
lines = append(lines, "")
|
||||
}
|
||||
lines = append(lines, promptDebugMessageLines(i+1, msg, innerWidth)...)
|
||||
}
|
||||
}
|
||||
lines = append(lines, "", chatHeaderStyle.Render("Tools"))
|
||||
if len(req.Tools) == 0 {
|
||||
lines = append(lines, chatMetaStyle.Render("none"))
|
||||
return lines
|
||||
}
|
||||
for i, tool := range req.Tools {
|
||||
if i > 0 {
|
||||
lines = append(lines, "")
|
||||
}
|
||||
lines = append(lines, promptDebugToolLines(i+1, tool, innerWidth)...)
|
||||
}
|
||||
return lines
|
||||
}
|
||||
|
||||
func promptDebugFieldLine(label, value string, width int) string {
|
||||
labelText := label + ":"
|
||||
value = strings.TrimSpace(value)
|
||||
if value == "" {
|
||||
value = "_empty_"
|
||||
}
|
||||
line := chatHistoryLabelStyle.Render(labelText) + " " + chatHistoryTextStyle.Render(value)
|
||||
return truncateRenderedLine(line, width)
|
||||
}
|
||||
|
||||
func promptDebugMessageLines(index int, msg api.Message, width int) []string {
|
||||
role := promptMessageLabel(msg)
|
||||
header := fmt.Sprintf("%d. %s", index, role)
|
||||
lines := []string{historyRoleStyle(msg.Role).Render(header)}
|
||||
|
||||
if strings.TrimSpace(msg.Thinking) != "" {
|
||||
lines = append(lines, promptDebugBlockLines("thinking", msg.Thinking, width, chatHistoryTextStyle)...)
|
||||
}
|
||||
if msg.Role != "tool" && (strings.TrimSpace(msg.Content) != "" || (msg.Role != "assistant" && len(msg.ToolCalls) == 0 && len(msg.Images) == 0 && msg.Thinking == "")) {
|
||||
lines = append(lines, promptDebugBlockLines("content", msg.Content, width, chatHistoryTextStyle)...)
|
||||
}
|
||||
if len(msg.ToolCalls) > 0 {
|
||||
for i, call := range msg.ToolCalls {
|
||||
lines = append(lines, promptDebugToolCallLines(i+1, call, width)...)
|
||||
}
|
||||
}
|
||||
if msg.Role == "tool" {
|
||||
if msg.ToolName != "" {
|
||||
lines = append(lines, " "+chatHistoryLabelStyle.Render("tool_name:")+" "+chatHistoryTextStyle.Render(msg.ToolName))
|
||||
}
|
||||
if msg.ToolCallID != "" {
|
||||
lines = append(lines, " "+chatHistoryLabelStyle.Render("tool_call_id:")+" "+chatHistoryTextStyle.Render(msg.ToolCallID))
|
||||
}
|
||||
lines = append(lines, promptDebugBlockLines("tool result", promptDebugToolResult(msg.Content), width, chatHistoryTextStyle)...)
|
||||
}
|
||||
if len(msg.Images) > 0 {
|
||||
lines = append(lines, " "+chatHistoryLabelStyle.Render(fmt.Sprintf("%d image%s", len(msg.Images), pluralSuffix(len(msg.Images)))))
|
||||
}
|
||||
return lines
|
||||
}
|
||||
|
||||
func promptDebugToolResult(content string) string {
|
||||
runes := []rune(content)
|
||||
if len(runes) <= maxPromptDebugToolResultRunes {
|
||||
return content
|
||||
}
|
||||
return string(runes[:maxPromptDebugToolResultRunes-3]) + "..."
|
||||
}
|
||||
|
||||
func promptDebugMapLines(label string, values map[string]any, width int) []string {
|
||||
lines := []string{" " + chatHistoryLabelStyle.Render(label+":")}
|
||||
if len(values) == 0 {
|
||||
return append(lines, " "+chatMetaStyle.Render("_empty_"))
|
||||
}
|
||||
keys := make([]string, 0, len(values))
|
||||
for key := range values {
|
||||
keys = append(keys, key)
|
||||
}
|
||||
slices.Sort(keys)
|
||||
for _, key := range keys {
|
||||
lines = append(lines, promptDebugValueLine(4, key, values[key], width)...)
|
||||
}
|
||||
return lines
|
||||
}
|
||||
|
||||
func promptDebugToolLines(index int, tool api.Tool, width int) []string {
|
||||
name := strings.TrimSpace(tool.Function.Name)
|
||||
if name == "" {
|
||||
name = "_unnamed_"
|
||||
}
|
||||
lines := []string{historyRoleStyle("tool").Render(fmt.Sprintf("%d. %s", index, name))}
|
||||
if strings.TrimSpace(tool.Function.Description) != "" {
|
||||
lines = append(lines, promptDebugBlockLines("description", tool.Function.Description, width, chatHistoryTextStyle)...)
|
||||
}
|
||||
|
||||
params := tool.Function.Parameters
|
||||
if params.Type != "" || params.Properties != nil {
|
||||
kind := params.Type
|
||||
if kind == "" {
|
||||
kind = "object"
|
||||
}
|
||||
lines = append(lines, " "+chatHistoryLabelStyle.Render("parameters:")+" "+chatHistoryTextStyle.Render(kind))
|
||||
}
|
||||
if params.Properties == nil || params.Properties.Len() == 0 {
|
||||
return lines
|
||||
}
|
||||
|
||||
lines = append(lines, " "+chatHistoryLabelStyle.Render("properties:"))
|
||||
required := map[string]bool{}
|
||||
for _, name := range params.Required {
|
||||
required[name] = true
|
||||
}
|
||||
for name, property := range params.Properties.All() {
|
||||
label := name
|
||||
propertyType := property.ToTypeScriptType()
|
||||
switch {
|
||||
case propertyType != "" && required[name]:
|
||||
label += " (" + propertyType + ", required)"
|
||||
case propertyType != "":
|
||||
label += " (" + propertyType + ")"
|
||||
case required[name]:
|
||||
label += " (required)"
|
||||
}
|
||||
value := strings.TrimSpace(property.Description)
|
||||
if value == "" {
|
||||
value = promptDebugPropertyDetails(property)
|
||||
}
|
||||
lines = append(lines, promptDebugTextLine(4, label, value, width)...)
|
||||
}
|
||||
return lines
|
||||
}
|
||||
|
||||
func promptDebugToolCallLines(index int, call api.ToolCall, width int) []string {
|
||||
name := strings.TrimSpace(call.Function.Name)
|
||||
if name == "" {
|
||||
name = "_unnamed_"
|
||||
}
|
||||
lines := []string{" " + chatHistoryLabelStyle.Render(fmt.Sprintf("tool call %d:", index)) + " " + chatHistoryTextStyle.Render(name)}
|
||||
if strings.TrimSpace(call.ID) != "" {
|
||||
lines = append(lines, promptDebugTextLine(4, "id", call.ID, width)...)
|
||||
}
|
||||
if call.Function.Arguments.Len() == 0 {
|
||||
lines = append(lines, " "+chatHistoryLabelStyle.Render("arguments:")+" "+chatMetaStyle.Render("none"))
|
||||
return lines
|
||||
}
|
||||
lines = append(lines, " "+chatHistoryLabelStyle.Render("arguments:"))
|
||||
for key, value := range call.Function.Arguments.All() {
|
||||
lines = append(lines, promptDebugValueLine(6, key, value, width)...)
|
||||
}
|
||||
return lines
|
||||
}
|
||||
|
||||
func promptDebugPropertyDetails(property api.ToolProperty) string {
|
||||
var parts []string
|
||||
if len(property.Enum) > 0 {
|
||||
values := make([]string, 0, len(property.Enum))
|
||||
for _, value := range property.Enum {
|
||||
values = append(values, promptDebugValueText(value))
|
||||
}
|
||||
parts = append(parts, "one of "+strings.Join(values, ", "))
|
||||
}
|
||||
if property.Properties != nil && property.Properties.Len() > 0 {
|
||||
count := property.Properties.Len()
|
||||
noun := "property"
|
||||
if count != 1 {
|
||||
noun = "properties"
|
||||
}
|
||||
parts = append(parts, fmt.Sprintf("%d nested %s", count, noun))
|
||||
}
|
||||
if property.Items != nil {
|
||||
parts = append(parts, "array items: "+promptDebugValueText(property.Items))
|
||||
}
|
||||
if len(parts) == 0 {
|
||||
return "_empty_"
|
||||
}
|
||||
return strings.Join(parts, "; ")
|
||||
}
|
||||
|
||||
func promptDebugValueLine(indent int, label string, value any, width int) []string {
|
||||
return promptDebugTextLine(indent, label, promptDebugValueText(value), width)
|
||||
}
|
||||
|
||||
func promptDebugTextLine(indent int, label, value string, width int) []string {
|
||||
prefix := strings.Repeat(" ", indent) + chatHistoryLabelStyle.Render(label+":")
|
||||
value = strings.TrimSpace(value)
|
||||
if value == "" {
|
||||
value = "_empty_"
|
||||
}
|
||||
wrapWidth := max(20, width-indent-lipgloss.Width(label)-2)
|
||||
wrapped := wrapChatText(value, wrapWidth)
|
||||
if len(wrapped) == 0 {
|
||||
return []string{prefix + " " + chatMetaStyle.Render("_empty_")}
|
||||
}
|
||||
lines := []string{prefix + " " + chatHistoryTextStyle.Render(wrapped[0])}
|
||||
for _, line := range wrapped[1:] {
|
||||
lines = append(lines, strings.Repeat(" ", indent+2)+chatHistoryTextStyle.Render(line))
|
||||
}
|
||||
return lines
|
||||
}
|
||||
|
||||
func promptDebugValueText(value any) string {
|
||||
switch v := value.(type) {
|
||||
case nil:
|
||||
return "null"
|
||||
case string:
|
||||
return v
|
||||
case fmt.Stringer:
|
||||
return v.String()
|
||||
case []any:
|
||||
parts := make([]string, 0, len(v))
|
||||
for _, item := range v {
|
||||
parts = append(parts, promptDebugValueText(item))
|
||||
}
|
||||
return strings.Join(parts, ", ")
|
||||
case map[string]any:
|
||||
keys := make([]string, 0, len(v))
|
||||
for key := range v {
|
||||
keys = append(keys, key)
|
||||
}
|
||||
slices.Sort(keys)
|
||||
parts := make([]string, 0, len(keys))
|
||||
for _, key := range keys {
|
||||
parts = append(parts, key+": "+promptDebugValueText(v[key]))
|
||||
}
|
||||
return strings.Join(parts, ", ")
|
||||
default:
|
||||
return fmt.Sprint(value)
|
||||
}
|
||||
}
|
||||
|
||||
func promptDebugBlockLines(label, value string, width int, style lipgloss.Style) []string {
|
||||
lines := []string{" " + chatHistoryLabelStyle.Render(label+":")}
|
||||
if value == "" {
|
||||
return append(lines, " "+chatMetaStyle.Render("_empty_"))
|
||||
}
|
||||
for _, raw := range strings.Split(strings.TrimRight(value, "\n"), "\n") {
|
||||
if raw == "" {
|
||||
lines = append(lines, "")
|
||||
continue
|
||||
}
|
||||
for _, wrapped := range wrapChatText(raw, max(20, width-4)) {
|
||||
lines = append(lines, " "+style.Render(wrapped))
|
||||
}
|
||||
}
|
||||
return lines
|
||||
}
|
||||
|
||||
func (m chatModel) promptTokenText(tokens int) string {
|
||||
window := m.displayContextWindowTokens()
|
||||
if window > 0 {
|
||||
return fmt.Sprintf("%s / %s tokens", formatPromptTokenCount(max(tokens, 0)), formatPromptTokenCount(window))
|
||||
}
|
||||
return formatTokenCount(tokens)
|
||||
}
|
||||
|
||||
func formatPromptTokenCount(count int) string {
|
||||
sign := ""
|
||||
if count < 0 {
|
||||
sign = "-"
|
||||
count = -count
|
||||
}
|
||||
if count < 100_000 {
|
||||
return sign + fmt.Sprint(count)
|
||||
}
|
||||
if count >= 950_000 {
|
||||
return fmt.Sprintf("%s%dM", sign, int(float64(count)/1_000_000+0.5))
|
||||
}
|
||||
return fmt.Sprintf("%s%dk", sign, int(float64(count)/1024+0.5))
|
||||
}
|
||||
|
||||
func promptMessageLabel(msg api.Message) string {
|
||||
if msg.Role == "tool" && msg.ToolName != "" {
|
||||
return msg.Role + ":" + msg.ToolName
|
||||
}
|
||||
return msg.Role
|
||||
}
|
||||
|
||||
func saveRequestFilename(args string) (string, error) {
|
||||
args = strings.TrimSpace(args)
|
||||
if args == "" {
|
||||
return "", fmt.Errorf("usage: /save <filename>")
|
||||
}
|
||||
if strings.HasPrefix(args, ">") {
|
||||
args = strings.TrimSpace(strings.TrimPrefix(args, ">"))
|
||||
}
|
||||
fields := strings.Fields(args)
|
||||
if len(fields) != 1 {
|
||||
return "", fmt.Errorf("usage: /save <filename>")
|
||||
}
|
||||
filename := strings.TrimSpace(fields[0])
|
||||
if filename == "" || filename == "." || filename == ".." || strings.ContainsAny(filename, `/\`) || filepath.IsAbs(filename) {
|
||||
return "", fmt.Errorf("save filename must be a file name, not a path")
|
||||
}
|
||||
if !strings.HasSuffix(strings.ToLower(filename), ".json") {
|
||||
filename += ".json"
|
||||
}
|
||||
return filename, nil
|
||||
}
|
||||
|
||||
func (m chatModel) debugWorkingDir() (string, error) {
|
||||
dir := strings.TrimSpace(m.currentWorkingDir())
|
||||
if dir != "" {
|
||||
return dir, nil
|
||||
}
|
||||
return os.Getwd()
|
||||
}
|
||||
@@ -0,0 +1,373 @@
|
||||
package chat
|
||||
|
||||
import (
|
||||
"context"
|
||||
"slices"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
tea "github.com/charmbracelet/bubbletea"
|
||||
|
||||
coreagent "github.com/ollama/ollama/agent"
|
||||
"github.com/ollama/ollama/api"
|
||||
)
|
||||
|
||||
type chatAgentMsg struct {
|
||||
event coreagent.Event
|
||||
}
|
||||
|
||||
type chatApprovalPromptMsg struct {
|
||||
request coreagent.ApprovalRequest
|
||||
reply chan<- coreagent.Approval
|
||||
}
|
||||
|
||||
type chatRunDoneMsg struct {
|
||||
result *coreagent.RunResult
|
||||
err error
|
||||
newMessagesPersisted bool
|
||||
persistedMessages []api.Message
|
||||
}
|
||||
|
||||
type chatCompactDoneMsg struct {
|
||||
result coreagent.CompactionResult
|
||||
err error
|
||||
}
|
||||
|
||||
type chatCompactProgressMsg struct {
|
||||
tokens int
|
||||
}
|
||||
|
||||
// resetStreamingState clears the transient streaming flags that every
|
||||
// non-streaming event resets before applying its own state.
|
||||
func (m *chatModel) resetStreamingState() {
|
||||
m.finishThinkingEntry()
|
||||
m.awaitingModel = false
|
||||
m.thinking = false
|
||||
m.thinkingTokens = 0
|
||||
}
|
||||
|
||||
// resetRunState clears all run-progress flags (streaming plus compaction
|
||||
// progress) for terminal events that fully reset the run view.
|
||||
func (m *chatModel) resetRunState() {
|
||||
m.finishThinkingEntry()
|
||||
m.awaitingModel = false
|
||||
m.compacting = false
|
||||
m.compactingTokens = 0
|
||||
m.detectedToolCalls = nil
|
||||
m.thinking = false
|
||||
m.thinkingTokens = 0
|
||||
}
|
||||
|
||||
type chatModelPreloadDoneMsg struct {
|
||||
model string
|
||||
contextWindowTokens int
|
||||
err error
|
||||
}
|
||||
|
||||
type chatEventsClosedMsg struct{}
|
||||
|
||||
type chatTickMsg struct{}
|
||||
|
||||
func (m *chatModel) applyAgentEvent(event coreagent.Event) {
|
||||
contextChanged := false
|
||||
|
||||
switch event.Type {
|
||||
case coreagent.EventThinkingDelta:
|
||||
m.awaitingModel = false
|
||||
if event.Thinking != "" {
|
||||
m.thinking = true
|
||||
if event.Tokens > 0 {
|
||||
m.thinkingTokens = max(m.thinkingTokens, event.Tokens)
|
||||
} else {
|
||||
m.thinkingTokens += approximateTokenCount(event.Thinking)
|
||||
}
|
||||
idx := m.ensureLiveAssistantMessage()
|
||||
m.liveMessages[idx].Thinking += event.Thinking
|
||||
m.syncThinkingEntry()
|
||||
contextChanged = true
|
||||
}
|
||||
case coreagent.EventMessageDelta:
|
||||
m.resetStreamingState()
|
||||
m.spinner = 0
|
||||
m.detectedToolCalls = nil
|
||||
idx := m.ensureAssistantEntry()
|
||||
m.entries[idx].content += event.Content
|
||||
m.markEntryDirty(idx)
|
||||
msgIdx := m.ensureLiveAssistantMessage()
|
||||
m.liveMessages[msgIdx].Content += event.Content
|
||||
contextChanged = true
|
||||
case coreagent.EventToolCallDetected:
|
||||
m.finishThinkingEntry()
|
||||
m.awaitingModel = m.running
|
||||
m.thinking = false
|
||||
m.thinkingTokens = 0
|
||||
m.groupCompletedToolHistory()
|
||||
m.detectedToolCalls = nil
|
||||
m.addDetectedToolCalls(event.ToolCalls)
|
||||
idx := m.ensureLiveAssistantMessage()
|
||||
m.liveMessages[idx].ToolCalls = append(m.liveMessages[idx].ToolCalls, event.ToolCalls...)
|
||||
contextChanged = true
|
||||
case coreagent.EventToolStarted:
|
||||
m.resetStreamingState()
|
||||
m.refreshContextWindowTokens(m.opts.Model)
|
||||
startedAt := time.Now()
|
||||
idx := m.findActiveToolEntry(event.ToolCallID)
|
||||
if idx < 0 {
|
||||
m.groupCompletedToolHistory()
|
||||
m.entries = append(m.entries, newChatEntry(chatEntry{role: "tool"}))
|
||||
idx = len(m.entries) - 1
|
||||
}
|
||||
m.entries[idx].detail = event.ToolName
|
||||
m.entries[idx].label = toolInvocationLabel(event.ToolName, event.Args)
|
||||
m.entries[idx].status = "running"
|
||||
m.entries[idx].toolID = event.ToolCallID
|
||||
m.entries[idx].args = event.Args
|
||||
m.entries[idx].startedAt = startedAt
|
||||
m.applyToolOutputModeTo(idx)
|
||||
m.markEntryDirty(idx)
|
||||
case coreagent.EventToolFinished:
|
||||
m.resetStreamingState()
|
||||
m.refreshContextWindowTokens(m.opts.Model)
|
||||
if event.WorkingDir != "" {
|
||||
m.workingDir = event.WorkingDir
|
||||
}
|
||||
startedAt := m.toolStartedAt(event.ToolCallID)
|
||||
status := toolFinishedStatus(event)
|
||||
idx := m.findToolEntry(event.ToolCallID)
|
||||
if idx < 0 {
|
||||
m.entries = append(m.entries, newChatEntry(chatEntry{role: "tool"}))
|
||||
idx = len(m.entries) - 1
|
||||
}
|
||||
m.entries[idx].content = event.Content
|
||||
m.entries[idx].label = toolInvocationLabel(event.ToolName, event.Args)
|
||||
m.entries[idx].detail = event.ToolName
|
||||
m.entries[idx].status = status
|
||||
if status != "denied" {
|
||||
m.entries[idx].err = event.Error
|
||||
}
|
||||
m.entries[idx].toolID = event.ToolCallID
|
||||
m.entries[idx].args = event.Args
|
||||
m.entries[idx].startedAt = startedAt
|
||||
m.entries[idx].finishedAt = time.Now()
|
||||
m.applyToolOutputModeTo(idx)
|
||||
m.markEntryDirty(idx)
|
||||
m.liveMessages = append(m.liveMessages, api.Message{
|
||||
Role: "tool",
|
||||
Content: event.Content,
|
||||
ToolName: event.ToolName,
|
||||
ToolCallID: event.ToolCallID,
|
||||
})
|
||||
if m.running && status != "denied" && !m.hasPendingDetectedToolCalls() {
|
||||
m.awaitingModel = true
|
||||
}
|
||||
contextChanged = true
|
||||
case coreagent.EventCompacted:
|
||||
m.resetRunState()
|
||||
if len(event.Messages) > 0 {
|
||||
m.liveMessages = slices.Clone(event.Messages)
|
||||
m.messages = slices.Clone(event.Messages)
|
||||
contextChanged = true
|
||||
}
|
||||
m.status = "compacted"
|
||||
case coreagent.EventCompactionStarted:
|
||||
m.awaitingModel = false
|
||||
m.compacting = true
|
||||
m.compactingTokens = 0
|
||||
m.thinking = false
|
||||
m.thinkingTokens = 0
|
||||
m.status = "compacting"
|
||||
case coreagent.EventCompactionProgress:
|
||||
m.awaitingModel = false
|
||||
m.compacting = true
|
||||
m.thinking = false
|
||||
m.thinkingTokens = 0
|
||||
if event.Tokens > m.compactingTokens {
|
||||
m.compactingTokens = event.Tokens
|
||||
}
|
||||
case coreagent.EventCompactionSkipped:
|
||||
m.resetRunState()
|
||||
message := event.Content
|
||||
if strings.TrimSpace(message) == "" {
|
||||
message = coreagent.CompactionSkippedMessage(event.Error)
|
||||
}
|
||||
m.entries = append(m.entries, newChatEntry(chatEntry{role: "system", content: message}))
|
||||
m.status = "compact skipped"
|
||||
case coreagent.EventError:
|
||||
m.resetRunState()
|
||||
m.eventErrorRendered = true
|
||||
m.entries = append(m.entries, newChatEntry(chatEntry{role: "error", content: event.Error, err: event.Error}))
|
||||
}
|
||||
|
||||
if contextChanged {
|
||||
m.refreshLiveContextEstimate()
|
||||
}
|
||||
}
|
||||
|
||||
func (m *chatModel) addDetectedToolCalls(calls []api.ToolCall) {
|
||||
if len(calls) == 0 {
|
||||
return
|
||||
}
|
||||
seen := make(map[string]struct{}, len(m.detectedToolCalls)+len(calls))
|
||||
for _, entry := range m.detectedToolCalls {
|
||||
if entry.toolID != "" {
|
||||
seen[entry.toolID] = struct{}{}
|
||||
}
|
||||
}
|
||||
for _, call := range calls {
|
||||
if call.ID != "" {
|
||||
if _, ok := seen[call.ID]; ok {
|
||||
continue
|
||||
}
|
||||
seen[call.ID] = struct{}{}
|
||||
}
|
||||
args := call.Function.Arguments.ToMap()
|
||||
m.detectedToolCalls = append(m.detectedToolCalls, newChatEntry(chatEntry{
|
||||
role: "tool",
|
||||
label: toolInvocationLabel(call.Function.Name, args),
|
||||
detail: call.Function.Name,
|
||||
status: "queued",
|
||||
toolID: call.ID,
|
||||
args: args,
|
||||
}))
|
||||
}
|
||||
}
|
||||
|
||||
func toolFinishedStatus(event coreagent.Event) string {
|
||||
switch event.ToolStatus {
|
||||
case coreagent.ToolStatusDenied:
|
||||
return "denied"
|
||||
case coreagent.ToolStatusDisabled:
|
||||
return "disabled"
|
||||
case coreagent.ToolStatusDone:
|
||||
return "done"
|
||||
}
|
||||
// failed/skipped/unknown: derive from content and error fields.
|
||||
if isDeniedToolResult(event.Content) || isDeniedToolResult(event.Error) {
|
||||
return "denied"
|
||||
}
|
||||
if event.Error != "" {
|
||||
return "error"
|
||||
}
|
||||
return "done"
|
||||
}
|
||||
|
||||
func messagesEndWithCompactionResult(messages []api.Message) bool {
|
||||
if len(messages) == 0 {
|
||||
return false
|
||||
}
|
||||
return coreagent.IsCompactionToolResult(messages[len(messages)-1])
|
||||
}
|
||||
|
||||
func (m chatModel) awaitingToolStart() bool {
|
||||
for i := len(m.liveMessages) - 1; i >= 0; i-- {
|
||||
msg := m.liveMessages[i]
|
||||
if msg.Role != "assistant" {
|
||||
continue
|
||||
}
|
||||
if len(msg.ToolCalls) == 0 {
|
||||
return false
|
||||
}
|
||||
for _, call := range msg.ToolCalls {
|
||||
if call.ID == "" || m.findToolEntry(call.ID) < 0 {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func (m *chatModel) ensureLiveAssistantMessage() int {
|
||||
if len(m.liveMessages) > 0 && m.liveMessages[len(m.liveMessages)-1].Role == "assistant" {
|
||||
return len(m.liveMessages) - 1
|
||||
}
|
||||
m.liveMessages = append(m.liveMessages, api.Message{Role: "assistant"})
|
||||
return len(m.liveMessages) - 1
|
||||
}
|
||||
|
||||
func (m *chatModel) refreshLiveContextEstimate() {
|
||||
messages := m.liveMessages
|
||||
if len(messages) == 0 {
|
||||
messages = m.messages
|
||||
}
|
||||
m.contextTokens = m.estimatePromptTokens(messages, "")
|
||||
m.contextEstimate = true
|
||||
}
|
||||
|
||||
//nolint:containedctx // event sinks need the session context to unblock sends on cancellation.
|
||||
type chatEventSink struct {
|
||||
ctx context.Context
|
||||
ch chan<- tea.Msg
|
||||
newMessagesPersisted *bool
|
||||
}
|
||||
|
||||
func (s chatEventSink) Emit(event coreagent.Event) error {
|
||||
if s.newMessagesPersisted != nil {
|
||||
*s.newMessagesPersisted = true
|
||||
}
|
||||
select {
|
||||
case s.ch <- chatAgentMsg{event: event}:
|
||||
return nil
|
||||
case <-s.ctx.Done():
|
||||
return s.ctx.Err()
|
||||
}
|
||||
}
|
||||
|
||||
func waitForChatMsg(ch <-chan tea.Msg) tea.Cmd {
|
||||
if ch == nil {
|
||||
return nil
|
||||
}
|
||||
return func() tea.Msg {
|
||||
msg, ok := <-ch
|
||||
if !ok {
|
||||
return chatEventsClosedMsg{}
|
||||
}
|
||||
return msg
|
||||
}
|
||||
}
|
||||
|
||||
func (m *chatModel) scheduleTick() tea.Cmd {
|
||||
if m.tickActive {
|
||||
return nil
|
||||
}
|
||||
m.tickActive = true
|
||||
return chatTickCmd()
|
||||
}
|
||||
|
||||
func chatTickCmd() tea.Cmd {
|
||||
return tea.Tick(350*time.Millisecond, func(time.Time) tea.Msg {
|
||||
return chatTickMsg{}
|
||||
})
|
||||
}
|
||||
|
||||
func preloadModelCmd(ctx context.Context, preload func(context.Context, string, *api.ThinkValue) (int, error), model string, think *api.ThinkValue) tea.Cmd {
|
||||
if preload == nil || strings.TrimSpace(model) == "" {
|
||||
return nil
|
||||
}
|
||||
if think != nil {
|
||||
copied := *think
|
||||
think = &copied
|
||||
}
|
||||
return func() tea.Msg {
|
||||
if ctx == nil {
|
||||
ctx = context.Background()
|
||||
}
|
||||
tokens, err := preload(ctx, model, think)
|
||||
return chatModelPreloadDoneMsg{model: model, contextWindowTokens: tokens, err: err}
|
||||
}
|
||||
}
|
||||
|
||||
func isUnsupportedThinkingError(err error) bool {
|
||||
if err == nil {
|
||||
return false
|
||||
}
|
||||
text := strings.ToLower(err.Error())
|
||||
return strings.Contains(text, "does not support thinking")
|
||||
}
|
||||
|
||||
func thinkRequestsThinking(think *api.ThinkValue) bool {
|
||||
if think == nil {
|
||||
return false
|
||||
}
|
||||
return think.Bool()
|
||||
}
|
||||
@@ -0,0 +1,376 @@
|
||||
package chat
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
coreagent "github.com/ollama/ollama/agent"
|
||||
"github.com/ollama/ollama/api"
|
||||
)
|
||||
|
||||
func TestApplyAgentEventStreamsAssistantContent(t *testing.T) {
|
||||
m := chatModel{running: true}
|
||||
|
||||
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventMessageDelta, Content: "hello"})
|
||||
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventMessageDelta, Content: " world"})
|
||||
|
||||
if len(m.entries) != 1 || m.entries[0].role != "assistant" || m.entries[0].content != "hello world" {
|
||||
t.Fatalf("entries = %#v", m.entries)
|
||||
}
|
||||
if len(m.liveMessages) != 1 || m.liveMessages[0].Content != "hello world" {
|
||||
t.Fatalf("live messages = %#v", m.liveMessages)
|
||||
}
|
||||
}
|
||||
|
||||
func TestApplyAgentEventTracksToolLifecycle(t *testing.T) {
|
||||
m := chatModel{running: true}
|
||||
args := map[string]any{"command": "pwd"}
|
||||
|
||||
m.applyAgentEvent(coreagent.Event{
|
||||
Type: coreagent.EventToolStarted,
|
||||
ToolCallID: "call-1",
|
||||
ToolName: "bash",
|
||||
Args: args,
|
||||
})
|
||||
m.applyAgentEvent(coreagent.Event{
|
||||
Type: coreagent.EventToolFinished,
|
||||
ToolCallID: "call-1",
|
||||
ToolName: "bash",
|
||||
Args: args,
|
||||
Content: "ok",
|
||||
})
|
||||
|
||||
if len(m.entries) != 1 {
|
||||
t.Fatalf("entries = %#v", m.entries)
|
||||
}
|
||||
entry := m.entries[0]
|
||||
if entry.status != "done" || entry.content != "ok" || !strings.Contains(entry.label, "Bash") {
|
||||
t.Fatalf("tool entry = %#v", entry)
|
||||
}
|
||||
if line := stripANSI(toolStatusLine(entry)); line != `Bash("pwd")` {
|
||||
t.Fatalf("tool status line = %q, want command label", line)
|
||||
}
|
||||
if len(m.liveMessages) != 1 || m.liveMessages[0].Role != "tool" || m.liveMessages[0].Content != "ok" {
|
||||
t.Fatalf("live messages = %#v", m.liveMessages)
|
||||
}
|
||||
}
|
||||
|
||||
func TestApplyAgentEventRendersDeniedCommandAsDenied(t *testing.T) {
|
||||
m := chatModel{running: true}
|
||||
args := map[string]any{"command": "pwd"}
|
||||
|
||||
m.applyAgentEvent(coreagent.Event{
|
||||
Type: coreagent.EventToolFinished,
|
||||
ToolStatus: coreagent.ToolStatusDenied,
|
||||
ToolCallID: "call-1",
|
||||
ToolName: "bash",
|
||||
Args: args,
|
||||
Content: "Tool execution denied.",
|
||||
Error: "Tool execution denied.",
|
||||
})
|
||||
|
||||
if len(m.entries) != 1 {
|
||||
t.Fatalf("entries = %#v", m.entries)
|
||||
}
|
||||
entry := m.entries[0]
|
||||
if entry.status != "denied" {
|
||||
t.Fatalf("tool status = %q, want denied: %#v", entry.status, entry)
|
||||
}
|
||||
if line := stripANSI(toolStatusLine(entry)); line != `Bash("pwd") denied` {
|
||||
t.Fatalf("tool status line = %q, want denied command label", line)
|
||||
}
|
||||
}
|
||||
|
||||
func TestApplyAgentEventShowsWorkingWhileAwaitingCloudToolStart(t *testing.T) {
|
||||
args := api.NewToolCallFunctionArguments()
|
||||
args.Set("command", "pwd")
|
||||
m := chatModel{
|
||||
running: true,
|
||||
}
|
||||
|
||||
m.applyAgentEvent(coreagent.Event{
|
||||
Type: coreagent.EventToolCallDetected,
|
||||
ToolCalls: []api.ToolCall{{
|
||||
ID: "call-1",
|
||||
Function: api.ToolCallFunction{
|
||||
Name: "bash",
|
||||
Arguments: args,
|
||||
},
|
||||
}},
|
||||
})
|
||||
|
||||
if line := stripANSI(m.activityLine()); !strings.Contains(line, "Working") {
|
||||
t.Fatalf("activityLine = %q, want Working while tool call is pending", line)
|
||||
}
|
||||
}
|
||||
|
||||
func TestActivityLineShowsWorkingWhileAwaitingModelBeforeFirstEvent(t *testing.T) {
|
||||
m := chatModel{
|
||||
running: true,
|
||||
awaitingModel: true,
|
||||
spinner: 0,
|
||||
}
|
||||
|
||||
if line := stripANSI(m.activityLine()); !strings.Contains(line, "Working") {
|
||||
t.Fatalf("activityLine = %q, want Working while stream is open before first event", line)
|
||||
}
|
||||
}
|
||||
|
||||
func TestActivityLineShowsWorkingAfterAssistantContentGoesIdle(t *testing.T) {
|
||||
m := chatModel{
|
||||
running: true,
|
||||
spinner: idleWorkingDelayTicks,
|
||||
}
|
||||
|
||||
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventMessageDelta, Content: "I will inspect that next."})
|
||||
if line := strings.TrimSpace(stripANSI(m.activityLine())); line != "" {
|
||||
t.Fatalf("activityLine immediately after content = %q, want quiet until the idle delay", line)
|
||||
}
|
||||
|
||||
m.spinner = idleWorkingDelayTicks
|
||||
if line := stripANSI(m.activityLine()); !strings.Contains(line, "Working") {
|
||||
t.Fatalf("activityLine after idle content stream = %q, want Working while stream remains open", line)
|
||||
}
|
||||
}
|
||||
|
||||
func TestApplyAgentEventKeepsDetectedBatchStableUntilComplete(t *testing.T) {
|
||||
firstArgs := api.NewToolCallFunctionArguments()
|
||||
firstArgs.Set("command", "pwd")
|
||||
secondArgs := api.NewToolCallFunctionArguments()
|
||||
secondArgs.Set("command", "ls")
|
||||
m := chatModel{
|
||||
running: true,
|
||||
spinner: idleWorkingDelayTicks,
|
||||
}
|
||||
|
||||
m.applyAgentEvent(coreagent.Event{
|
||||
Type: coreagent.EventToolCallDetected,
|
||||
ToolCalls: []api.ToolCall{
|
||||
{ID: "call-1", Function: api.ToolCallFunction{Name: "bash", Arguments: firstArgs}},
|
||||
{ID: "call-2", Function: api.ToolCallFunction{Name: "bash", Arguments: secondArgs}},
|
||||
},
|
||||
})
|
||||
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventToolStarted, ToolCallID: "call-1", ToolName: "bash", Args: firstArgs.ToMap()})
|
||||
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventToolFinished, ToolCallID: "call-1", ToolName: "bash", Args: firstArgs.ToMap(), Content: "one"})
|
||||
|
||||
if len(m.entries) != 1 {
|
||||
t.Fatalf("entries = %d, want first completed command row: %#v", len(m.entries), m.entries)
|
||||
}
|
||||
if m.entries[0].role != "tool" || m.entries[0].status != "done" {
|
||||
t.Fatalf("first command should remain stable while second is pending: %#v", m.entries[0])
|
||||
}
|
||||
if line := stripANSI(toolStatusLine(m.entries[0])); line != `Bash("pwd")` {
|
||||
t.Fatalf("completed command line = %q", line)
|
||||
}
|
||||
if line := stripANSI(m.activityLine()); !strings.Contains(line, "Working") {
|
||||
t.Fatalf("activityLine = %q, want Working while second command is pending", line)
|
||||
}
|
||||
|
||||
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventToolStarted, ToolCallID: "call-2", ToolName: "bash", Args: secondArgs.ToMap()})
|
||||
if len(m.entries) != 2 {
|
||||
t.Fatalf("entries after second start = %d, want finished command plus running command: %#v", len(m.entries), m.entries)
|
||||
}
|
||||
if line := stripANSI(toolStatusLine(m.entries[0])); line != `Bash("pwd")` {
|
||||
t.Fatalf("finished command line after second start = %q", line)
|
||||
}
|
||||
if line := stripANSI(toolStatusLine(m.entries[1])); line != `Bash("ls")` {
|
||||
t.Fatalf("running command line = %q", line)
|
||||
}
|
||||
|
||||
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventToolFinished, ToolCallID: "call-2", ToolName: "bash", Args: secondArgs.ToMap(), Content: "two"})
|
||||
if line := stripANSI(m.activityLine()); !strings.Contains(line, "Working") {
|
||||
t.Fatalf("activityLine after completed batch = %q, want Working while waiting for next model response", line)
|
||||
}
|
||||
if len(m.entries) != 2 {
|
||||
t.Fatalf("entries after batch completion = %d, want stable command rows until the next tool boundary: %#v", len(m.entries), m.entries)
|
||||
}
|
||||
for i, want := range []string{`Bash("pwd")`, `Bash("ls")`} {
|
||||
if line := stripANSI(toolStatusLine(m.entries[i])); line != want {
|
||||
t.Fatalf("completed command row %d = %q, want %q", i, line, want)
|
||||
}
|
||||
}
|
||||
|
||||
thirdArgs := api.NewToolCallFunctionArguments()
|
||||
thirdArgs.Set("command", "date")
|
||||
m.applyAgentEvent(coreagent.Event{
|
||||
Type: coreagent.EventToolCallDetected,
|
||||
ToolCalls: []api.ToolCall{
|
||||
{ID: "call-3", Function: api.ToolCallFunction{Name: "bash", Arguments: thirdArgs}},
|
||||
},
|
||||
})
|
||||
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventToolStarted, ToolCallID: "call-3", ToolName: "bash", Args: thirdArgs.ToMap()})
|
||||
|
||||
if len(m.entries) != 2 {
|
||||
t.Fatalf("entries after next tool boundary = %d, want grouped history plus running command: %#v", len(m.entries), m.entries)
|
||||
}
|
||||
if m.entries[0].role != "tool_group" || len(m.entries[0].tools) != 2 {
|
||||
t.Fatalf("completed detected batch should collapse at the next tool boundary: %#v", m.entries[0])
|
||||
}
|
||||
if line := stripANSI(toolGroupStatusLine(m.entries[0])); line != "Ran 2 commands" {
|
||||
t.Fatalf("grouped command line = %q", line)
|
||||
}
|
||||
if line := stripANSI(toolStatusLine(m.entries[1])); line != `Bash("date")` {
|
||||
t.Fatalf("running command line = %q", line)
|
||||
}
|
||||
}
|
||||
|
||||
func TestApplyAgentEventDoesNotCollapsePartialDetectedBatch(t *testing.T) {
|
||||
firstArgs := api.NewToolCallFunctionArguments()
|
||||
firstArgs.Set("command", "pwd")
|
||||
secondArgs := api.NewToolCallFunctionArguments()
|
||||
secondArgs.Set("command", "ls")
|
||||
thirdArgs := api.NewToolCallFunctionArguments()
|
||||
thirdArgs.Set("command", "date")
|
||||
m := chatModel{running: true}
|
||||
|
||||
m.applyAgentEvent(coreagent.Event{
|
||||
Type: coreagent.EventToolCallDetected,
|
||||
ToolCalls: []api.ToolCall{
|
||||
{ID: "call-1", Function: api.ToolCallFunction{Name: "bash", Arguments: firstArgs}},
|
||||
{ID: "call-2", Function: api.ToolCallFunction{Name: "bash", Arguments: secondArgs}},
|
||||
{ID: "call-3", Function: api.ToolCallFunction{Name: "bash", Arguments: thirdArgs}},
|
||||
},
|
||||
})
|
||||
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventToolStarted, ToolCallID: "call-1", ToolName: "bash", Args: firstArgs.ToMap()})
|
||||
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventToolFinished, ToolCallID: "call-1", ToolName: "bash", Args: firstArgs.ToMap(), Content: "one"})
|
||||
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventToolStarted, ToolCallID: "call-2", ToolName: "bash", Args: secondArgs.ToMap()})
|
||||
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventToolFinished, ToolCallID: "call-2", ToolName: "bash", Args: secondArgs.ToMap(), Content: "two"})
|
||||
|
||||
if line := stripANSI(m.activityLine()); !strings.Contains(line, "Working") {
|
||||
t.Fatalf("activityLine before final detected call = %q, want Working while final tool is pending", line)
|
||||
}
|
||||
if len(m.entries) != 2 {
|
||||
t.Fatalf("entries before final detected call = %d, want two stable rows: %#v", len(m.entries), m.entries)
|
||||
}
|
||||
for i, want := range []string{`Bash("pwd")`, `Bash("ls")`} {
|
||||
if line := stripANSI(toolStatusLine(m.entries[i])); line != want {
|
||||
t.Fatalf("tool row %d = %q, want %q", i, line, want)
|
||||
}
|
||||
}
|
||||
|
||||
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventToolStarted, ToolCallID: "call-3", ToolName: "bash", Args: thirdArgs.ToMap()})
|
||||
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventToolFinished, ToolCallID: "call-3", ToolName: "bash", Args: thirdArgs.ToMap(), Content: "three"})
|
||||
|
||||
if len(m.entries) != 3 {
|
||||
t.Fatalf("entries after full detected batch = %#v, want stable tool rows until the next tool boundary", m.entries)
|
||||
}
|
||||
for i, want := range []string{`Bash("pwd")`, `Bash("ls")`, `Bash("date")`} {
|
||||
if line := stripANSI(toolStatusLine(m.entries[i])); line != want {
|
||||
t.Fatalf("tool row %d = %q, want %q", i, line, want)
|
||||
}
|
||||
}
|
||||
|
||||
fourthArgs := api.NewToolCallFunctionArguments()
|
||||
fourthArgs.Set("command", "whoami")
|
||||
m.applyAgentEvent(coreagent.Event{
|
||||
Type: coreagent.EventToolCallDetected,
|
||||
ToolCalls: []api.ToolCall{
|
||||
{ID: "call-4", Function: api.ToolCallFunction{Name: "bash", Arguments: fourthArgs}},
|
||||
},
|
||||
})
|
||||
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventToolStarted, ToolCallID: "call-4", ToolName: "bash", Args: fourthArgs.ToMap()})
|
||||
|
||||
if len(m.entries) != 2 || m.entries[0].role != "tool_group" || len(m.entries[0].tools) != 3 {
|
||||
t.Fatalf("entries after next detected batch starts = %#v, want one grouped history entry plus active tool", m.entries)
|
||||
}
|
||||
if line := stripANSI(toolGroupStatusLine(m.entries[0])); line != "Ran 3 commands" {
|
||||
t.Fatalf("grouped command line = %q", line)
|
||||
}
|
||||
}
|
||||
|
||||
func TestApplyAgentEventGroupsCompletedCommandsAtNextToolBoundary(t *testing.T) {
|
||||
m := chatModel{running: true}
|
||||
firstArgs := map[string]any{"command": "pwd"}
|
||||
secondArgs := map[string]any{"command": "ls"}
|
||||
thirdArgs := map[string]any{"command": "date"}
|
||||
|
||||
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventToolStarted, ToolCallID: "call-1", ToolName: "bash", Args: firstArgs})
|
||||
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventToolFinished, ToolCallID: "call-1", ToolName: "bash", Args: firstArgs, Content: "one"})
|
||||
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventToolStarted, ToolCallID: "call-2", ToolName: "bash", Args: secondArgs})
|
||||
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventToolFinished, ToolCallID: "call-2", ToolName: "bash", Args: secondArgs, Content: "two"})
|
||||
|
||||
if len(m.entries) != 2 {
|
||||
t.Fatalf("entries after second finish = %d, want two stable command rows: %#v", len(m.entries), m.entries)
|
||||
}
|
||||
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventToolStarted, ToolCallID: "call-3", ToolName: "bash", Args: thirdArgs})
|
||||
|
||||
if len(m.entries) != 2 {
|
||||
t.Fatalf("entries = %d, want grouped command history plus active command: %#v", len(m.entries), m.entries)
|
||||
}
|
||||
if m.entries[0].role != "tool_group" || len(m.entries[0].tools) != 2 {
|
||||
t.Fatalf("completed commands should be grouped when the next command starts: %#v", m.entries[0])
|
||||
}
|
||||
if line := stripANSI(toolGroupStatusLine(m.entries[0])); line != "Ran 2 commands" {
|
||||
t.Fatalf("grouped command line = %q", line)
|
||||
}
|
||||
|
||||
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventToolFinished, ToolCallID: "call-3", ToolName: "bash", Args: thirdArgs, Content: "three"})
|
||||
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventMessageDelta, Content: "done"})
|
||||
|
||||
if len(m.entries) != 3 {
|
||||
t.Fatalf("entries after assistant content = %d, want grouped history, last command, assistant: %#v", len(m.entries), m.entries)
|
||||
}
|
||||
transcript := stripANSI(m.renderTranscript(100))
|
||||
if !strings.Contains(transcript, "• Ran 2 commands\n\n• Bash(\"date\")\n\n done") {
|
||||
t.Fatalf("tool history should stay visually separated from assistant content:\n%s", transcript)
|
||||
}
|
||||
}
|
||||
|
||||
func TestApplyAgentEventDoesNotGroupCompletedCommandsOnMessageDelta(t *testing.T) {
|
||||
m := chatModel{running: true}
|
||||
firstArgs := map[string]any{"command": "pwd"}
|
||||
secondArgs := map[string]any{"command": "ls"}
|
||||
|
||||
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventToolStarted, ToolCallID: "call-1", ToolName: "bash", Args: firstArgs})
|
||||
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventToolFinished, ToolCallID: "call-1", ToolName: "bash", Args: firstArgs, Content: "one"})
|
||||
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventToolStarted, ToolCallID: "call-2", ToolName: "bash", Args: secondArgs})
|
||||
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventToolFinished, ToolCallID: "call-2", ToolName: "bash", Args: secondArgs, Content: "two"})
|
||||
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventMessageDelta, Content: "done"})
|
||||
|
||||
if len(m.entries) != 3 {
|
||||
t.Fatalf("entries = %d, want two command rows plus assistant content: %#v", len(m.entries), m.entries)
|
||||
}
|
||||
if m.entries[0].role != "tool" || m.entries[1].role != "tool" || m.entries[2].role != "assistant" {
|
||||
t.Fatalf("completed commands should not collapse on assistant content: %#v", m.entries)
|
||||
}
|
||||
}
|
||||
|
||||
func TestApplyAgentEventGroupsPreviouslyDeniedCommandsAtNextToolBoundary(t *testing.T) {
|
||||
m := chatModel{running: true}
|
||||
firstArgs := map[string]any{"command": "pwd"}
|
||||
secondArgs := map[string]any{"command": "ls"}
|
||||
thirdArgs := map[string]any{"command": "date"}
|
||||
|
||||
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventToolFinished, ToolStatus: coreagent.ToolStatusDenied, ToolCallID: "call-1", ToolName: "bash", Args: firstArgs, Content: "Tool execution denied.", Error: "Tool execution denied."})
|
||||
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventToolFinished, ToolStatus: coreagent.ToolStatusDenied, ToolCallID: "call-2", ToolName: "bash", Args: secondArgs, Content: "Tool execution denied.", Error: "Tool execution denied."})
|
||||
|
||||
if len(m.entries) != 2 {
|
||||
t.Fatalf("entries = %d, want two stable denied command rows: %#v", len(m.entries), m.entries)
|
||||
}
|
||||
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventToolStarted, ToolCallID: "call-3", ToolName: "bash", Args: thirdArgs})
|
||||
|
||||
if len(m.entries) != 2 {
|
||||
t.Fatalf("entries = %d, want grouped denied command entry plus active command: %#v", len(m.entries), m.entries)
|
||||
}
|
||||
if m.entries[0].role != "tool_group" || len(m.entries[0].tools) != 2 {
|
||||
t.Fatalf("denied commands should be grouped at the next tool boundary: %#v", m.entries[0])
|
||||
}
|
||||
if line := stripANSI(toolGroupStatusLine(m.entries[0])); line != "Denied 2 commands" {
|
||||
t.Fatalf("grouped command line = %q", line)
|
||||
}
|
||||
if line := stripANSI(toolStatusLine(m.entries[1])); line != `Bash("date")` {
|
||||
t.Fatalf("running command line = %q", line)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMessagesEndWithCompactionResult(t *testing.T) {
|
||||
messages := []api.Message{{
|
||||
Role: "tool",
|
||||
ToolName: coreagent.CompactionToolName,
|
||||
ToolCallID: coreagent.CompactionToolCallID,
|
||||
Content: coreagent.CompactionSummaryMessagePrefix + "summary",
|
||||
}}
|
||||
if !messagesEndWithCompactionResult(messages) {
|
||||
t.Fatal("expected compaction result")
|
||||
}
|
||||
}
|
||||
File diff suppressed because it is too large.
Load diff
File diff suppressed because it is too large.
Load diff
@@ -0,0 +1,294 @@
|
||||
package chat
|
||||
|
||||
import (
|
||||
"strings"
|
||||
|
||||
"github.com/charmbracelet/lipgloss"
|
||||
)
|
||||
|
||||
func renderMarkdownForView(markdown string, width int) string {
|
||||
if width < 20 {
|
||||
width = 20
|
||||
}
|
||||
|
||||
source := strings.Split(strings.TrimRight(markdown, "\n"), "\n")
|
||||
var rendered []string
|
||||
inCodeBlock := false
|
||||
for i := 0; i < len(source); i++ {
|
||||
line := strings.TrimRight(source[i], "\r")
|
||||
trimmed := strings.TrimSpace(line)
|
||||
|
||||
if strings.HasPrefix(trimmed, "```") {
|
||||
inCodeBlock = !inCodeBlock
|
||||
continue
|
||||
}
|
||||
if inCodeBlock {
|
||||
rendered = append(rendered, renderMarkdownCodeLine(line, width)...)
|
||||
continue
|
||||
}
|
||||
|
||||
if table, consumed := renderMarkdownTable(source[i:], width); consumed > 0 {
|
||||
rendered = append(rendered, table...)
|
||||
i += consumed - 1
|
||||
continue
|
||||
}
|
||||
|
||||
if heading, ok := markdownHeading(trimmed); ok {
|
||||
rendered = append(rendered, chatHeaderStyle.Render(heading))
|
||||
continue
|
||||
}
|
||||
|
||||
if trimmed == "" {
|
||||
rendered = append(rendered, "")
|
||||
continue
|
||||
}
|
||||
for _, wrapped := range wrapChatText(line, width) {
|
||||
rendered = append(rendered, renderMarkdownInline(wrapped))
|
||||
}
|
||||
}
|
||||
return strings.Join(rendered, "\n")
|
||||
}
|
||||
|
||||
func splitRenderedBody(body string) []string {
|
||||
body = strings.TrimRight(body, "\n")
|
||||
if body == "" {
|
||||
return []string{""}
|
||||
}
|
||||
return strings.Split(body, "\n")
|
||||
}
|
||||
|
||||
func markdownHeading(line string) (string, bool) {
|
||||
if !strings.HasPrefix(line, "#") {
|
||||
return "", false
|
||||
}
|
||||
level := 0
|
||||
for level < len(line) && line[level] == '#' {
|
||||
level++
|
||||
}
|
||||
if level == 0 || level > 6 || level >= len(line) || line[level] != ' ' {
|
||||
return "", false
|
||||
}
|
||||
return strings.TrimSpace(line[level:]), true
|
||||
}
|
||||
|
||||
func renderMarkdownInline(line string) string {
|
||||
var b strings.Builder
|
||||
for {
|
||||
before, rest, ok := strings.Cut(line, "`")
|
||||
b.WriteString(before)
|
||||
if !ok {
|
||||
break
|
||||
}
|
||||
code, after, ok := strings.Cut(rest, "`")
|
||||
if !ok {
|
||||
b.WriteString("`")
|
||||
b.WriteString(rest)
|
||||
break
|
||||
}
|
||||
b.WriteString(chatInlineCodeStyle.Render(code))
|
||||
line = after
|
||||
}
|
||||
return b.String()
|
||||
}
|
||||
|
||||
func renderMarkdownCodeLine(line string, width int) []string {
|
||||
codeWidth := max(1, width-2)
|
||||
lines := wrapChatText(line, codeWidth)
|
||||
for i, wrapped := range lines {
|
||||
lines[i] = " " + chatCodeBlockStyle.Render(wrapped)
|
||||
}
|
||||
return lines
|
||||
}
|
||||
|
||||
func renderMarkdownTable(lines []string, width int) ([]string, int) {
|
||||
if len(lines) < 2 || !looksLikeMarkdownTableRow(lines[0]) || !isMarkdownTableSeparator(lines[1]) {
|
||||
return nil, 0
|
||||
}
|
||||
|
||||
var rows [][]string
|
||||
consumed := 0
|
||||
for consumed < len(lines) && looksLikeMarkdownTableRow(lines[consumed]) {
|
||||
if consumed == 1 && isMarkdownTableSeparator(lines[consumed]) {
|
||||
consumed++
|
||||
continue
|
||||
}
|
||||
rows = append(rows, parseMarkdownTableRow(lines[consumed]))
|
||||
consumed++
|
||||
}
|
||||
if len(rows) == 0 {
|
||||
return nil, 0
|
||||
}
|
||||
|
||||
columnCount := 0
|
||||
for _, row := range rows {
|
||||
columnCount = max(columnCount, len(row))
|
||||
}
|
||||
naturalWidths := make([]int, columnCount)
|
||||
for _, row := range rows {
|
||||
for i := range columnCount {
|
||||
cell := ""
|
||||
if i < len(row) {
|
||||
cell = row[i]
|
||||
}
|
||||
naturalWidths[i] = max(naturalWidths[i], lipglossWidth(cell))
|
||||
}
|
||||
}
|
||||
widths := markdownTableColumnWidths(naturalWidths, width)
|
||||
|
||||
var rendered []string
|
||||
for rowIndex, row := range rows {
|
||||
wrappedCells := make([][]string, columnCount)
|
||||
rowHeight := 1
|
||||
for i := range columnCount {
|
||||
cell := ""
|
||||
if i < len(row) {
|
||||
cell = row[i]
|
||||
}
|
||||
wrappedCells[i] = wrapMarkdownTableCell(cell, widths[i])
|
||||
rowHeight = max(rowHeight, len(wrappedCells[i]))
|
||||
}
|
||||
for lineIndex := range rowHeight {
|
||||
cells := make([]string, columnCount)
|
||||
for i := range columnCount {
|
||||
cellLine := ""
|
||||
if lineIndex < len(wrappedCells[i]) {
|
||||
cellLine = wrappedCells[i][lineIndex]
|
||||
}
|
||||
cells[i] = padPlainLine(cellLine, widths[i])
|
||||
}
|
||||
line := strings.Join(cells, chatTableBorderStyle.Render(" | "))
|
||||
if rowIndex == 0 {
|
||||
line = chatHeaderStyle.Render(stripANSIForWidth(line))
|
||||
}
|
||||
rendered = append(rendered, line)
|
||||
}
|
||||
}
|
||||
return rendered, consumed
|
||||
}
|
||||
|
||||
func markdownTableColumnWidths(naturalWidths []int, width int) []int {
|
||||
if len(naturalWidths) == 0 {
|
||||
return nil
|
||||
}
|
||||
separatorWidth := max(0, len(naturalWidths)-1) * lipglossWidth(" | ")
|
||||
available := max(1, width-separatorWidth)
|
||||
widths := make([]int, len(naturalWidths))
|
||||
minWidths := make([]int, len(naturalWidths))
|
||||
for i, natural := range naturalWidths {
|
||||
widths[i] = max(1, natural)
|
||||
minWidth := min(widths[i], 12)
|
||||
if i == 0 {
|
||||
minWidth = min(widths[i], 4)
|
||||
}
|
||||
minWidths[i] = max(1, minWidth)
|
||||
}
|
||||
|
||||
for sumInts(widths) > available {
|
||||
index := widestShrinkableColumn(widths, minWidths)
|
||||
if index < 0 {
|
||||
break
|
||||
}
|
||||
widths[index]--
|
||||
}
|
||||
for sumInts(widths) > available {
|
||||
index := widestColumn(widths)
|
||||
if index < 0 || widths[index] <= 1 {
|
||||
break
|
||||
}
|
||||
widths[index]--
|
||||
}
|
||||
return widths
|
||||
}
|
||||
|
||||
func widestShrinkableColumn(widths, minWidths []int) int {
|
||||
index := -1
|
||||
for i, width := range widths {
|
||||
if width <= minWidths[i] {
|
||||
continue
|
||||
}
|
||||
if index < 0 || width > widths[index] {
|
||||
index = i
|
||||
}
|
||||
}
|
||||
return index
|
||||
}
|
||||
|
||||
func widestColumn(widths []int) int {
|
||||
index := -1
|
||||
for i, width := range widths {
|
||||
if index < 0 || width > widths[index] {
|
||||
index = i
|
||||
}
|
||||
}
|
||||
return index
|
||||
}
|
||||
|
||||
func sumInts(values []int) int {
|
||||
sum := 0
|
||||
for _, value := range values {
|
||||
sum += value
|
||||
}
|
||||
return sum
|
||||
}
|
||||
|
||||
func wrapMarkdownTableCell(cell string, width int) []string {
|
||||
width = max(1, width)
|
||||
var out []string
|
||||
line := strings.TrimSpace(cell)
|
||||
for lipglossWidth(line) > width {
|
||||
cut := chatDisplayWidthCut(line, width)
|
||||
out = append(out, strings.TrimSpace(line[:cut]))
|
||||
line = strings.TrimSpace(line[cut:])
|
||||
}
|
||||
out = append(out, line)
|
||||
if len(out) == 0 {
|
||||
return []string{""}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func looksLikeMarkdownTableRow(line string) bool {
|
||||
line = strings.TrimSpace(line)
|
||||
return strings.Contains(line, "|") && strings.Count(line, "|") >= 1
|
||||
}
|
||||
|
||||
func isMarkdownTableSeparator(line string) bool {
|
||||
cells := parseMarkdownTableRow(line)
|
||||
if len(cells) == 0 {
|
||||
return false
|
||||
}
|
||||
for _, cell := range cells {
|
||||
cell = strings.Trim(cell, " :-")
|
||||
if cell != "" {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
func parseMarkdownTableRow(line string) []string {
|
||||
line = strings.TrimSpace(line)
|
||||
line = strings.TrimPrefix(line, "|")
|
||||
line = strings.TrimSuffix(line, "|")
|
||||
raw := strings.Split(line, "|")
|
||||
cells := make([]string, 0, len(raw))
|
||||
for _, cell := range raw {
|
||||
cells = append(cells, strings.TrimSpace(cell))
|
||||
}
|
||||
return cells
|
||||
}
|
||||
|
||||
func padPlainLine(line string, width int) string {
|
||||
if extra := width - lipglossWidth(line); extra > 0 {
|
||||
return line + strings.Repeat(" ", extra)
|
||||
}
|
||||
return line
|
||||
}
|
||||
|
||||
func stripANSIForWidth(line string) string {
|
||||
return stripChatANSI(line)
|
||||
}
|
||||
|
||||
func lipglossWidth(line string) int {
|
||||
return lipgloss.Width(line)
|
||||
}
|
||||
@@ -0,0 +1,263 @@
|
||||
package chat
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"slices"
|
||||
"strings"
|
||||
|
||||
tea "github.com/charmbracelet/bubbletea"
|
||||
|
||||
apptui "github.com/ollama/ollama/cmd/tui"
|
||||
)
|
||||
|
||||
type chatModelPicker = apptui.SelectorModel
|
||||
|
||||
func (m *chatModel) openModelPicker(filter string) (tea.Model, tea.Cmd) {
|
||||
if m.opts.ModelOptions == nil {
|
||||
m.entries = append(m.entries, newChatEntry(chatEntry{role: "error", content: "Model picker is unavailable.", err: "Model picker is unavailable."}))
|
||||
m.status = "error"
|
||||
return *m, nil
|
||||
}
|
||||
|
||||
ctx := m.ctx
|
||||
if ctx == nil {
|
||||
ctx = context.Background()
|
||||
}
|
||||
models, err := m.opts.ModelOptions(ctx)
|
||||
if err != nil {
|
||||
m.entries = append(m.entries, newChatEntry(chatEntry{role: "error", content: fmt.Sprintf("Could not list models: %v", err), err: err.Error()}))
|
||||
m.status = "error"
|
||||
return *m, nil
|
||||
}
|
||||
models = normalizeModelOptions(models)
|
||||
if len(models) == 0 {
|
||||
m.entries = append(m.entries, newChatEntry(chatEntry{role: "system", content: "No models available."}))
|
||||
m.status = "ready"
|
||||
return *m, nil
|
||||
}
|
||||
|
||||
items := modelSelectorItems(models, m.opts.Model)
|
||||
current := m.opts.Model
|
||||
if !m.openModelOnInit {
|
||||
items = compactModelSelectorItems(models, m.opts.Model)
|
||||
current = ""
|
||||
}
|
||||
picker := apptui.NewModelSelectorModel("Select model", items, current, filter)
|
||||
picker.SetHelpText("↑/↓ navigate • enter select • type search • esc cancel")
|
||||
m.modelPicker = &picker
|
||||
m.modelPickerModels = models
|
||||
m.status = "model"
|
||||
return *m, nil
|
||||
}
|
||||
|
||||
func normalizeModelOptions(models []ModelOption) []ModelOption {
|
||||
seen := make(map[string]struct{}, len(models))
|
||||
out := make([]ModelOption, 0, len(models))
|
||||
for _, model := range models {
|
||||
model.Name = strings.TrimSpace(model.Name)
|
||||
model.Description = strings.TrimSpace(model.Description)
|
||||
if model.Name == "" {
|
||||
continue
|
||||
}
|
||||
key := strings.ToLower(model.Name)
|
||||
if _, ok := seen[key]; ok {
|
||||
continue
|
||||
}
|
||||
seen[key] = struct{}{}
|
||||
out = append(out, model)
|
||||
}
|
||||
slices.SortStableFunc(out, func(a, b ModelOption) int {
|
||||
if a.Recommended == b.Recommended {
|
||||
return 0
|
||||
}
|
||||
if a.Recommended {
|
||||
return -1
|
||||
}
|
||||
return 1
|
||||
})
|
||||
return out
|
||||
}
|
||||
|
||||
func modelSelectorItems(models []ModelOption, current string) []apptui.SelectItem {
|
||||
return modelSelectorItemsWithCurrentPriority(models, current, true)
|
||||
}
|
||||
|
||||
func compactModelSelectorItems(models []ModelOption, current string) []apptui.SelectItem {
|
||||
return modelSelectorItemsWithCurrentPriority(models, current, false)
|
||||
}
|
||||
|
||||
func modelSelectorItemsWithCurrentPriority(models []ModelOption, current string, pinCurrent bool) []apptui.SelectItem {
|
||||
ordered := slices.Clone(models)
|
||||
slices.SortStableFunc(ordered, func(a, b ModelOption) int {
|
||||
if cmp := compareModelPickerGroup(modelPickerGroup(a, current, pinCurrent), modelPickerGroup(b, current, pinCurrent)); cmp != 0 {
|
||||
return cmp
|
||||
}
|
||||
return 0
|
||||
})
|
||||
|
||||
items := make([]apptui.SelectItem, 0, len(ordered))
|
||||
for _, model := range ordered {
|
||||
items = append(items, apptui.SelectItem{
|
||||
Name: model.Name,
|
||||
Description: modelOptionMeta(model),
|
||||
Recommended: model.Name == current || !model.Cloud || model.Recommended,
|
||||
AvailabilityBadge: model.AvailabilityBadge,
|
||||
})
|
||||
}
|
||||
return items
|
||||
}
|
||||
|
||||
func modelPickerGroup(model ModelOption, current string, pinCurrent bool) int {
|
||||
if pinCurrent && model.Name == current {
|
||||
return 0
|
||||
}
|
||||
if model.Recommended {
|
||||
return 1
|
||||
}
|
||||
if model.Name == current {
|
||||
return 2
|
||||
}
|
||||
if !model.Cloud {
|
||||
return 3
|
||||
}
|
||||
return 4
|
||||
}
|
||||
|
||||
func compareModelPickerGroup(a, b int) int {
|
||||
switch {
|
||||
case a < b:
|
||||
return -1
|
||||
case a > b:
|
||||
return 1
|
||||
default:
|
||||
return 0
|
||||
}
|
||||
}
|
||||
|
||||
func (m chatModel) updateModelPicker(msg tea.KeyMsg) (tea.Model, tea.Cmd) {
|
||||
if m.modelPicker == nil {
|
||||
return m, nil
|
||||
}
|
||||
switch msg.Type {
|
||||
case tea.KeyCtrlC, tea.KeyEsc:
|
||||
m.modelPicker = nil
|
||||
m.modelPickerModels = nil
|
||||
m.openModelOnInit = false
|
||||
m.status = "ready"
|
||||
return m, nil
|
||||
case tea.KeyEnter:
|
||||
return m.selectModel()
|
||||
default:
|
||||
m.modelPicker.UpdateNavigation(msg)
|
||||
}
|
||||
return m, nil
|
||||
}
|
||||
|
||||
func (m chatModel) selectModel() (tea.Model, tea.Cmd) {
|
||||
if m.modelPicker == nil {
|
||||
return m, nil
|
||||
}
|
||||
selectedItem, ok := m.modelPicker.SelectedItem()
|
||||
if !ok {
|
||||
return m, nil
|
||||
}
|
||||
selected, ok := m.modelOptionForSelection(selectedItem.Name)
|
||||
if !ok {
|
||||
return m, nil
|
||||
}
|
||||
|
||||
// Cloud models need auth + plan check before switching. If we already
|
||||
// know the badge state from the model list, go directly to the right
|
||||
// prompt — no "checking" spinner.
|
||||
if selected.Cloud && m.opts.CheckCloudModel != nil {
|
||||
switch selected.AvailabilityBadge {
|
||||
case "Sign in required":
|
||||
return m.startCloudAuthSignIn(selected.Name, selected.RequiredPlan, selected.SignInURL)
|
||||
case "Upgrade required":
|
||||
return m.startCloudAuthUpgrade(selected.Name, selected.RequiredPlan)
|
||||
}
|
||||
// Badge is empty — auth is satisfied (confirmed via Whoami when the
|
||||
// list was built). Apply directly.
|
||||
}
|
||||
|
||||
m.modelPicker = nil
|
||||
m.modelPickerModels = nil
|
||||
m.openModelOnInit = false
|
||||
if err := m.applyModelSelection(selected.Name, true); err != nil {
|
||||
m.entries = append(m.entries, newChatEntry(chatEntry{role: "error", content: fmt.Sprintf("Could not switch model: %v", err), err: err.Error()}))
|
||||
m.status = "error"
|
||||
return m, nil
|
||||
}
|
||||
m.status = "ready"
|
||||
return m, tea.Batch(m.startModelPreload(selected.Name), cloudModelPreflightCmd(m.ctx, m.opts, selected.Name, selected.RequiredPlan))
|
||||
}
|
||||
|
||||
func (m chatModel) modelOptionForSelection(name string) (ModelOption, bool) {
|
||||
for _, model := range m.modelPickerModels {
|
||||
if model.Name == name {
|
||||
return model, true
|
||||
}
|
||||
}
|
||||
return ModelOption{}, false
|
||||
}
|
||||
|
||||
func (m *chatModel) applyModelSelection(modelName string, persist bool) error {
|
||||
modelName = strings.TrimSpace(modelName)
|
||||
if modelName == "" {
|
||||
return nil
|
||||
}
|
||||
m.opts.Model = modelName
|
||||
m.opts.ContextWindowTokens = 0
|
||||
if m.opts.ToolRegistryForModel != nil {
|
||||
m.opts.Tools = m.opts.ToolRegistryForModel(m.ctx, modelName)
|
||||
}
|
||||
if m.opts.SystemPromptForModel != nil {
|
||||
m.opts.SystemPrompt = m.opts.SystemPromptForModel(m.ctx, modelName, m.opts.Tools, m.opts.ToolsDisabled)
|
||||
}
|
||||
if m.opts.MultiModalForModel != nil {
|
||||
ctx := m.ctx
|
||||
if ctx == nil {
|
||||
ctx = context.Background()
|
||||
}
|
||||
m.opts.MultiModal = m.opts.MultiModalForModel(ctx, modelName)
|
||||
}
|
||||
m.refreshContextWindowTokens(modelName)
|
||||
m.contextTokens = m.estimatePromptTokens(m.messages, "")
|
||||
m.contextEstimate = true
|
||||
if persist && m.opts.OnModelSelected != nil {
|
||||
ctx := m.ctx
|
||||
if ctx == nil {
|
||||
ctx = context.Background()
|
||||
}
|
||||
return m.opts.OnModelSelected(ctx, modelName)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *chatModel) startModelPreload(modelName string) tea.Cmd {
|
||||
modelName = strings.TrimSpace(modelName)
|
||||
if m == nil || modelName == "" || m.opts.PreloadModel == nil {
|
||||
return nil
|
||||
}
|
||||
m.preloadingModel = modelName
|
||||
m.spinner = 0
|
||||
return tea.Batch(preloadModelCmd(m.ctx, m.opts.PreloadModel, modelName, m.opts.Think), m.scheduleTick())
|
||||
}
|
||||
|
||||
func (m chatModel) renderModelPicker(width int) string {
|
||||
return m.modelPicker.RenderContent()
|
||||
}
|
||||
|
||||
func (m chatModel) renderInlineModelPicker(width int) []string {
|
||||
rendered := m.modelPicker.RenderCompactContent(maxInlineModelPickerItems)
|
||||
lines := strings.Split(strings.TrimRight(rendered, "\n"), "\n")
|
||||
for i := range lines {
|
||||
lines[i] = truncateRenderedLine(lines[i], width)
|
||||
}
|
||||
return lines
|
||||
}
|
||||
|
||||
func modelOptionMeta(model ModelOption) string {
|
||||
return strings.TrimSpace(model.Description)
|
||||
}
|
||||
@@ -0,0 +1,441 @@
|
||||
package chat
|
||||
|
||||
import (
|
||||
"context"
|
||||
"slices"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
tea "github.com/charmbracelet/bubbletea"
|
||||
|
||||
coreagent "github.com/ollama/ollama/agent"
|
||||
"github.com/ollama/ollama/api"
|
||||
apptui "github.com/ollama/ollama/cmd/tui"
|
||||
)
|
||||
|
||||
func TestChatModelCommandOpensPicker(t *testing.T) {
|
||||
m := chatModel{
|
||||
ctx: context.Background(),
|
||||
input: []rune("/model"),
|
||||
width: 100,
|
||||
height: 20,
|
||||
opts: Options{
|
||||
Model: "llama3.2",
|
||||
ContextWindowTokens: 131072,
|
||||
ModelOptions: func(context.Context) ([]ModelOption, error) {
|
||||
return []ModelOption{
|
||||
{Name: "kimi-k2.6:cloud", Description: "cloud coding"},
|
||||
{Name: "llama3.2", Description: "local"},
|
||||
}, nil
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
updated, cmd := m.handleSubmit()
|
||||
if cmd != nil {
|
||||
t.Fatal("model command should not return a command")
|
||||
}
|
||||
m = updated.(chatModel)
|
||||
if m.modelPicker == nil {
|
||||
t.Fatal("model picker was not opened")
|
||||
}
|
||||
view := stripANSI(m.View())
|
||||
if !strings.Contains(view, "Select model") ||
|
||||
!strings.Contains(view, "Type to filter") ||
|
||||
!strings.Contains(view, "kimi-k2.6:cloud") ||
|
||||
!strings.Contains(view, "llama3.2") {
|
||||
t.Fatalf("model picker view missing content: %q", view)
|
||||
}
|
||||
if strings.Contains(view, "Search...") {
|
||||
t.Fatalf("model picker should render inline without full search box: %q", view)
|
||||
}
|
||||
if strings.Contains(view, "local") || strings.Contains(view, "cloud coding") {
|
||||
t.Fatalf("inline model picker should stay compact without descriptions: %q", view)
|
||||
}
|
||||
if !strings.Contains(view, "│ █") {
|
||||
t.Fatalf("inline model picker should keep input box visible: %q", view)
|
||||
}
|
||||
}
|
||||
|
||||
func TestChatModelCommandShowsRecommendedFirstWithoutSections(t *testing.T) {
|
||||
m := chatModel{
|
||||
ctx: context.Background(),
|
||||
input: []rune("/model"),
|
||||
width: 100,
|
||||
height: 20,
|
||||
opts: Options{
|
||||
Model: "llama3.2",
|
||||
ModelOptions: func(context.Context) ([]ModelOption, error) {
|
||||
return []ModelOption{
|
||||
{Name: "llama3.2", Description: "selected local"},
|
||||
{Name: "gemma4", Description: "local"},
|
||||
{Name: "glm-5.2:cloud", Description: "recommended cloud", Recommended: true, Cloud: true},
|
||||
{Name: "kimi-k2.7-code:cloud", Description: "another recommended cloud", Recommended: true, Cloud: true},
|
||||
}, nil
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
updated, cmd := m.handleSubmit()
|
||||
if cmd != nil {
|
||||
t.Fatal("model command should not return a command")
|
||||
}
|
||||
view := stripANSI(updated.(chatModel).View())
|
||||
for _, unwanted := range []string{"Recommended", "More", "recommended cloud", "selected local"} {
|
||||
if strings.Contains(view, unwanted) {
|
||||
t.Fatalf("compact model picker should be flat and description-free; found %q in %q", unwanted, view)
|
||||
}
|
||||
}
|
||||
firstRecommended := strings.Index(view, "glm-5.2:cloud")
|
||||
secondRecommended := strings.Index(view, "kimi-k2.7-code:cloud")
|
||||
current := strings.Index(view, "llama3.2")
|
||||
local := strings.Index(view, "gemma4")
|
||||
if firstRecommended < 0 || secondRecommended < 0 || current < 0 || local < 0 {
|
||||
t.Fatalf("compact model picker missing expected models: %q", view)
|
||||
}
|
||||
if !(firstRecommended < current && secondRecommended < current && current < local) {
|
||||
t.Fatalf("compact model picker order should be recommended, current, local: %q", view)
|
||||
}
|
||||
}
|
||||
|
||||
func TestChatModelCommandOpensSmallPicker(t *testing.T) {
|
||||
m := chatModel{
|
||||
ctx: context.Background(),
|
||||
input: []rune("/model"),
|
||||
width: 100,
|
||||
height: 24,
|
||||
opts: Options{
|
||||
Model: "model-1",
|
||||
ModelOptions: func(context.Context) ([]ModelOption, error) {
|
||||
return []ModelOption{
|
||||
{Name: "model-1"},
|
||||
{Name: "model-2"},
|
||||
{Name: "model-3"},
|
||||
{Name: "model-4"},
|
||||
{Name: "model-5"},
|
||||
{Name: "model-6"},
|
||||
{Name: "model-7"},
|
||||
}, nil
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
updated, cmd := m.handleSubmit()
|
||||
if cmd != nil {
|
||||
t.Fatal("model command should not return a command")
|
||||
}
|
||||
m = updated.(chatModel)
|
||||
view := stripANSI(m.View())
|
||||
for _, want := range []string{"model-1", "model-5", "... and 2 more"} {
|
||||
if !strings.Contains(view, want) {
|
||||
t.Fatalf("small model picker missing %q: %q", want, view)
|
||||
}
|
||||
}
|
||||
if strings.Contains(view, "model-6") || strings.Contains(view, "model-7") {
|
||||
t.Fatalf("small model picker rendered too many items: %q", view)
|
||||
}
|
||||
}
|
||||
|
||||
func TestChatModelPickerStaysInlineWhenSmall(t *testing.T) {
|
||||
m := chatModel{
|
||||
ctx: context.Background(),
|
||||
input: []rune("/model"),
|
||||
width: 44,
|
||||
height: 10,
|
||||
opts: Options{
|
||||
Model: "llama3.2",
|
||||
ModelOptions: func(context.Context) ([]ModelOption, error) {
|
||||
return []ModelOption{
|
||||
{Name: "kimi-k2.6:cloud", Description: "cloud coding"},
|
||||
{Name: "llama3.2", Description: "local"},
|
||||
}, nil
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
updated, cmd := m.handleSubmit()
|
||||
if cmd != nil {
|
||||
t.Fatal("model command should not return a command")
|
||||
}
|
||||
m = updated.(chatModel)
|
||||
view := stripANSI(m.View())
|
||||
if !strings.Contains(view, "Select model") || !strings.Contains(view, "Type to filter") {
|
||||
t.Fatalf("small model picker should stay inline: %q", view)
|
||||
}
|
||||
if strings.Contains(view, "Search...") {
|
||||
t.Fatalf("small model picker should not use bespoke full-frame search: %q", view)
|
||||
}
|
||||
}
|
||||
|
||||
func TestChatModelPickerShowsRecommendedModelsFirst(t *testing.T) {
|
||||
models := normalizeModelOptions([]ModelOption{
|
||||
{Name: "llama3.2", Description: "local"},
|
||||
{Name: "kimi-k2.6:cloud", Description: "cloud coding", Recommended: true},
|
||||
{Name: "qwen3.5:cloud", Description: "cloud reasoning", Recommended: true},
|
||||
{Name: "gemma4", Description: "local"},
|
||||
})
|
||||
got := make([]string, 0, len(models))
|
||||
for _, model := range models {
|
||||
got = append(got, model.Name)
|
||||
}
|
||||
want := []string{"kimi-k2.6:cloud", "qwen3.5:cloud", "llama3.2", "gemma4"}
|
||||
if !slices.Equal(got, want) {
|
||||
t.Fatalf("model order = %#v, want %#v", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestChatModelPickerPinsCurrentThenRecommendedModels(t *testing.T) {
|
||||
models := normalizeModelOptions([]ModelOption{
|
||||
{Name: "llama3.2", Description: "local"},
|
||||
{Name: "glm-5.2:cloud", Description: "cloud selected", Recommended: true, Cloud: true},
|
||||
{Name: "kimi-k2.7-code:cloud", Description: "cloud coding", Recommended: true, Cloud: true},
|
||||
{Name: "gemma4", Description: "local"},
|
||||
})
|
||||
|
||||
items := modelSelectorItems(models, "glm-5.2:cloud")
|
||||
got := make([]string, 0, len(items))
|
||||
for _, item := range items {
|
||||
got = append(got, item.Name)
|
||||
}
|
||||
want := []string{"glm-5.2:cloud", "kimi-k2.7-code:cloud", "llama3.2", "gemma4"}
|
||||
if !slices.Equal(got, want) {
|
||||
t.Fatalf("selector item order = %#v, want %#v", got, want)
|
||||
}
|
||||
for _, item := range items[:3] {
|
||||
if !item.Recommended {
|
||||
t.Fatalf("%q should be pinned in the first picker section", item.Name)
|
||||
}
|
||||
}
|
||||
if items[0].Description != "cloud selected" {
|
||||
t.Fatalf("current model description = %q, want plain model description", items[0].Description)
|
||||
}
|
||||
}
|
||||
|
||||
func TestInitialModelPickerRendersBeforeChatShell(t *testing.T) {
|
||||
models := normalizeModelOptions([]ModelOption{
|
||||
{Name: "glm-5.2:cloud", Description: "cloud selected", Recommended: true, Cloud: true},
|
||||
{Name: "llama3.2", Description: "local"},
|
||||
})
|
||||
picker := apptui.NewModelSelectorModel("Select model", modelSelectorItems(models, "glm-5.2:cloud"), "glm-5.2:cloud", "")
|
||||
m := chatModel{
|
||||
width: 100,
|
||||
height: 20,
|
||||
openModelOnInit: true,
|
||||
modelPicker: &picker,
|
||||
entries: []chatEntry{{role: "assistant", content: "old chat content"}},
|
||||
}
|
||||
|
||||
view := stripANSI(m.View())
|
||||
if !strings.Contains(view, "Select model") || !strings.Contains(view, "llama3.2") {
|
||||
t.Fatalf("initial picker view missing model content: %q", view)
|
||||
}
|
||||
if strings.Contains(view, "old chat content") || strings.Contains(view, "│ █") {
|
||||
t.Fatalf("initial picker should render before chat shell: %q", view)
|
||||
}
|
||||
}
|
||||
|
||||
func TestChatModelPickerRanksClosestFilteredModelFirst(t *testing.T) {
|
||||
models := normalizeModelOptions([]ModelOption{
|
||||
{Name: "gemma3:27b", Description: "recommended but longer", Recommended: true},
|
||||
{Name: "llama3.2", Description: "mentions gemm in description"},
|
||||
{Name: "gemma4:27b", Description: "longer local"},
|
||||
{Name: "gemma4", Description: "short local"},
|
||||
})
|
||||
picker := apptui.NewModelSelectorModel("Select model", modelSelectorItems(models, ""), "", "gemm")
|
||||
|
||||
filtered := picker.FilteredItems()
|
||||
got := make([]string, 0, len(filtered))
|
||||
for _, model := range filtered {
|
||||
got = append(got, model.Name)
|
||||
}
|
||||
want := []string{"gemma4", "gemma3:27b", "gemma4:27b", "llama3.2"}
|
||||
if !slices.Equal(got, want) {
|
||||
t.Fatalf("filtered model order = %#v, want %#v", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestChatModelPickerFiltersAndSwitchesModel(t *testing.T) {
|
||||
var savedModel string
|
||||
originalMessages := []api.Message{{Role: "user", Content: "keep me"}}
|
||||
m := chatModel{
|
||||
ctx: context.Background(),
|
||||
chatID: "chat-1",
|
||||
input: []rune("/model qwen"),
|
||||
width: 100,
|
||||
height: 20,
|
||||
messages: slices.Clone(originalMessages),
|
||||
opts: Options{
|
||||
Model: "llama3.2",
|
||||
ModelOptions: func(context.Context) ([]ModelOption, error) {
|
||||
return []ModelOption{
|
||||
{Name: "llama3.2", Description: "local"},
|
||||
{Name: "qwen3.5:cloud", Description: "cloud reasoning"},
|
||||
}, nil
|
||||
},
|
||||
ToolRegistryForModel: func(ctx context.Context, model string) *coreagent.Registry {
|
||||
if model != "qwen3.5:cloud" {
|
||||
t.Fatalf("tool registry model = %q, want qwen3.5:cloud", model)
|
||||
}
|
||||
registry := &coreagent.Registry{}
|
||||
registry.Register(chatTestTool{})
|
||||
return registry
|
||||
},
|
||||
ContextWindowTokensForModel: func(ctx context.Context, model string, fallback int) int {
|
||||
if model != "qwen3.5:cloud" {
|
||||
t.Fatalf("context model = %q, want qwen3.5:cloud", model)
|
||||
}
|
||||
if fallback != 0 {
|
||||
t.Fatalf("context fallback = %d, want 0 after model switch", fallback)
|
||||
}
|
||||
return 262144
|
||||
},
|
||||
SystemPromptForModel: func(ctx context.Context, model string, registry *coreagent.Registry, toolsDisabled bool) string {
|
||||
if model != "qwen3.5:cloud" {
|
||||
t.Fatalf("system prompt model = %q, want qwen3.5:cloud", model)
|
||||
}
|
||||
if registry == nil {
|
||||
t.Fatalf("system prompt registry missing fake tool: %#v", registry)
|
||||
}
|
||||
if _, ok := registry.Get("fake_tool"); !ok {
|
||||
t.Fatalf("system prompt registry missing fake tool: %#v", registry)
|
||||
}
|
||||
return "system for " + model
|
||||
},
|
||||
OnModelSelected: func(ctx context.Context, model string) error {
|
||||
savedModel = model
|
||||
return nil
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
updated, cmd := m.handleSubmit()
|
||||
if cmd != nil {
|
||||
t.Fatal("model command should not return a command")
|
||||
}
|
||||
m = updated.(chatModel)
|
||||
if m.modelPicker == nil || m.modelPicker.Filter() != "qwen" {
|
||||
t.Fatalf("model picker = %#v, want qwen filter", m.modelPicker)
|
||||
}
|
||||
if view := stripANSI(m.View()); !strings.Contains(view, "qwen3.5:cloud") || strings.Contains(view, "llama3.2") {
|
||||
t.Fatalf("filtered model picker view = %q", view)
|
||||
}
|
||||
|
||||
updated, cmd = m.Update(tea.KeyMsg{Type: tea.KeyEnter})
|
||||
m = updated.(chatModel)
|
||||
if cmd != nil {
|
||||
t.Fatal("switching models should not start a command")
|
||||
}
|
||||
if m.modelPicker != nil {
|
||||
t.Fatal("model picker should close after selection")
|
||||
}
|
||||
if m.status != "ready" || m.notificationLine() != "" {
|
||||
t.Fatalf("model switch should not show action status, status=%q notification=%q", m.status, m.notificationLine())
|
||||
}
|
||||
if m.opts.Model != "qwen3.5:cloud" {
|
||||
t.Fatalf("model = %q, want qwen3.5:cloud", m.opts.Model)
|
||||
}
|
||||
if m.chatID != "chat-1" {
|
||||
t.Fatalf("chatID = %q, want chat-1", m.chatID)
|
||||
}
|
||||
if len(m.messages) != len(originalMessages) || m.messages[0].Content != originalMessages[0].Content {
|
||||
t.Fatalf("messages changed on model switch: %#v", m.messages)
|
||||
}
|
||||
if len(m.entries) != 0 {
|
||||
t.Fatalf("model switch should not append transcript entries: %#v", m.entries)
|
||||
}
|
||||
if savedModel != "qwen3.5:cloud" {
|
||||
t.Fatalf("saved model = %q, want qwen3.5:cloud", savedModel)
|
||||
}
|
||||
if m.opts.Tools == nil {
|
||||
t.Fatalf("tools registry was not rebuilt for model: %#v", m.opts.Tools)
|
||||
}
|
||||
if _, ok := m.opts.Tools.Get("fake_tool"); !ok {
|
||||
t.Fatalf("tools registry was not rebuilt for model: %#v", m.opts.Tools)
|
||||
}
|
||||
if m.opts.ContextWindowTokens != 262144 {
|
||||
t.Fatalf("context window = %d, want 262144", m.opts.ContextWindowTokens)
|
||||
}
|
||||
if m.opts.SystemPrompt != "system for qwen3.5:cloud" {
|
||||
t.Fatalf("system prompt = %q", m.opts.SystemPrompt)
|
||||
}
|
||||
}
|
||||
|
||||
func TestChatModelSelectionStartsBackgroundPreload(t *testing.T) {
|
||||
m := chatModel{
|
||||
ctx: context.Background(),
|
||||
opts: Options{
|
||||
Model: "llama3.2",
|
||||
PreloadModel: func(context.Context, string, *api.ThinkValue) (int, error) {
|
||||
return 0, nil
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
if err := m.applyModelSelection("qwen3", false); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
cmd := m.startModelPreload("qwen3")
|
||||
if cmd == nil {
|
||||
t.Fatal("model switch should start background preload when configured")
|
||||
}
|
||||
if m.preloadingModel != "qwen3" {
|
||||
t.Fatalf("preloadingModel = %q, want qwen3", m.preloadingModel)
|
||||
}
|
||||
}
|
||||
|
||||
func TestChatModelSwitchNextRunKeepsHistory(t *testing.T) {
|
||||
client := &chatCaptureClient{}
|
||||
history := []api.Message{
|
||||
{Role: "user", Content: "old question"},
|
||||
{Role: "assistant", Content: "old answer"},
|
||||
}
|
||||
m := chatModel{
|
||||
ctx: context.Background(),
|
||||
chatID: "chat-1",
|
||||
messages: slices.Clone(history),
|
||||
input: []rune("continue"),
|
||||
opts: Options{
|
||||
Model: "llama3.2",
|
||||
Client: client,
|
||||
SystemPromptForModel: func(_ context.Context, model string, _ *coreagent.Registry, _ bool) string {
|
||||
return "system for " + model
|
||||
},
|
||||
},
|
||||
}
|
||||
if err := m.applyModelSelection("qwen3", true); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
updated, cmd := m.handleSubmit()
|
||||
m = updated.(chatModel)
|
||||
if cmd == nil {
|
||||
t.Fatal("next prompt should start a model run")
|
||||
}
|
||||
done := waitForRunDone(t, m.events)
|
||||
if done.err != nil {
|
||||
t.Fatal(done.err)
|
||||
}
|
||||
|
||||
if len(client.requests) != 1 {
|
||||
t.Fatalf("requests = %d, want 1", len(client.requests))
|
||||
}
|
||||
req := client.requests[0]
|
||||
if req.Model != "qwen3" {
|
||||
t.Fatalf("request model = %q, want qwen3", req.Model)
|
||||
}
|
||||
if len(req.Messages) != 4 {
|
||||
t.Fatalf("request messages = %#v, want system + 2 history + new user", req.Messages)
|
||||
}
|
||||
if req.Messages[0].Role != "system" || req.Messages[0].Content != "system for qwen3" {
|
||||
t.Fatalf("system message = %#v", req.Messages[0])
|
||||
}
|
||||
for i, want := range history {
|
||||
got := req.Messages[i+1]
|
||||
if got.Role != want.Role || got.Content != want.Content {
|
||||
t.Fatalf("history message %d = %#v, want %#v", i, got, want)
|
||||
}
|
||||
}
|
||||
if req.Messages[3].Role != "user" || req.Messages[3].Content != "continue" {
|
||||
t.Fatalf("new user message = %#v", req.Messages[3])
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,307 @@
|
||||
package chat
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/url"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
tea "github.com/charmbracelet/bubbletea"
|
||||
)
|
||||
|
||||
func TestChatStartRunAttachesDroppedImagePath(t *testing.T) {
|
||||
fp := writeTestPNG(t)
|
||||
m := chatModel{
|
||||
ctx: context.Background(),
|
||||
opts: Options{
|
||||
Model: "test",
|
||||
Client: chatTestClient{},
|
||||
MultiModal: true,
|
||||
},
|
||||
}
|
||||
|
||||
updated, cmd := m.startRun("describe " + fp)
|
||||
m = updated.(chatModel)
|
||||
|
||||
if cmd == nil {
|
||||
t.Fatal("startRun should return a command")
|
||||
}
|
||||
if len(m.liveMessages) != 1 {
|
||||
t.Fatalf("liveMessages = %d, want 1", len(m.liveMessages))
|
||||
}
|
||||
if got := m.liveMessages[0].Content; got != "describe" {
|
||||
t.Fatalf("content = %q, want describe", got)
|
||||
}
|
||||
if got := len(m.liveMessages[0].Images); got != 1 {
|
||||
t.Fatalf("images = %d, want 1", got)
|
||||
}
|
||||
if len(m.entries) == 0 {
|
||||
t.Fatal("missing user transcript entry")
|
||||
}
|
||||
entry := m.entries[0].content
|
||||
if strings.Contains(entry, fp) {
|
||||
t.Fatalf("transcript entry should hide local file path: %q", entry)
|
||||
}
|
||||
if !strings.Contains(entry, "describe") || !strings.Contains(entry, "[attached 1 file]") {
|
||||
t.Fatalf("transcript entry = %q, want prompt plus attachment note", entry)
|
||||
}
|
||||
}
|
||||
|
||||
func TestChatStartRunAttachesDroppedFileURL(t *testing.T) {
|
||||
fp := writeTestPNG(t)
|
||||
fileURL := (&url.URL{Scheme: "file", Path: fp}).String()
|
||||
m := chatModel{
|
||||
ctx: context.Background(),
|
||||
opts: Options{
|
||||
Model: "test",
|
||||
Client: chatTestClient{},
|
||||
MultiModal: true,
|
||||
},
|
||||
}
|
||||
|
||||
updated, _ := m.startRun(fileURL)
|
||||
m = updated.(chatModel)
|
||||
|
||||
if got := m.liveMessages[0].Content; got != "" {
|
||||
t.Fatalf("content = %q, want empty prompt after extracting file URL", got)
|
||||
}
|
||||
if got := len(m.liveMessages[0].Images); got != 1 {
|
||||
t.Fatalf("images = %d, want 1", got)
|
||||
}
|
||||
if got := m.entries[0].content; got != "[attached 1 file]" {
|
||||
t.Fatalf("transcript entry = %q, want attachment-only note", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestChatPasteImagePathAttachesOnSubmit(t *testing.T) {
|
||||
fp := writeTestPNG(t)
|
||||
m := chatModel{
|
||||
ctx: context.Background(),
|
||||
opts: Options{
|
||||
Model: "test",
|
||||
Client: chatTestClient{},
|
||||
MultiModal: true,
|
||||
},
|
||||
}
|
||||
|
||||
updated, _ := m.Update(tea.KeyMsg{Type: tea.KeyRunes, Runes: []rune("describe " + fp), Paste: true})
|
||||
m = updated.(chatModel)
|
||||
if got := string(m.input); got != "describe [Image #0]" {
|
||||
t.Fatalf("pasted path input = %q, want placeholder", got)
|
||||
}
|
||||
if got := m.notificationLine(); got != "" {
|
||||
t.Fatalf("notification = %q, want no attachment notification", got)
|
||||
}
|
||||
if got := string(m.input); strings.Contains(got, fp) {
|
||||
t.Fatalf("pasted path should be hidden behind placeholder, input = %q", got)
|
||||
}
|
||||
if completions := m.slashCompletions(); len(completions) != 0 {
|
||||
t.Fatalf("placeholder input should not show slash completions: %#v", completions)
|
||||
}
|
||||
|
||||
updated, cmd := m.Update(tea.KeyMsg{Type: tea.KeyEnter})
|
||||
m = updated.(chatModel)
|
||||
if cmd == nil {
|
||||
t.Fatal("submit should start a run")
|
||||
}
|
||||
if got := m.liveMessages[0].Content; got != "describe [Image #0]" {
|
||||
t.Fatalf("content = %q, want prompt with placeholder", got)
|
||||
}
|
||||
if got := len(m.liveMessages[0].Images); got != 1 {
|
||||
t.Fatalf("images = %d, want 1", got)
|
||||
}
|
||||
if strings.Contains(m.entries[0].content, fp) {
|
||||
t.Fatalf("transcript entry should hide pasted file path: %q", m.entries[0].content)
|
||||
}
|
||||
if !strings.Contains(m.entries[0].content, "[Image #0]") {
|
||||
t.Fatalf("transcript entry should show placeholder: %q", m.entries[0].content)
|
||||
}
|
||||
}
|
||||
|
||||
func TestChatPasteImagePathAfterSwitchingToMultimodalModel(t *testing.T) {
|
||||
fp := writeTestPNG(t)
|
||||
m := chatModel{
|
||||
ctx: context.Background(),
|
||||
input: []rune("/model vision"),
|
||||
opts: Options{
|
||||
Model: "text",
|
||||
ModelOptions: func(context.Context) ([]ModelOption, error) {
|
||||
return []ModelOption{
|
||||
{Name: "text", Description: "local"},
|
||||
{Name: "vision", Description: "local vision"},
|
||||
}, nil
|
||||
},
|
||||
MultiModalForModel: func(_ context.Context, model string) bool {
|
||||
return model == "vision"
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
updated, cmd := m.handleSubmit()
|
||||
if cmd != nil {
|
||||
t.Fatal("model picker should not return a command")
|
||||
}
|
||||
m = updated.(chatModel)
|
||||
|
||||
updated, cmd = m.Update(tea.KeyMsg{Type: tea.KeyEnter})
|
||||
if cmd != nil {
|
||||
t.Fatal("model switch should not preload without a preload hook")
|
||||
}
|
||||
m = updated.(chatModel)
|
||||
if m.opts.Model != "vision" {
|
||||
t.Fatalf("model = %q, want vision", m.opts.Model)
|
||||
}
|
||||
if !m.opts.MultiModal {
|
||||
t.Fatal("switching to a multimodal model should enable image paste handling")
|
||||
}
|
||||
|
||||
updated, _ = m.Update(tea.KeyMsg{Type: tea.KeyRunes, Runes: []rune("describe " + fp), Paste: true})
|
||||
m = updated.(chatModel)
|
||||
if got := string(m.input); got != "describe [Image #0]" {
|
||||
t.Fatalf("pasted path input = %q, want placeholder", got)
|
||||
}
|
||||
if strings.Contains(string(m.input), fp) {
|
||||
t.Fatalf("pasted path should be hidden behind placeholder, input = %q", string(m.input))
|
||||
}
|
||||
}
|
||||
|
||||
func TestChatImagePlaceholdersUseSessionNumbers(t *testing.T) {
|
||||
first := writeTestPNG(t)
|
||||
second := writeTestPNG(t)
|
||||
m := chatModel{
|
||||
ctx: context.Background(),
|
||||
opts: Options{
|
||||
Model: "test",
|
||||
Client: chatTestClient{},
|
||||
MultiModal: true,
|
||||
},
|
||||
}
|
||||
|
||||
updated, _ := m.Update(tea.KeyMsg{Type: tea.KeyRunes, Runes: []rune(first), Paste: true})
|
||||
m = updated.(chatModel)
|
||||
if got := string(m.input); got != "[Image #0]" {
|
||||
t.Fatalf("first placeholder = %q, want [Image #0]", got)
|
||||
}
|
||||
|
||||
updated, cmd := m.Update(tea.KeyMsg{Type: tea.KeyEnter})
|
||||
m = updated.(chatModel)
|
||||
if cmd == nil {
|
||||
t.Fatal("submit should start a run")
|
||||
}
|
||||
|
||||
updated, _ = m.Update(tea.KeyMsg{Type: tea.KeyRunes, Runes: []rune(second), Paste: true})
|
||||
m = updated.(chatModel)
|
||||
if got := string(m.input); got != "[Image #1]" {
|
||||
t.Fatalf("second placeholder = %q, want [Image #1]", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestChatAbsoluteImagePathBypassesSlashCommandParsing(t *testing.T) {
|
||||
fp := writeTestPNG(t)
|
||||
m := chatModel{
|
||||
ctx: context.Background(),
|
||||
opts: Options{
|
||||
Model: "test",
|
||||
Client: chatTestClient{},
|
||||
MultiModal: true,
|
||||
},
|
||||
}
|
||||
|
||||
m.input = []rune(fp)
|
||||
m.inputCursor = len(m.input)
|
||||
m.inputCursorSet = true
|
||||
updated, cmd := m.handleSubmit()
|
||||
m = updated.(chatModel)
|
||||
|
||||
if cmd == nil {
|
||||
t.Fatal("absolute image path should start a run instead of being parsed as a slash command")
|
||||
}
|
||||
if got := len(m.liveMessages[0].Images); got != 1 {
|
||||
t.Fatalf("images = %d, want 1", got)
|
||||
}
|
||||
if got := m.entries[0].role; got != "user" {
|
||||
t.Fatalf("entry role = %q, want user", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestChatDeletingImagePlaceholderRemovesAttachment(t *testing.T) {
|
||||
fp := writeTestPNG(t)
|
||||
m := chatModel{
|
||||
ctx: context.Background(),
|
||||
opts: Options{
|
||||
Model: "test",
|
||||
Client: chatTestClient{},
|
||||
MultiModal: true,
|
||||
},
|
||||
}
|
||||
|
||||
updated, _ := m.Update(tea.KeyMsg{Type: tea.KeyRunes, Runes: []rune("describe " + fp), Paste: true})
|
||||
m = updated.(chatModel)
|
||||
if got := len(m.inputAttachments); got != 1 {
|
||||
t.Fatalf("input attachments = %d, want 1", got)
|
||||
}
|
||||
|
||||
updated, _ = m.Update(tea.KeyMsg{Type: tea.KeyBackspace})
|
||||
m = updated.(chatModel)
|
||||
if got := string(m.input); got != "describe " {
|
||||
t.Fatalf("input after backspace = %q, want image placeholder removed", got)
|
||||
}
|
||||
if got := len(m.inputAttachments); got != 0 {
|
||||
t.Fatalf("input attachments after editing placeholder = %d, want 0", got)
|
||||
}
|
||||
|
||||
m.input = []rune("describe")
|
||||
m.inputCursor = len(m.input)
|
||||
m.inputCursorSet = true
|
||||
updated, cmd := m.Update(tea.KeyMsg{Type: tea.KeyEnter})
|
||||
m = updated.(chatModel)
|
||||
if cmd == nil {
|
||||
t.Fatal("submit should start a run")
|
||||
}
|
||||
if got := len(m.liveMessages[0].Images); got != 0 {
|
||||
t.Fatalf("images = %d, want 0 after deleting placeholder", got)
|
||||
}
|
||||
if got := m.liveMessages[0].Content; got != "describe" {
|
||||
t.Fatalf("content = %q, want describe", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestChatWordDeletingImagePlaceholderRemovesAttachment(t *testing.T) {
|
||||
fp := writeTestPNG(t)
|
||||
m := chatModel{
|
||||
ctx: context.Background(),
|
||||
opts: Options{
|
||||
Model: "test",
|
||||
Client: chatTestClient{},
|
||||
MultiModal: true,
|
||||
},
|
||||
}
|
||||
|
||||
updated, _ := m.Update(tea.KeyMsg{Type: tea.KeyRunes, Runes: []rune("describe " + fp), Paste: true})
|
||||
m = updated.(chatModel)
|
||||
updated, _ = m.Update(tea.KeyMsg{Type: tea.KeySpace})
|
||||
m = updated.(chatModel)
|
||||
updated, _ = m.Update(tea.KeyMsg{Type: tea.KeyBackspace, Alt: true})
|
||||
m = updated.(chatModel)
|
||||
|
||||
if got := string(m.input); got != "describe " {
|
||||
t.Fatalf("input after word backspace = %q, want image placeholder removed", got)
|
||||
}
|
||||
if got := len(m.inputAttachments); got != 0 {
|
||||
t.Fatalf("input attachments after word backspace = %d, want 0", got)
|
||||
}
|
||||
}
|
||||
|
||||
func writeTestPNG(t *testing.T) string {
|
||||
t.Helper()
|
||||
dir := t.TempDir()
|
||||
fp := filepath.Join(dir, "dragged image.png")
|
||||
data := make([]byte, 600)
|
||||
copy(data, []byte{0x89, 'P', 'N', 'G', '\r', '\n', 0x1a, '\n'})
|
||||
if err := os.WriteFile(fp, data, 0o600); err != nil {
|
||||
t.Fatalf("failed to write test image: %v", err)
|
||||
}
|
||||
return fp
|
||||
}
|
||||
File diff suppressed because it is too large.
Load diff
File diff suppressed because it is too large.
Load diff
@@ -0,0 +1,113 @@
|
||||
package chat
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"regexp"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
tea "github.com/charmbracelet/bubbletea"
|
||||
|
||||
coreagent "github.com/ollama/ollama/agent"
|
||||
"github.com/ollama/ollama/api"
|
||||
)
|
||||
|
||||
type chatTestTool struct{}
|
||||
|
||||
type chatTestClient struct{}
|
||||
|
||||
type chatCaptureClient struct {
|
||||
requests []*api.ChatRequest
|
||||
}
|
||||
|
||||
type chatToolLoopClient struct {
|
||||
calls int
|
||||
toolRounds int
|
||||
}
|
||||
|
||||
func (chatTestClient) Chat(ctx context.Context, req *api.ChatRequest, fn api.ChatResponseFunc) error {
|
||||
if err := ctx.Err(); err != nil {
|
||||
return err
|
||||
}
|
||||
return fn(api.ChatResponse{
|
||||
Message: api.Message{Role: "assistant", Content: "ok"},
|
||||
Done: true,
|
||||
})
|
||||
}
|
||||
|
||||
func (c *chatCaptureClient) Chat(ctx context.Context, req *api.ChatRequest, fn api.ChatResponseFunc) error {
|
||||
if err := ctx.Err(); err != nil {
|
||||
return err
|
||||
}
|
||||
c.requests = append(c.requests, req)
|
||||
return fn(api.ChatResponse{
|
||||
Message: api.Message{Role: "assistant", Content: "ok"},
|
||||
Done: true,
|
||||
})
|
||||
}
|
||||
|
||||
func (c *chatToolLoopClient) Chat(ctx context.Context, req *api.ChatRequest, fn api.ChatResponseFunc) error {
|
||||
if err := ctx.Err(); err != nil {
|
||||
return err
|
||||
}
|
||||
c.calls++
|
||||
if c.calls > c.toolRounds {
|
||||
return fn(api.ChatResponse{Message: api.Message{Role: "assistant", Content: "done"}, Done: true})
|
||||
}
|
||||
|
||||
args := api.NewToolCallFunctionArguments()
|
||||
args.Set("value", "keep going")
|
||||
return fn(api.ChatResponse{Message: api.Message{Role: "assistant", ToolCalls: []api.ToolCall{{
|
||||
ID: fmt.Sprintf("call-%d", c.calls),
|
||||
Function: api.ToolCallFunction{
|
||||
Name: "fake_tool",
|
||||
Arguments: args,
|
||||
},
|
||||
}}}})
|
||||
}
|
||||
|
||||
func (chatTestTool) Name() string {
|
||||
return "fake_tool"
|
||||
}
|
||||
|
||||
func (chatTestTool) Description() string {
|
||||
return "does test work"
|
||||
}
|
||||
|
||||
func (chatTestTool) Schema() api.ToolFunction {
|
||||
return api.ToolFunction{
|
||||
Name: "fake_tool",
|
||||
Description: "does test work",
|
||||
Parameters: api.ToolFunctionParameters{
|
||||
Type: "object",
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func (chatTestTool) Execute(context.Context, coreagent.ToolContext, map[string]any) (coreagent.ToolResult, error) {
|
||||
return coreagent.ToolResult{Content: "ok"}, nil
|
||||
}
|
||||
|
||||
func waitForRunDone(t *testing.T, events <-chan tea.Msg) chatRunDoneMsg {
|
||||
t.Helper()
|
||||
timeout := time.After(2 * time.Second)
|
||||
for {
|
||||
select {
|
||||
case msg, ok := <-events:
|
||||
if !ok {
|
||||
t.Fatal("events closed before run done")
|
||||
}
|
||||
if done, ok := msg.(chatRunDoneMsg); ok {
|
||||
return done
|
||||
}
|
||||
case <-timeout:
|
||||
t.Fatal("timed out waiting for run done")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func stripANSI(s string) string {
|
||||
re := regexp.MustCompile(`\x1b\[[0-9;:]*[A-Za-z]`)
|
||||
return re.ReplaceAllString(s, "")
|
||||
}
|
||||
@@ -0,0 +1,126 @@
|
||||
package chat
|
||||
|
||||
import "github.com/charmbracelet/lipgloss"
|
||||
|
||||
const (
|
||||
chatAnsiRed = "1"
|
||||
chatAnsiGreen = "2"
|
||||
chatAnsiYellow = "3"
|
||||
chatAnsiBlue = "4"
|
||||
chatAnsiCyan = "6"
|
||||
chatAnsiBrightBlack = "8"
|
||||
)
|
||||
|
||||
var (
|
||||
chatHeaderStyle = lipgloss.NewStyle().
|
||||
Bold(true)
|
||||
|
||||
chatMetaStyle = lipgloss.NewStyle().
|
||||
Faint(true)
|
||||
|
||||
chatFooterStyle = lipgloss.NewStyle().
|
||||
Faint(true)
|
||||
|
||||
chatInputBorderStyle = lipgloss.NewStyle().
|
||||
Faint(true)
|
||||
|
||||
chatInputPlaceholderStyle = lipgloss.NewStyle().
|
||||
Foreground(lipgloss.Color("8"))
|
||||
|
||||
chatCursorStyle = lipgloss.NewStyle().
|
||||
Reverse(true)
|
||||
|
||||
chatBlankCursorStyle = lipgloss.NewStyle().
|
||||
Faint(true)
|
||||
|
||||
chatNotificationStyle = chatMetaStyle
|
||||
|
||||
chatUserStyle = lipgloss.NewStyle()
|
||||
|
||||
chatUserBlockStyle = lipgloss.NewStyle().
|
||||
Foreground(lipgloss.AdaptiveColor{Light: "#777777", Dark: "#8a8a8a"})
|
||||
|
||||
chatToolStyle = lipgloss.NewStyle()
|
||||
|
||||
chatInlineCodeStyle = lipgloss.NewStyle().
|
||||
Bold(true)
|
||||
|
||||
chatCodeBlockStyle = lipgloss.NewStyle()
|
||||
|
||||
chatTableBorderStyle = lipgloss.NewStyle().
|
||||
Faint(true)
|
||||
|
||||
chatToolRunningStyle = lipgloss.NewStyle().
|
||||
Foreground(lipgloss.Color(chatAnsiYellow))
|
||||
|
||||
chatToolDoneStyle = lipgloss.NewStyle().
|
||||
Foreground(lipgloss.Color(chatAnsiGreen))
|
||||
|
||||
// chatToolMixedStyle marks a tool group with both succeeded and failed
|
||||
// calls (partial success). Amber/orange is distinct from green (success),
|
||||
// red (failure), and yellow (running).
|
||||
chatToolMixedStyle = lipgloss.NewStyle().
|
||||
Foreground(lipgloss.Color("208"))
|
||||
|
||||
chatToolOutputStyle = lipgloss.NewStyle().
|
||||
Foreground(lipgloss.AdaptiveColor{Light: "#666666", Dark: "#a0a0a0"})
|
||||
|
||||
chatDiffMetaStyle = lipgloss.NewStyle().
|
||||
Faint(true)
|
||||
|
||||
chatDiffFileStyle = lipgloss.NewStyle().
|
||||
Bold(true).
|
||||
Foreground(lipgloss.Color(chatAnsiCyan))
|
||||
|
||||
chatDiffHunkStyle = lipgloss.NewStyle().
|
||||
Foreground(lipgloss.Color(chatAnsiBlue))
|
||||
|
||||
chatDiffAddStyle = lipgloss.NewStyle().
|
||||
Foreground(lipgloss.Color(chatAnsiGreen))
|
||||
|
||||
chatDiffDeleteStyle = lipgloss.NewStyle().
|
||||
Foreground(lipgloss.Color(chatAnsiRed))
|
||||
|
||||
chatErrorStyle = lipgloss.NewStyle().
|
||||
Foreground(lipgloss.Color(chatAnsiRed))
|
||||
|
||||
chatFullAccessStyle = lipgloss.NewStyle().
|
||||
Foreground(lipgloss.AdaptiveColor{Light: "#9f5f5f", Dark: "#b87373"})
|
||||
|
||||
chatCommandNameStyle = lipgloss.NewStyle()
|
||||
|
||||
chatPickerTextStyle = lipgloss.NewStyle()
|
||||
|
||||
chatPickerTitleStyle = lipgloss.NewStyle().
|
||||
Bold(true)
|
||||
|
||||
chatPickerSelectedStyle = lipgloss.NewStyle().
|
||||
Bold(true)
|
||||
|
||||
chatPickerMetaStyle = lipgloss.NewStyle().
|
||||
Faint(true)
|
||||
|
||||
chatHistoryTitleStyle = lipgloss.NewStyle().
|
||||
Bold(true)
|
||||
|
||||
chatHistorySystemRoleStyle = lipgloss.NewStyle().
|
||||
Bold(true).
|
||||
Faint(true)
|
||||
|
||||
chatHistoryUserRoleStyle = lipgloss.NewStyle().
|
||||
Bold(true).
|
||||
Foreground(lipgloss.Color(chatAnsiBlue))
|
||||
|
||||
chatHistoryAssistantRoleStyle = lipgloss.NewStyle().
|
||||
Bold(true).
|
||||
Foreground(lipgloss.Color(chatAnsiYellow))
|
||||
|
||||
chatHistoryToolRoleStyle = lipgloss.NewStyle().
|
||||
Bold(true).
|
||||
Foreground(lipgloss.Color(chatAnsiGreen))
|
||||
|
||||
chatHistoryLabelStyle = lipgloss.NewStyle().
|
||||
Faint(true)
|
||||
|
||||
chatHistoryTextStyle = lipgloss.NewStyle()
|
||||
)
|
||||
@@ -0,0 +1,166 @@
|
||||
package chat
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
tea "github.com/charmbracelet/bubbletea"
|
||||
|
||||
"github.com/ollama/ollama/api"
|
||||
)
|
||||
|
||||
type chatThinkOption struct {
|
||||
value string
|
||||
label string
|
||||
description string
|
||||
}
|
||||
|
||||
type chatThinkPicker struct {
|
||||
options []chatThinkOption
|
||||
cursor int
|
||||
}
|
||||
|
||||
var chatThinkOptions = []chatThinkOption{
|
||||
{value: "auto", label: "auto", description: "use the model default"},
|
||||
{value: "on", label: "on", description: "enable thinking"},
|
||||
{value: "off", label: "off", description: "disable thinking"},
|
||||
{value: "low", label: "low", description: "use low thinking effort"},
|
||||
{value: "medium", label: "medium", description: "use medium thinking effort"},
|
||||
{value: "high", label: "high", description: "use high thinking effort"},
|
||||
{value: "max", label: "max", description: "use maximum thinking effort"},
|
||||
}
|
||||
|
||||
func (m *chatModel) openThinkPicker() (tea.Model, tea.Cmd) {
|
||||
m.thinkPicker = newChatThinkPicker(m.opts.Think)
|
||||
m.status = "think"
|
||||
return *m, nil
|
||||
}
|
||||
|
||||
func newChatThinkPicker(current *api.ThinkValue) *chatThinkPicker {
|
||||
picker := &chatThinkPicker{options: append([]chatThinkOption(nil), chatThinkOptions...)}
|
||||
currentValue := thinkValueLabel(current)
|
||||
for i, option := range picker.options {
|
||||
if option.value == currentValue {
|
||||
picker.cursor = i
|
||||
break
|
||||
}
|
||||
}
|
||||
return picker
|
||||
}
|
||||
|
||||
func (m chatModel) updateThinkPicker(msg tea.KeyMsg) (tea.Model, tea.Cmd) {
|
||||
switch msg.Type {
|
||||
case tea.KeyCtrlC, tea.KeyEsc:
|
||||
m.thinkPicker = nil
|
||||
m.status = "ready"
|
||||
case tea.KeyEnter:
|
||||
return m.selectThinkOption()
|
||||
case tea.KeyUp:
|
||||
m.thinkPicker.move(-1)
|
||||
case tea.KeyDown:
|
||||
m.thinkPicker.move(1)
|
||||
}
|
||||
return m, nil
|
||||
}
|
||||
|
||||
func (p *chatThinkPicker) move(delta int) {
|
||||
if p == nil || len(p.options) == 0 || delta == 0 {
|
||||
return
|
||||
}
|
||||
p.cursor = clamp(p.cursor+delta, 0, len(p.options)-1)
|
||||
}
|
||||
|
||||
func (p *chatThinkPicker) selected() (chatThinkOption, bool) {
|
||||
if p == nil || len(p.options) == 0 {
|
||||
return chatThinkOption{}, false
|
||||
}
|
||||
return p.options[clamp(p.cursor, 0, len(p.options)-1)], true
|
||||
}
|
||||
|
||||
func (m chatModel) selectThinkOption() (tea.Model, tea.Cmd) {
|
||||
option, ok := m.thinkPicker.selected()
|
||||
if !ok {
|
||||
return m, nil
|
||||
}
|
||||
m.thinkPicker = nil
|
||||
return m.applyThinkValue(option.value)
|
||||
}
|
||||
|
||||
func (m *chatModel) handleThinkCommand(value string) (tea.Model, tea.Cmd) {
|
||||
return m.applyThinkValue(value)
|
||||
}
|
||||
|
||||
func (m *chatModel) applyThinkValue(value string) (tea.Model, tea.Cmd) {
|
||||
think, label, err := parseThinkValue(value)
|
||||
if err != nil {
|
||||
m.entries = append(m.entries, newChatEntry(chatEntry{role: "error", content: err.Error(), err: err.Error()}))
|
||||
m.status = "error"
|
||||
return *m, nil
|
||||
}
|
||||
m.opts.Think = think
|
||||
m.status = "think " + label
|
||||
return *m, nil
|
||||
}
|
||||
|
||||
func parseThinkValue(value string) (*api.ThinkValue, string, error) {
|
||||
switch strings.ToLower(strings.TrimSpace(value)) {
|
||||
case "", "auto", "default", "unset":
|
||||
return nil, "auto", nil
|
||||
case "on", "true", "think", "thinking":
|
||||
return &api.ThinkValue{Value: true}, "on", nil
|
||||
case "off", "false", "nothink", "no-think":
|
||||
return &api.ThinkValue{Value: false}, "off", nil
|
||||
case "low", "medium", "high", "max":
|
||||
value = strings.ToLower(strings.TrimSpace(value))
|
||||
return &api.ThinkValue{Value: value}, value, nil
|
||||
default:
|
||||
return nil, "", fmt.Errorf("Usage: /think [auto|on|off|low|medium|high|max]")
|
||||
}
|
||||
}
|
||||
|
||||
func thinkValueLabel(value *api.ThinkValue) string {
|
||||
if value == nil || value.Value == nil {
|
||||
return "auto"
|
||||
}
|
||||
switch v := value.Value.(type) {
|
||||
case bool:
|
||||
if v {
|
||||
return "on"
|
||||
}
|
||||
return "off"
|
||||
case string:
|
||||
return strings.ToLower(v)
|
||||
default:
|
||||
return "auto"
|
||||
}
|
||||
}
|
||||
|
||||
func (m chatModel) renderThinkPicker(width int) string {
|
||||
picker := m.thinkPicker
|
||||
if picker == nil {
|
||||
return ""
|
||||
}
|
||||
|
||||
var b strings.Builder
|
||||
b.WriteString(chatPickerTitleStyle.Render("Thinking mode"))
|
||||
b.WriteString("\n\n")
|
||||
for i, option := range picker.options {
|
||||
selected := i == picker.cursor
|
||||
if selected {
|
||||
b.WriteString(chatPickerSelectedStyle.Render("› " + option.label))
|
||||
} else {
|
||||
b.WriteString(" ")
|
||||
b.WriteString(chatPickerTextStyle.Render(option.label))
|
||||
}
|
||||
b.WriteByte('\n')
|
||||
b.WriteString(chatPickerMetaStyle.Render(" " + option.description))
|
||||
b.WriteByte('\n')
|
||||
if i < len(picker.options)-1 {
|
||||
b.WriteByte('\n')
|
||||
}
|
||||
}
|
||||
|
||||
b.WriteString("\n")
|
||||
b.WriteString(chatPickerMetaStyle.Render("↑/↓ navigate • enter select • esc cancel"))
|
||||
return b.String()
|
||||
}
|
||||
+2
-2
@@ -39,7 +39,7 @@ func (m confirmModel) Update(msg tea.Msg) (tea.Model, tea.Cmd) {
|
||||
wasSet := m.width > 0
|
||||
m.width = msg.Width
|
||||
if wasSet {
|
||||
return m, tea.EnterAltScreen
|
||||
return m, tea.ClearScreen
|
||||
}
|
||||
return m, nil
|
||||
|
||||
@@ -115,7 +115,7 @@ func RunConfirmWithOptions(prompt string, options ConfirmOptions) (bool, error)
|
||||
prompt: prompt,
|
||||
yesLabel: yesLabel,
|
||||
noLabel: noLabel,
|
||||
yes: true, // default to yes
|
||||
yes: options.Default != launch.ConfirmDefaultNo,
|
||||
}
|
||||
|
||||
p := tea.NewProgram(m)
|
||||
|
||||
+245
-10
@@ -2,6 +2,7 @@ package tui
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"sort"
|
||||
"strings"
|
||||
|
||||
tea "github.com/charmbracelet/bubbletea"
|
||||
@@ -63,6 +64,8 @@ type SelectItem struct {
|
||||
AvailabilityBadge string
|
||||
}
|
||||
|
||||
type SelectorModel = selectorModel
|
||||
|
||||
type selectorItemsUpdatedMsg struct {
|
||||
items []SelectItem
|
||||
}
|
||||
@@ -121,6 +124,7 @@ type selectorModel struct {
|
||||
cancelled bool
|
||||
helpText string
|
||||
width int
|
||||
rankFiltered bool
|
||||
}
|
||||
|
||||
func selectorModelWithCurrent(title string, items []SelectItem, current string) selectorModel {
|
||||
@@ -133,6 +137,21 @@ func selectorModelWithCurrent(title string, items []SelectItem, current string)
|
||||
return m
|
||||
}
|
||||
|
||||
func NewSelectorModel(title string, items []SelectItem, current string) SelectorModel {
|
||||
return selectorModelWithCurrent(title, items, current)
|
||||
}
|
||||
|
||||
func NewModelSelectorModel(title string, items []SelectItem, current, filter string) SelectorModel {
|
||||
m := selectorModelWithCurrent(title, items, current)
|
||||
m.filter = strings.TrimSpace(filter)
|
||||
m.rankFiltered = true
|
||||
if m.filter != "" {
|
||||
m.cursor = 0
|
||||
m.scrollOffset = 0
|
||||
}
|
||||
return m
|
||||
}
|
||||
|
||||
func currentItemName(items []SelectItem, cursor int) string {
|
||||
if cursor < 0 || cursor >= len(items) {
|
||||
return ""
|
||||
@@ -140,15 +159,22 @@ func currentItemName(items []SelectItem, cursor int) string {
|
||||
return items[cursor].Name
|
||||
}
|
||||
|
||||
func indexOfItemName(items []SelectItem, name string) int {
|
||||
for i, item := range items {
|
||||
if item.Name == name {
|
||||
return i
|
||||
}
|
||||
}
|
||||
return -1
|
||||
}
|
||||
|
||||
func cursorForItemName(items []SelectItem, name string, fallback int) int {
|
||||
if len(items) == 0 {
|
||||
return 0
|
||||
}
|
||||
if name != "" {
|
||||
for i, item := range items {
|
||||
if item.Name == name {
|
||||
return i
|
||||
}
|
||||
if i := indexOfItemName(items, name); i >= 0 {
|
||||
return i
|
||||
}
|
||||
}
|
||||
if fallback < 0 {
|
||||
@@ -167,10 +193,19 @@ func (m selectorModel) filteredItems() []SelectItem {
|
||||
filterLower := strings.ToLower(m.filter)
|
||||
var result []SelectItem
|
||||
for _, item := range m.items {
|
||||
if m.rankFiltered {
|
||||
if selectItemMatchScore(item, filterLower).ok {
|
||||
result = append(result, item)
|
||||
}
|
||||
continue
|
||||
}
|
||||
if strings.Contains(strings.ToLower(item.Name), filterLower) {
|
||||
result = append(result, item)
|
||||
}
|
||||
}
|
||||
if m.rankFiltered {
|
||||
sortSelectItemsForFilter(result, filterLower)
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
@@ -241,6 +276,54 @@ func (m *selectorModel) updateNavigation(msg tea.KeyMsg) {
|
||||
}
|
||||
}
|
||||
|
||||
func (m *selectorModel) UpdateNavigation(msg tea.KeyMsg) {
|
||||
m.updateNavigation(msg)
|
||||
}
|
||||
|
||||
func (m *selectorModel) Move(delta int) {
|
||||
if delta == 0 {
|
||||
return
|
||||
}
|
||||
filtered := m.filteredItems()
|
||||
if len(filtered) == 0 {
|
||||
m.cursor = 0
|
||||
m.scrollOffset = 0
|
||||
return
|
||||
}
|
||||
m.cursor += delta
|
||||
if m.cursor < 0 {
|
||||
m.cursor = 0
|
||||
}
|
||||
if m.cursor >= len(filtered) {
|
||||
m.cursor = len(filtered) - 1
|
||||
}
|
||||
m.updateScroll(m.otherStart())
|
||||
}
|
||||
|
||||
func (m *selectorModel) SetHelpText(help string) {
|
||||
m.helpText = help
|
||||
}
|
||||
|
||||
func (m selectorModel) Filter() string {
|
||||
return m.filter
|
||||
}
|
||||
|
||||
func (m selectorModel) FilteredItems() []SelectItem {
|
||||
return append([]SelectItem(nil), m.filteredItems()...)
|
||||
}
|
||||
|
||||
func (m selectorModel) SelectedItem() (SelectItem, bool) {
|
||||
filtered := m.filteredItems()
|
||||
if len(filtered) == 0 || m.cursor < 0 || m.cursor >= len(filtered) {
|
||||
return SelectItem{}, false
|
||||
}
|
||||
return filtered[m.cursor], true
|
||||
}
|
||||
|
||||
func (m selectorModel) RenderContent() string {
|
||||
return m.renderContent()
|
||||
}
|
||||
|
||||
// updateScroll adjusts scrollOffset based on cursor position.
|
||||
// When not filtering, scrollOffset is relative to the "More" (non-recommended) section.
|
||||
// When filtering, it's relative to the full filtered list.
|
||||
@@ -281,7 +364,7 @@ func (m selectorModel) Update(msg tea.Msg) (tea.Model, tea.Cmd) {
|
||||
wasSet := m.width > 0
|
||||
m.width = msg.Width
|
||||
if wasSet {
|
||||
return m, tea.EnterAltScreen
|
||||
return m, tea.ClearScreen
|
||||
}
|
||||
return m, nil
|
||||
|
||||
@@ -338,6 +421,16 @@ func (m selectorModel) renderItem(s *strings.Builder, item SelectItem, idx int)
|
||||
}
|
||||
}
|
||||
|
||||
func (m selectorModel) renderCompactItem(s *strings.Builder, item SelectItem, idx int) {
|
||||
if idx == m.cursor {
|
||||
s.WriteString(selectorSelectedItemStyle.Render("▸ " + item.Name))
|
||||
s.WriteString(cursorItemSuffix(item))
|
||||
} else {
|
||||
s.WriteString(selectorItemStyle.Render(item.Name))
|
||||
}
|
||||
s.WriteString("\n")
|
||||
}
|
||||
|
||||
// renderContent renders the selector content (title, items, help text) without
|
||||
// checking the cancelled/selected state. This is used by both View() (standalone mode)
|
||||
// and by the TUI modal which embeds a selectorModel.
|
||||
@@ -431,6 +524,57 @@ func (m selectorModel) renderContent() string {
|
||||
return s.String()
|
||||
}
|
||||
|
||||
func (m selectorModel) RenderCompactContent(maxItems int) string {
|
||||
var s strings.Builder
|
||||
|
||||
s.WriteString(selectorTitleStyle.Render(m.title))
|
||||
s.WriteString(" ")
|
||||
if m.filter == "" {
|
||||
s.WriteString(selectorFilterStyle.Render("Type to filter..."))
|
||||
} else {
|
||||
s.WriteString(selectorInputStyle.Render(m.filter))
|
||||
}
|
||||
s.WriteString("\n")
|
||||
|
||||
filtered := m.filteredItems()
|
||||
if len(filtered) == 0 {
|
||||
s.WriteString(selectorItemStyle.Render(selectorDescStyle.Render("(no matches)")))
|
||||
s.WriteString("\n")
|
||||
} else {
|
||||
maxItems = max(1, maxItems)
|
||||
start := 0
|
||||
if len(filtered) > maxItems {
|
||||
start = m.cursor - maxItems/2
|
||||
if start < 0 {
|
||||
start = 0
|
||||
}
|
||||
if maxStart := len(filtered) - maxItems; start > maxStart {
|
||||
start = maxStart
|
||||
}
|
||||
}
|
||||
end := min(len(filtered), start+maxItems)
|
||||
if start > 0 {
|
||||
s.WriteString(selectorMoreStyle.Render(fmt.Sprintf("... %d more above", start)))
|
||||
s.WriteString("\n")
|
||||
}
|
||||
for idx := start; idx < end; idx++ {
|
||||
m.renderCompactItem(&s, filtered[idx], idx)
|
||||
}
|
||||
if remaining := len(filtered) - end; remaining > 0 {
|
||||
s.WriteString(selectorMoreStyle.Render(fmt.Sprintf("... and %d more", remaining)))
|
||||
s.WriteString("\n")
|
||||
}
|
||||
}
|
||||
|
||||
help := "↑/↓ navigate • enter select • esc cancel"
|
||||
if m.helpText != "" {
|
||||
help = m.helpText
|
||||
}
|
||||
s.WriteString(selectorHelpStyle.Render(help))
|
||||
|
||||
return s.String()
|
||||
}
|
||||
|
||||
func (m selectorModel) View() string {
|
||||
if m.cancelled || m.selected != "" {
|
||||
return ""
|
||||
@@ -443,6 +587,99 @@ func (m selectorModel) View() string {
|
||||
return s
|
||||
}
|
||||
|
||||
type selectItemScore struct {
|
||||
ok bool
|
||||
rank int
|
||||
index int
|
||||
lengthDelta int
|
||||
recommended int
|
||||
name string
|
||||
}
|
||||
|
||||
func sortSelectItemsForFilter(items []SelectItem, filter string) {
|
||||
filter = strings.ToLower(strings.TrimSpace(filter))
|
||||
sort.SliceStable(items, func(i, j int) bool {
|
||||
return compareSelectItemsForFilter(items[i], items[j], filter) < 0
|
||||
})
|
||||
}
|
||||
|
||||
func compareSelectItemsForFilter(a, b SelectItem, filter string) int {
|
||||
aScore := selectItemMatchScore(a, filter)
|
||||
bScore := selectItemMatchScore(b, filter)
|
||||
for _, cmp := range []int{
|
||||
compareSelectorInt(aScore.rank, bScore.rank),
|
||||
compareSelectorInt(aScore.index, bScore.index),
|
||||
compareSelectorInt(aScore.lengthDelta, bScore.lengthDelta),
|
||||
compareSelectorInt(aScore.recommended, bScore.recommended),
|
||||
strings.Compare(aScore.name, bScore.name),
|
||||
} {
|
||||
if cmp != 0 {
|
||||
return cmp
|
||||
}
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
func selectItemMatchScore(item SelectItem, filter string) selectItemScore {
|
||||
filter = strings.ToLower(strings.TrimSpace(filter))
|
||||
name := strings.ToLower(strings.TrimSpace(item.Name))
|
||||
description := strings.ToLower(strings.TrimSpace(item.Description))
|
||||
score := selectItemScore{
|
||||
rank: 4,
|
||||
index: 1 << 20,
|
||||
lengthDelta: 1 << 20,
|
||||
name: name,
|
||||
}
|
||||
if item.Recommended {
|
||||
score.recommended = -1
|
||||
}
|
||||
if filter == "" {
|
||||
score.ok = true
|
||||
return score
|
||||
}
|
||||
nameRunes := len([]rune(name))
|
||||
filterRunes := len([]rune(filter))
|
||||
if name == filter {
|
||||
score.ok = true
|
||||
score.rank = 0
|
||||
score.index = 0
|
||||
score.lengthDelta = 0
|
||||
return score
|
||||
}
|
||||
if strings.HasPrefix(name, filter) {
|
||||
score.ok = true
|
||||
score.rank = 1
|
||||
score.index = 0
|
||||
score.lengthDelta = max(0, nameRunes-filterRunes)
|
||||
return score
|
||||
}
|
||||
if index := strings.Index(name, filter); index >= 0 {
|
||||
score.ok = true
|
||||
score.rank = 2
|
||||
score.index = len([]rune(name[:index]))
|
||||
score.lengthDelta = max(0, nameRunes-filterRunes)
|
||||
return score
|
||||
}
|
||||
if index := strings.Index(description, filter); index >= 0 {
|
||||
score.ok = true
|
||||
score.rank = 3
|
||||
score.index = len([]rune(description[:index]))
|
||||
score.lengthDelta = max(0, nameRunes-filterRunes)
|
||||
}
|
||||
return score
|
||||
}
|
||||
|
||||
func compareSelectorInt(a, b int) int {
|
||||
switch {
|
||||
case a < b:
|
||||
return -1
|
||||
case a > b:
|
||||
return 1
|
||||
default:
|
||||
return 0
|
||||
}
|
||||
}
|
||||
|
||||
// cursorForCurrent returns the item index matching current, or 0 if not found.
|
||||
func cursorForCurrent(items []SelectItem, current string) int {
|
||||
if current == "" {
|
||||
@@ -451,10 +688,8 @@ func cursorForCurrent(items []SelectItem, current string) int {
|
||||
|
||||
// Prefer exact name matches before tag-prefix fallback so "qwen3.5" does not
|
||||
// incorrectly select "qwen3.5:cloud" (and vice versa) based on list order.
|
||||
for i, item := range items {
|
||||
if item.Name == current {
|
||||
return i
|
||||
}
|
||||
if i := indexOfItemName(items, current); i >= 0 {
|
||||
return i
|
||||
}
|
||||
|
||||
for i, item := range items {
|
||||
@@ -700,7 +935,7 @@ func (m multiSelectorModel) Update(msg tea.Msg) (tea.Model, tea.Cmd) {
|
||||
wasSet := m.width > 0
|
||||
m.width = msg.Width
|
||||
if wasSet {
|
||||
return m, tea.EnterAltScreen
|
||||
return m, tea.ClearScreen
|
||||
}
|
||||
return m, nil
|
||||
|
||||
|
||||
+2
-2
@@ -60,7 +60,7 @@ func (m signInModel) Update(msg tea.Msg) (tea.Model, tea.Cmd) {
|
||||
wasSet := m.width > 0
|
||||
m.width = msg.Width
|
||||
if wasSet {
|
||||
return m, tea.EnterAltScreen
|
||||
return m, tea.ClearScreen
|
||||
}
|
||||
return m, nil
|
||||
|
||||
@@ -115,7 +115,7 @@ func (m upgradeModel) Update(msg tea.Msg) (tea.Model, tea.Cmd) {
|
||||
wasSet := m.width > 0
|
||||
m.width = msg.Width
|
||||
if wasSet {
|
||||
return m, tea.EnterAltScreen
|
||||
return m, tea.ClearScreen
|
||||
}
|
||||
return m, nil
|
||||
|
||||
|
||||
@@ -0,0 +1,110 @@
|
||||
package tui
|
||||
|
||||
import (
|
||||
"os"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
tea "github.com/charmbracelet/bubbletea"
|
||||
"github.com/charmbracelet/lipgloss"
|
||||
"github.com/ollama/ollama/cmd/launch"
|
||||
"golang.org/x/term"
|
||||
)
|
||||
|
||||
// spinnerStyle dims the spinner so it reads as ancillary status text, matching
|
||||
// the sign-in/upgrade spinners in signin.go.
|
||||
var spinnerStyle = lipgloss.NewStyle().
|
||||
Foreground(lipgloss.AdaptiveColor{Light: "242", Dark: "246"})
|
||||
|
||||
type spinnerTickMsg struct{}
|
||||
|
||||
// spinnerQuitMsg is sent by Stop to ask the program to quit cleanly.
|
||||
type spinnerQuitMsg struct{}
|
||||
|
||||
type spinnerModel struct {
|
||||
message string
|
||||
frame int
|
||||
quitting bool
|
||||
cancelled chan struct{}
|
||||
once sync.Once
|
||||
}
|
||||
|
||||
func (m *spinnerModel) Init() tea.Cmd {
|
||||
return tea.Tick(100*time.Millisecond, func(time.Time) tea.Msg { return spinnerTickMsg{} })
|
||||
}
|
||||
|
||||
func (m *spinnerModel) Update(msg tea.Msg) (tea.Model, tea.Cmd) {
|
||||
switch msg := msg.(type) {
|
||||
case spinnerTickMsg:
|
||||
if m.quitting {
|
||||
return m, nil
|
||||
}
|
||||
m.frame++
|
||||
return m, tea.Tick(100*time.Millisecond, func(time.Time) tea.Msg { return spinnerTickMsg{} })
|
||||
case tea.KeyMsg:
|
||||
// bubbletea runs the terminal in raw mode, so Ctrl+C is delivered here
|
||||
// as a key rather than as a SIGINT. Treat it as a user cancellation:
|
||||
// close the cancelled channel (so the caller's wait loop can abort) and
|
||||
// quit the program so bubbletea restores the terminal before control
|
||||
// returns to the caller.
|
||||
if msg.String() == "ctrl+c" {
|
||||
m.once.Do(func() { close(m.cancelled) })
|
||||
m.quitting = true
|
||||
return m, tea.Quit
|
||||
}
|
||||
case spinnerQuitMsg:
|
||||
m.quitting = true
|
||||
// Returning "" from View on quit clears the spinner line, mirroring how
|
||||
// confirm.go blanks its view when it quits.
|
||||
return m, tea.Quit
|
||||
}
|
||||
return m, nil
|
||||
}
|
||||
|
||||
func (m *spinnerModel) View() string {
|
||||
if m.quitting {
|
||||
return ""
|
||||
}
|
||||
frame := launch.SpinnerFrames[m.frame%len(launch.SpinnerFrames)]
|
||||
return spinnerStyle.Render(frame + " " + m.message)
|
||||
}
|
||||
|
||||
// RunSpinner runs a bubbletea spinner displaying message until the returned
|
||||
// Spinner's Stop is called. Stop signals the program to quit and blocks until
|
||||
// it has exited and cleared its line. If the user presses Ctrl+C while the
|
||||
// spinner is running, Spinner.Cancelled() is closed so the caller can abort
|
||||
// its wait; the program quits and the terminal is restored before Stop
|
||||
// returns. RunSpinner returns nil when there is no interactive terminal, so
|
||||
// launch.StartSpinner can fall back to its ANSI spinner for headless/--yes
|
||||
// runs.
|
||||
func RunSpinner(message string) *launch.Spinner {
|
||||
if !term.IsTerminal(int(os.Stdin.Fd())) || !term.IsTerminal(int(os.Stderr.Fd())) {
|
||||
return nil
|
||||
}
|
||||
|
||||
cancelled := make(chan struct{})
|
||||
m := &spinnerModel{message: message, cancelled: cancelled}
|
||||
p := tea.NewProgram(m, tea.WithOutput(os.Stderr))
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
_, _ = p.Run()
|
||||
close(done)
|
||||
}()
|
||||
|
||||
var once sync.Once
|
||||
stop := func() {
|
||||
once.Do(func() {
|
||||
select {
|
||||
case <-done:
|
||||
// Program already finished (e.g. the user cancelled), so don't
|
||||
// send to it; just ensure it has exited.
|
||||
return
|
||||
default:
|
||||
}
|
||||
p.Send(spinnerQuitMsg{})
|
||||
<-done
|
||||
})
|
||||
}
|
||||
|
||||
return launch.NewSpinner(stop, cancelled)
|
||||
}
|
||||
+22
-82
@@ -42,68 +42,41 @@ type menuItem struct {
|
||||
description string
|
||||
integration string
|
||||
isRunModel bool
|
||||
isOthers bool
|
||||
}
|
||||
|
||||
const pinnedIntegrationCount = 4
|
||||
|
||||
var runModelMenuItem = menuItem{
|
||||
title: "Chat with a model",
|
||||
description: "Start an interactive chat with a model",
|
||||
title: "Chat, Code, & Work",
|
||||
description: "Chat with models, code, search the web, and delegate real work",
|
||||
isRunModel: true,
|
||||
}
|
||||
|
||||
var othersMenuItem = menuItem{
|
||||
title: "More...",
|
||||
description: "Show additional integrations",
|
||||
isOthers: true,
|
||||
}
|
||||
// launcherMenuIntegrations is intentionally short: the root ollama command is
|
||||
// a quick path to the most common launch targets. Other registered
|
||||
// integrations remain available through `ollama launch <integration>`.
|
||||
var launcherMenuIntegrations = []string{"claude", "opencode", "hermes", "openclaw"}
|
||||
|
||||
type model struct {
|
||||
state *launch.LauncherState
|
||||
items []menuItem
|
||||
cursor int
|
||||
showOthers bool
|
||||
width int
|
||||
quitting bool
|
||||
selected bool
|
||||
action TUIAction
|
||||
state *launch.LauncherState
|
||||
items []menuItem
|
||||
cursor int
|
||||
width int
|
||||
quitting bool
|
||||
selected bool
|
||||
action TUIAction
|
||||
}
|
||||
|
||||
func newModel(state *launch.LauncherState) model {
|
||||
m := model{
|
||||
state: state,
|
||||
}
|
||||
m.showOthers = shouldExpandOthers(state)
|
||||
m.items = buildMenuItems(state, m.showOthers)
|
||||
m.items = buildMenuItems(state)
|
||||
m.cursor = initialCursor(state, m.items)
|
||||
return m
|
||||
}
|
||||
|
||||
func shouldExpandOthers(state *launch.LauncherState) bool {
|
||||
if state == nil {
|
||||
return false
|
||||
}
|
||||
for _, item := range otherIntegrationItems(state) {
|
||||
if item.integration == state.LastSelection {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func buildMenuItems(state *launch.LauncherState, showOthers bool) []menuItem {
|
||||
func buildMenuItems(state *launch.LauncherState) []menuItem {
|
||||
items := []menuItem{runModelMenuItem}
|
||||
items = append(items, pinnedIntegrationItems(state)...)
|
||||
|
||||
otherItems := otherIntegrationItems(state)
|
||||
switch {
|
||||
case showOthers:
|
||||
items = append(items, otherItems...)
|
||||
case len(otherItems) > 0:
|
||||
items = append(items, othersMenuItem)
|
||||
}
|
||||
|
||||
items = append(items, launcherIntegrationItems(state)...)
|
||||
return items
|
||||
}
|
||||
|
||||
@@ -119,30 +92,14 @@ func integrationMenuItem(state launch.LauncherIntegrationState) menuItem {
|
||||
}
|
||||
}
|
||||
|
||||
func otherIntegrationItems(state *launch.LauncherState) []menuItem {
|
||||
ordered := orderedIntegrationItems(state)
|
||||
if len(ordered) <= pinnedIntegrationCount {
|
||||
return nil
|
||||
}
|
||||
return ordered[pinnedIntegrationCount:]
|
||||
}
|
||||
|
||||
func pinnedIntegrationItems(state *launch.LauncherState) []menuItem {
|
||||
ordered := orderedIntegrationItems(state)
|
||||
if len(ordered) <= pinnedIntegrationCount {
|
||||
return ordered
|
||||
}
|
||||
return ordered[:pinnedIntegrationCount]
|
||||
}
|
||||
|
||||
func orderedIntegrationItems(state *launch.LauncherState) []menuItem {
|
||||
func launcherIntegrationItems(state *launch.LauncherState) []menuItem {
|
||||
if state == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
items := make([]menuItem, 0, len(state.Integrations))
|
||||
for _, info := range launch.ListIntegrationInfos() {
|
||||
integrationState, ok := state.Integrations[info.Name]
|
||||
items := make([]menuItem, 0, len(launcherMenuIntegrations))
|
||||
for _, name := range launcherMenuIntegrations {
|
||||
integrationState, ok := state.Integrations[name]
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
@@ -151,10 +108,6 @@ func orderedIntegrationItems(state *launch.LauncherState) []menuItem {
|
||||
return items
|
||||
}
|
||||
|
||||
func primaryMenuItemCount(state *launch.LauncherState) int {
|
||||
return 1 + len(pinnedIntegrationItems(state))
|
||||
}
|
||||
|
||||
func initialCursor(state *launch.LauncherState, items []menuItem) int {
|
||||
if state == nil || state.LastSelection == "" {
|
||||
return 0
|
||||
@@ -190,21 +143,12 @@ func (m model) Update(msg tea.Msg) (tea.Model, tea.Cmd) {
|
||||
if m.cursor > 0 {
|
||||
m.cursor--
|
||||
}
|
||||
if m.showOthers && m.cursor < primaryMenuItemCount(m.state) {
|
||||
m.showOthers = false
|
||||
m.items = buildMenuItems(m.state, false)
|
||||
m.cursor = min(m.cursor, len(m.items)-1)
|
||||
}
|
||||
return m, nil
|
||||
|
||||
case "down", "j":
|
||||
if m.cursor < len(m.items)-1 {
|
||||
m.cursor++
|
||||
}
|
||||
if m.cursor < len(m.items) && m.items[m.cursor].isOthers && !m.showOthers {
|
||||
m.showOthers = true
|
||||
m.items = buildMenuItems(m.state, true)
|
||||
}
|
||||
return m, nil
|
||||
|
||||
case "enter", " ":
|
||||
@@ -235,7 +179,7 @@ func (m model) selectableItem(item menuItem) bool {
|
||||
if item.isRunModel {
|
||||
return true
|
||||
}
|
||||
if item.integration == "" || item.isOthers {
|
||||
if item.integration == "" {
|
||||
return false
|
||||
}
|
||||
state, ok := m.state.Integrations[item.integration]
|
||||
@@ -243,7 +187,7 @@ func (m model) selectableItem(item menuItem) bool {
|
||||
}
|
||||
|
||||
func (m model) changeableItem(item menuItem) bool {
|
||||
if item.integration == "" || item.isOthers {
|
||||
if item.integration == "" {
|
||||
return false
|
||||
}
|
||||
state, ok := m.state.Integrations[item.integration]
|
||||
@@ -287,10 +231,6 @@ func (m model) renderMenuItem(index int, item menuItem) string {
|
||||
if m.cursor == index {
|
||||
style = menuSelectedItemStyle
|
||||
}
|
||||
} else if item.isOthers {
|
||||
if m.cursor == index {
|
||||
style = menuSelectedItemStyle
|
||||
}
|
||||
} else {
|
||||
integrationState := m.state.Integrations[item.integration]
|
||||
if !integrationState.Selectable {
|
||||
|
||||
+30
-82
@@ -29,10 +29,10 @@ func launcherTestState() *launch.LauncherState {
|
||||
Selectable: true,
|
||||
Changeable: true,
|
||||
},
|
||||
"codex-app": {
|
||||
Name: "codex-app",
|
||||
DisplayName: "Codex App",
|
||||
Description: "An AI agent you can delegate real work to, by OpenAI",
|
||||
"chatgpt": {
|
||||
Name: "chatgpt",
|
||||
DisplayName: "ChatGPT",
|
||||
Description: "Complete work with ChatGPT",
|
||||
Selectable: true,
|
||||
Changeable: true,
|
||||
},
|
||||
@@ -91,8 +91,6 @@ func integrationSequence(items []menuItem) []string {
|
||||
switch {
|
||||
case item.isRunModel:
|
||||
sequence = append(sequence, "run")
|
||||
case item.isOthers:
|
||||
sequence = append(sequence, "more")
|
||||
case item.integration != "":
|
||||
sequence = append(sequence, item.integration)
|
||||
}
|
||||
@@ -104,81 +102,31 @@ func compareStrings(got, want []string) string {
|
||||
return cmp.Diff(want, got)
|
||||
}
|
||||
|
||||
func expectedCollapsedSequence(state *launch.LauncherState) []string {
|
||||
sequence := []string{"run"}
|
||||
for _, item := range pinnedIntegrationItems(state) {
|
||||
sequence = append(sequence, item.integration)
|
||||
}
|
||||
if len(otherIntegrationItems(state)) > 0 {
|
||||
sequence = append(sequence, "more")
|
||||
}
|
||||
return sequence
|
||||
}
|
||||
|
||||
func expectedExpandedSequence(state *launch.LauncherState) []string {
|
||||
sequence := []string{"run"}
|
||||
for _, item := range pinnedIntegrationItems(state) {
|
||||
sequence = append(sequence, item.integration)
|
||||
}
|
||||
for _, item := range otherIntegrationItems(state) {
|
||||
sequence = append(sequence, item.integration)
|
||||
}
|
||||
return sequence
|
||||
}
|
||||
|
||||
func TestMenuRendersPinnedItemsAndMore(t *testing.T) {
|
||||
func TestMenuRendersRootLaunchChoices(t *testing.T) {
|
||||
state := launcherTestState()
|
||||
menu := newModel(state)
|
||||
wantPrefix := []string{"run", "claude", "codex-app", "hermes", "openclaw"}
|
||||
if findMenuCursorByIntegration(menu.items, "codex-app") == -1 {
|
||||
wantPrefix = []string{"run", "claude", "hermes", "openclaw", "opencode"}
|
||||
}
|
||||
if got := integrationSequence(menu.items); len(got) < len(wantPrefix) {
|
||||
t.Fatalf("expected at least %d menu items, got %v", len(wantPrefix), got)
|
||||
} else if diff := compareStrings(got[:len(wantPrefix)], wantPrefix); diff != "" {
|
||||
t.Fatalf("unexpected primary TUI order: %s", diff)
|
||||
want := []string{"run", "claude", "opencode", "hermes", "openclaw"}
|
||||
if diff := compareStrings(integrationSequence(menu.items), want); diff != "" {
|
||||
t.Fatalf("unexpected root launch choices: %s", diff)
|
||||
}
|
||||
|
||||
view := menu.View()
|
||||
for _, want := range []string{"Chat with a model", "Launch Claude Code", "Launch Hermes Agent", "Launch OpenClaw", "More..."} {
|
||||
for _, want := range []string{
|
||||
"Chat, Code, & Work",
|
||||
"Chat with models, code, search the web, and delegate real work",
|
||||
"Launch Claude Code",
|
||||
"Launch OpenCode",
|
||||
"Launch Hermes Agent",
|
||||
"Launch OpenClaw",
|
||||
} {
|
||||
if !strings.Contains(view, want) {
|
||||
t.Fatalf("expected menu view to contain %q\n%s", want, view)
|
||||
}
|
||||
}
|
||||
if findMenuCursorByIntegration(menu.items, "codex-app") != -1 && !strings.Contains(view, "Launch Codex App") {
|
||||
t.Fatalf("expected menu view to contain Codex App\n%s", view)
|
||||
}
|
||||
if strings.Contains(view, "Launch Claude Desktop") {
|
||||
t.Fatalf("expected hidden Claude Desktop to be absent\n%s", view)
|
||||
}
|
||||
wantOrder := expectedCollapsedSequence(state)
|
||||
if diff := compareStrings(integrationSequence(menu.items), wantOrder); diff != "" {
|
||||
t.Fatalf("unexpected pinned order: %s", diff)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMenuExpandsOthersFromLastSelection(t *testing.T) {
|
||||
state := launcherTestState()
|
||||
overflow := otherIntegrationItems(state)
|
||||
if len(overflow) == 0 {
|
||||
t.Fatal("expected at least one overflow integration")
|
||||
}
|
||||
state.LastSelection = overflow[0].integration
|
||||
|
||||
menu := newModel(state)
|
||||
if !menu.showOthers {
|
||||
t.Fatal("expected others section to expand when last selection is in the overflow list")
|
||||
}
|
||||
view := menu.View()
|
||||
if !strings.Contains(view, overflow[0].title) {
|
||||
t.Fatalf("expected expanded view to contain overflow integration\n%s", view)
|
||||
}
|
||||
if strings.Contains(view, "More...") {
|
||||
t.Fatalf("expected expanded view to replace More... item\n%s", view)
|
||||
}
|
||||
wantOrder := expectedExpandedSequence(state)
|
||||
if diff := compareStrings(integrationSequence(menu.items), wantOrder); diff != "" {
|
||||
t.Fatalf("unexpected expanded order: %s", diff)
|
||||
for _, hidden := range []string{"Launch ChatGPT", "Launch Codex", "Launch Droid", "Launch Pi", "More..."} {
|
||||
if strings.Contains(view, hidden) {
|
||||
t.Fatalf("expected root menu to omit %q\n%s", hidden, view)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -273,24 +221,24 @@ func TestMenuShowsCurrentModelSuffixes(t *testing.T) {
|
||||
|
||||
func TestMenuShowsInstallStatusAndHint(t *testing.T) {
|
||||
state := launcherTestState()
|
||||
codex := state.Integrations["codex"]
|
||||
codex.Installed = false
|
||||
codex.Selectable = false
|
||||
codex.Changeable = false
|
||||
codex.InstallHint = "Install from https://example.com/codex"
|
||||
state.Integrations["codex"] = codex
|
||||
opencode := state.Integrations["opencode"]
|
||||
opencode.Installed = false
|
||||
opencode.Selectable = false
|
||||
opencode.Changeable = false
|
||||
opencode.InstallHint = "Install from https://example.com/opencode"
|
||||
state.Integrations["opencode"] = opencode
|
||||
|
||||
state.LastSelection = "codex"
|
||||
state.LastSelection = "opencode"
|
||||
menu := newModel(state)
|
||||
menu.cursor = findMenuCursorByIntegration(menu.items, "codex")
|
||||
menu.cursor = findMenuCursorByIntegration(menu.items, "opencode")
|
||||
if menu.cursor == -1 {
|
||||
t.Fatal("expected codex menu item in overflow section")
|
||||
t.Fatal("expected opencode menu item")
|
||||
}
|
||||
view := menu.View()
|
||||
if !strings.Contains(view, "(not installed)") {
|
||||
t.Fatalf("expected not-installed marker\n%s", view)
|
||||
}
|
||||
if !strings.Contains(view, codex.InstallHint) {
|
||||
if !strings.Contains(view, opencode.InstallHint) {
|
||||
t.Fatalf("expected install hint in description\n%s", view)
|
||||
}
|
||||
}
|
||||
@@ -39,10 +39,6 @@ func (q *qwen25VLModel) KV(t *Tokenizer) KV {
|
||||
}
|
||||
}
|
||||
|
||||
if q.VisionModel.FullAttentionBlocks == nil {
|
||||
kv["qwen25vl.vision.fullatt_block_indexes"] = []int32{7, 15, 23, 31}
|
||||
}
|
||||
|
||||
kv["qwen25vl.vision.block_count"] = cmp.Or(q.VisionModel.Depth, 32)
|
||||
kv["qwen25vl.vision.embedding_length"] = q.VisionModel.HiddenSize
|
||||
kv["qwen25vl.vision.attention.head_count"] = cmp.Or(q.VisionModel.NumHeads, 16)
|
||||
@@ -53,12 +49,19 @@ func (q *qwen25VLModel) KV(t *Tokenizer) KV {
|
||||
kv["qwen25vl.vision.window_size"] = cmp.Or(q.VisionModel.WindowSize, 112)
|
||||
kv["qwen25vl.vision.attention.layer_norm_epsilon"] = cmp.Or(q.VisionModel.RMSNormEps, 1e-6)
|
||||
kv["qwen25vl.vision.rope.freq_base"] = cmp.Or(q.VisionModel.RopeTheta, 1e4)
|
||||
kv["qwen25vl.vision.fullatt_block_indexes"] = q.VisionModel.FullAttentionBlocks
|
||||
kv["qwen25vl.vision.fullatt_block_indexes"] = q.fullAttentionBlocks()
|
||||
kv["qwen25vl.vision.temporal_patch_size"] = cmp.Or(q.VisionModel.TemporalPatchSize, 2)
|
||||
|
||||
return kv
|
||||
}
|
||||
|
||||
func (q *qwen25VLModel) fullAttentionBlocks() []int32 {
|
||||
if len(q.VisionModel.FullAttentionBlocks) > 0 {
|
||||
return q.VisionModel.FullAttentionBlocks
|
||||
}
|
||||
return []int32{7, 15, 23, 31}
|
||||
}
|
||||
|
||||
func (q *qwen25VLModel) Tensors(ts []Tensor) []*ggml.Tensor {
|
||||
var out []*ggml.Tensor
|
||||
|
||||
|
||||
@@ -0,0 +1,44 @@
|
||||
package convert
|
||||
|
||||
import (
|
||||
"slices"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestQwen25VLFullAttentionBlockDefaults(t *testing.T) {
|
||||
tokenizer := &Tokenizer{Vocabulary: &Vocabulary{}}
|
||||
|
||||
for _, tt := range []struct {
|
||||
name string
|
||||
blocks []int32
|
||||
want []int32
|
||||
}{
|
||||
{
|
||||
name: "nil",
|
||||
want: []int32{7, 15, 23, 31},
|
||||
},
|
||||
{
|
||||
name: "empty",
|
||||
blocks: []int32{},
|
||||
want: []int32{7, 15, 23, 31},
|
||||
},
|
||||
{
|
||||
name: "custom",
|
||||
blocks: []int32{5, 17},
|
||||
want: []int32{5, 17},
|
||||
},
|
||||
} {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
model := &qwen25VLModel{}
|
||||
model.VisionModel.FullAttentionBlocks = tt.blocks
|
||||
|
||||
got, ok := model.KV(tokenizer)["qwen25vl.vision.fullatt_block_indexes"].([]int32)
|
||||
if !ok {
|
||||
t.Fatalf("fullatt_block_indexes has unexpected type %T", got)
|
||||
}
|
||||
if !slices.Equal(got, tt.want) {
|
||||
t.Fatalf("fullatt_block_indexes = %v, want %v", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
+48
-7
@@ -7,34 +7,66 @@ import (
|
||||
"github.com/ollama/ollama/ml"
|
||||
)
|
||||
|
||||
const (
|
||||
cudaV12RuntimeMajor = 12
|
||||
|
||||
minFatbinCompressionCUDARuntimeMinor = 4
|
||||
minFatbinCompressionNVIDIADriverMajor = 550
|
||||
|
||||
minLegacyComputeJITCUDARuntimeMinor = 8
|
||||
// Older CUDA compute targets need newer drivers when they are JITed from PTX.
|
||||
minLegacyComputeJITNVIDIADriverMajor = 570
|
||||
)
|
||||
|
||||
func filterOldCUDADriver(_ context.Context, devices []ml.DeviceInfo) []ml.DeviceInfo {
|
||||
oldCUDA := func(dev ml.DeviceInfo) bool {
|
||||
return dev.Library == "CUDA" && dev.ComputeMajor > 0 && dev.ComputeMajor < 7
|
||||
}
|
||||
|
||||
needsCheck := false
|
||||
hasCUDA := false
|
||||
for _, dev := range devices {
|
||||
if oldCUDA(dev) {
|
||||
needsCheck = true
|
||||
if dev.Library == "CUDA" {
|
||||
hasCUDA = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if !needsCheck {
|
||||
if !hasCUDA {
|
||||
return devices
|
||||
}
|
||||
|
||||
driver := nvidiaDriverMajorFromDevices(devices)
|
||||
if driver == 0 {
|
||||
slog.Warn("could not verify NVIDIA driver compatibility for an older NVIDIA GPU")
|
||||
slog.Warn("could not verify NVIDIA driver compatibility for CUDA")
|
||||
return devices
|
||||
}
|
||||
if driver >= 570 {
|
||||
|
||||
// Match the driver floor to the CUDA runtime we are about to load, so source
|
||||
// builds with older CUDA runtimes can still run on matching older drivers.
|
||||
runtimeMajor, runtimeMinor, hasRuntime := cudaRuntimeVersionFromDevices(devices)
|
||||
runtimeMayUseCompressedFatbins := hasRuntime &&
|
||||
runtimeMajor == cudaV12RuntimeMajor &&
|
||||
runtimeMinor >= minFatbinCompressionCUDARuntimeMinor
|
||||
// CUDA v12.8+ source builds are expected to either use Ollama's PTX packaging
|
||||
// for older compute targets or be built against a matching local driver/toolkit.
|
||||
runtimeMayJITLegacyCompute := hasRuntime &&
|
||||
runtimeMajor == cudaV12RuntimeMajor &&
|
||||
runtimeMinor >= minLegacyComputeJITCUDARuntimeMinor
|
||||
if driver >= minLegacyComputeJITNVIDIADriverMajor || (!runtimeMayUseCompressedFatbins && !runtimeMayJITLegacyCompute) {
|
||||
return devices
|
||||
}
|
||||
|
||||
filtered := devices[:0]
|
||||
for _, dev := range devices {
|
||||
if oldCUDA(dev) {
|
||||
if dev.Library != "CUDA" {
|
||||
filtered = append(filtered, dev)
|
||||
continue
|
||||
}
|
||||
if runtimeMayUseCompressedFatbins && driver < minFatbinCompressionNVIDIADriverMajor {
|
||||
slog.Warn("NVIDIA driver too old",
|
||||
"device", dev.Description, "compute", dev.Compute(), "driver", driver, "required_driver", "550 or newer")
|
||||
continue
|
||||
}
|
||||
if runtimeMayJITLegacyCompute && oldCUDA(dev) {
|
||||
slog.Warn("NVIDIA driver too old",
|
||||
"device", dev.Description, "compute", dev.Compute(), "driver", driver, "required_driver", "570 or newer")
|
||||
continue
|
||||
@@ -52,3 +84,12 @@ func nvidiaDriverMajorFromDevices(devices []ml.DeviceInfo) int {
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
func cudaRuntimeVersionFromDevices(devices []ml.DeviceInfo) (int, int, bool) {
|
||||
for _, dev := range devices {
|
||||
if dev.Library == "CUDA" {
|
||||
return cudaRuntimeVersion(dev.LibraryPath)
|
||||
}
|
||||
}
|
||||
return 0, 0, false
|
||||
}
|
||||
Loaded 100 of 364 files, more files were not shown because too many files have changed in this diff.
Show more
Reference in new issue
Block a user