cmd: remove built-in agent (#18393)

This commit is contained in:
Parth Sareen
2026-09-11 12:21:52 -07:00
committed by GitHub
parent b68b112bd8
commit c16bf9892a
56 changed files with 77 additions and 24892 deletions
-198
View File
@@ -1,198 +0,0 @@
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)
}
-95
View File
@@ -1,95 +0,0 @@
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)
}
}
-667
View File
@@ -1,667 +0,0 @@
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
}
}
-773
View File
@@ -1,773 +0,0 @@
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
View File
@@ -1,177 +0,0 @@
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
}
-104
View File
@@ -1,104 +0,0 @@
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
}
-1092
View File
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
-57
View File
@@ -1,57 +0,0 @@
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
}
-74
View File
@@ -1,74 +0,0 @@
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)
}
}
-813
View File
@@ -1,813 +0,0 @@
package agent
import (
"bytes"
"errors"
"fmt"
"io"
"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
}
// SkillImportResult describes one import attempt. Failed skills do not prevent
// other valid skills in the same source root from being imported.
type SkillImportResult struct {
Source string
SourceDir string
Destination string
Imported []string
Existing []string
Failures []SkillImportFailure
}
// SkillImportFailure identifies a source skill that was deliberately skipped.
// The destination is never changed for a failed skill.
type SkillImportFailure struct {
Name string
Err error
}
// ImportSkills imports skills from a conventional coding-agent source into the
// canonical Ollama skills directory. Supported sources are codex, claude, and
// pi. Existing skills are left untouched: an identical directory is reported
// as existing, and a differing one is reported as a conflict.
func ImportSkills(source string) (SkillImportResult, error) {
home, err := os.UserHomeDir()
if err != nil {
return SkillImportResult{}, fmt.Errorf("resolve home directory: %w", err)
}
destination, err := SkillsDir()
if err != nil {
return SkillImportResult{}, fmt.Errorf("resolve Ollama skills directory: %w", err)
}
return importSkillsFromRoots(source, conventionalSkillImportRoots(home), destination)
}
func conventionalSkillImportRoots(home string) map[string]string {
return map[string]string{
"codex": filepath.Join(home, ".codex", "skills"),
"claude": filepath.Join(home, ".claude", "skills"),
"pi": filepath.Join(home, ".pi", "agent", "skills"),
}
}
func importSkillsFromRoots(source string, roots map[string]string, destination string) (SkillImportResult, error) {
source = strings.ToLower(strings.TrimSpace(source))
sourceDir, ok := roots[source]
if !ok {
return SkillImportResult{}, fmt.Errorf("unknown skill source %q", source)
}
return importSkillsFromDir(source, sourceDir, destination)
}
func importSkillsFromDir(source, sourceDir, destination string) (SkillImportResult, error) {
result := SkillImportResult{Source: source, SourceDir: sourceDir, Destination: destination}
info, err := os.Lstat(sourceDir)
if errors.Is(err, fs.ErrNotExist) {
return result, nil
}
if err != nil {
return result, fmt.Errorf("inspect %s skills directory: %w", source, err)
}
if info.Mode()&os.ModeSymlink != 0 {
return result, fmt.Errorf("inspect %s skills directory: symlinks are not supported", source)
}
if !info.IsDir() {
return result, fmt.Errorf("inspect %s skills directory: not a directory", source)
}
entries, err := os.ReadDir(sourceDir)
if err != nil {
return result, fmt.Errorf("read %s skills directory: %w", source, err)
}
for _, entry := range entries {
name := entry.Name()
path := filepath.Join(sourceDir, name)
if entry.Type()&os.ModeSymlink != 0 {
result.Failures = append(result.Failures, SkillImportFailure{Name: name, Err: errors.New("symlinked skill directories are not supported")})
continue
}
info, err := entry.Info()
if err != nil {
result.Failures = append(result.Failures, SkillImportFailure{Name: name, Err: fmt.Errorf("inspect source: %w", err)})
continue
}
if !info.IsDir() {
continue
}
if !skillName.MatchString(name) {
result.Failures = append(result.Failures, SkillImportFailure{Name: name, Err: errors.New("invalid skill directory name")})
continue
}
if err := validateImportSkill(path, name); err != nil {
result.Failures = append(result.Failures, SkillImportFailure{Name: name, Err: err})
continue
}
state, err := importSkillDirectory(path, filepath.Join(destination, name))
if err != nil {
result.Failures = append(result.Failures, SkillImportFailure{Name: name, Err: err})
continue
}
if state == skillImportExisting {
result.Existing = append(result.Existing, name)
} else {
result.Imported = append(result.Imported, name)
}
}
return result, nil
}
func validateImportSkill(dir, name string) error {
manifest := filepath.Join(dir, skillFilename)
info, err := os.Lstat(manifest)
if err != nil {
return fmt.Errorf("inspect %s: %w", skillFilename, err)
}
if info.Mode()&os.ModeSymlink != 0 || !info.Mode().IsRegular() {
return fmt.Errorf("%s must be a regular, non-symlinked file", skillFilename)
}
if _, err := parseSkill(manifest, name); err != nil {
return err
}
return walkImportTree(dir, func(path string, entry fs.DirEntry, info fs.FileInfo) error {
if info.IsDir() || path == dir {
return nil
}
if !info.Mode().IsRegular() {
return fmt.Errorf("only regular files may be imported: %s", path)
}
file, err := os.Open(path)
if err != nil {
return fmt.Errorf("read %s: %w", path, err)
}
return file.Close()
})
}
func walkImportTree(root string, visit func(string, fs.DirEntry, fs.FileInfo) error) error {
return filepath.WalkDir(root, func(path string, entry fs.DirEntry, err error) error {
if err != nil {
return err
}
rel, err := filepath.Rel(root, path)
if err != nil || rel == ".." || strings.HasPrefix(rel, ".."+string(filepath.Separator)) {
return fmt.Errorf("unsafe skill path %q", path)
}
if entry.Type()&os.ModeSymlink != 0 {
return fmt.Errorf("symlinks may not be imported: %s", path)
}
info, err := entry.Info()
if err != nil {
return err
}
return visit(path, entry, info)
})
}
type skillImportState int
const (
skillImportCopied skillImportState = iota
skillImportExisting
)
func importSkillDirectory(source, destination string) (skillImportState, error) {
if info, err := os.Lstat(destination); err == nil {
if info.Mode()&os.ModeSymlink != 0 || !info.IsDir() {
return 0, errors.New("destination exists but is not a regular directory")
}
same, err := sameImportTree(source, destination)
if err != nil {
return 0, fmt.Errorf("inspect existing destination: %w", err)
}
if same {
return skillImportExisting, nil
}
return 0, errors.New("destination skill already exists with different contents")
} else if !errors.Is(err, fs.ErrNotExist) {
return 0, fmt.Errorf("inspect destination: %w", err)
}
if err := ensureImportDestination(filepath.Dir(destination)); err != nil {
return 0, err
}
stage, err := os.MkdirTemp(filepath.Dir(destination), "."+filepath.Base(destination)+".import-")
if err != nil {
return 0, fmt.Errorf("create import staging directory: %w", err)
}
defer os.RemoveAll(stage)
if err := copyImportTree(source, stage); err != nil {
return 0, err
}
if _, err := os.Lstat(destination); err == nil {
return 0, errors.New("destination skill was created during import")
} else if !errors.Is(err, fs.ErrNotExist) {
return 0, fmt.Errorf("inspect destination before install: %w", err)
}
if err := os.Rename(stage, destination); err != nil {
return 0, fmt.Errorf("install imported skill: %w", err)
}
return skillImportCopied, nil
}
func ensureImportDestination(dir string) error {
if err := os.MkdirAll(dir, 0o755); err != nil {
return fmt.Errorf("create Ollama skills directory: %w", err)
}
info, err := os.Lstat(dir)
if err != nil {
return fmt.Errorf("inspect Ollama skills directory: %w", err)
}
if info.Mode()&os.ModeSymlink != 0 || !info.IsDir() {
return errors.New("Ollama skills directory must be a regular, non-symlinked directory")
}
return nil
}
func copyImportTree(source, destination string) error {
return walkImportTree(source, func(path string, entry fs.DirEntry, info fs.FileInfo) error {
rel, err := filepath.Rel(source, path)
if err != nil {
return err
}
target := destination
if rel != "." {
target = filepath.Join(destination, rel)
}
if info.IsDir() {
if rel == "." {
return nil
}
return os.Mkdir(target, info.Mode().Perm())
}
if !info.Mode().IsRegular() {
return fmt.Errorf("only regular files may be imported: %s", path)
}
return copyImportFile(path, target, info.Mode().Perm())
})
}
func copyImportFile(source, destination string, mode fs.FileMode) error {
in, err := os.Open(source)
if err != nil {
return fmt.Errorf("read %s: %w", source, err)
}
defer in.Close()
out, err := os.OpenFile(destination, os.O_WRONLY|os.O_CREATE|os.O_EXCL, mode)
if err != nil {
return fmt.Errorf("create %s: %w", destination, err)
}
_, copyErr := io.Copy(out, in)
closeErr := out.Close()
if copyErr != nil {
return fmt.Errorf("copy %s: %w", source, copyErr)
}
if closeErr != nil {
return fmt.Errorf("write %s: %w", destination, closeErr)
}
return nil
}
func sameImportTree(source, destination string) (bool, error) {
seen := make(map[string]struct{})
same := true
err := walkImportTree(source, func(path string, entry fs.DirEntry, info fs.FileInfo) error {
rel, err := filepath.Rel(source, path)
if err != nil {
return err
}
seen[rel] = struct{}{}
other := destination
if rel != "." {
other = filepath.Join(destination, rel)
}
otherInfo, err := os.Lstat(other)
if errors.Is(err, fs.ErrNotExist) {
same = false
return nil
}
if err != nil {
return err
}
if otherInfo.Mode()&os.ModeSymlink != 0 || otherInfo.IsDir() != info.IsDir() || (!info.IsDir() && !otherInfo.Mode().IsRegular()) {
same = false
return nil
}
if info.Mode().IsRegular() {
equal, err := sameImportFile(path, other)
if err != nil {
return err
}
if !equal {
same = false
}
}
return nil
})
if err != nil || !same {
return same, err
}
err = walkImportTree(destination, func(path string, entry fs.DirEntry, info fs.FileInfo) error {
rel, err := filepath.Rel(destination, path)
if err != nil {
return err
}
if _, ok := seen[rel]; !ok {
same = false
}
return nil
})
return same, err
}
func sameImportFile(first, second string) (bool, error) {
a, err := os.Open(first)
if err != nil {
return false, err
}
defer a.Close()
b, err := os.Open(second)
if err != nil {
return false, err
}
defer b.Close()
left := make([]byte, 32*1024)
right := make([]byte, len(left))
for {
n, errA := a.Read(left)
m, errB := b.Read(right)
if n != m || !bytes.Equal(left[:n], right[:m]) {
return false, nil
}
if errA == io.EOF && errB == io.EOF {
return true, nil
}
if errA != nil && errA != io.EOF {
return false, errA
}
if errB != nil && errB != io.EOF {
return false, errB
}
if errA == io.EOF || errB == io.EOF {
return false, nil
}
}
}
// 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...)
}
// ExcludeNames removes skills whose names are reserved by a caller. It returns
// the excluded names in sorted order.
func (c *SkillCatalog) ExcludeNames(names []string) []string {
if c == nil {
return nil
}
reserved := make(map[string]struct{}, len(names))
for _, name := range names {
name = strings.TrimPrefix(strings.ToLower(strings.TrimSpace(name)), "/")
if name != "" {
reserved[name] = struct{}{}
}
}
var excluded []string
for name := range c.skills {
if _, ok := reserved[name]; !ok {
continue
}
delete(c.skills, name)
excluded = append(excluded, name)
}
sort.Strings(excluded)
return excluded
}
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")
}
-516
View File
@@ -1,516 +0,0 @@
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 writeImportFixtureSkill(t *testing.T, dir string) {
t.Helper()
contents, err := os.ReadFile(filepath.Join("testdata", "import", "release-notes", skillFilename))
if err != nil {
t.Fatal(err)
}
path := filepath.Join(dir, "release-notes", skillFilename)
if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(path, contents, 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 TestSkillCatalogExcludeNames(t *testing.T) {
dir := t.TempDir()
for _, name := range []string{"release-notes", "system", "exit"} {
writeCatalogSkill(t, dir, name, "instructions")
}
catalog, err := DiscoverSkills(dir)
if err != nil {
t.Fatal(err)
}
if got, want := strings.Join(catalog.ExcludeNames([]string{"/system", "EXIT"}), ","), "exit,system"; got != want {
t.Fatalf("excluded skills = %q, want %q", got, want)
}
if _, err := catalog.Load("system"); err == nil {
t.Fatal("excluded system skill should not load")
}
if _, err := catalog.Load("exit"); err == nil {
t.Fatal("excluded exit skill should not load")
}
if _, err := catalog.Load("release-notes"); err != nil {
t.Fatalf("non-conflicting skill should remain available: %v", err)
}
}
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)
}
}
func TestImportSkillsCopiesFixtureAndIsIdempotent(t *testing.T) {
source := t.TempDir()
destination := t.TempDir()
writeImportFixtureSkill(t, source)
writeCatalogSkill(t, source, "broken", "---\nname: another-skill\ndescription: Deliberately invalid.\n---\nIgnore this.")
if err := os.MkdirAll(filepath.Join(source, "release-notes", "references"), 0o755); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(filepath.Join(source, "release-notes", "references", "style.txt"), []byte("Keep it short.\n"), 0o644); err != nil {
t.Fatal(err)
}
if err := os.MkdirAll(filepath.Join(source, "release-notes", "scripts"), 0o755); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(filepath.Join(source, "release-notes", "scripts", "prepare.sh"), []byte("#!/bin/sh\n"), 0o755); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(filepath.Join(source, "ignored.md"), []byte("Ignored root file.\n"), 0o644); err != nil {
t.Fatal(err)
}
result, err := importSkillsFromDir("codex", source, destination)
if err != nil {
t.Fatal(err)
}
if got, want := strings.Join(result.Imported, ","), "release-notes"; got != want {
t.Fatalf("imported = %q, want %q", got, want)
}
catalog, err := DiscoverSkills(destination)
if err != nil {
t.Fatal(err)
}
skill, err := catalog.Load("release-notes")
if err != nil || skill.Description != "Draft concise release notes." {
t.Fatalf("imported skill = %#v, %v", skill, err)
}
if got := len(result.Failures); got != 1 || result.Failures[0].Name != "broken" {
t.Fatalf("failures = %#v, want broken fixture failure", result.Failures)
}
for _, file := range []string{skillFilename, filepath.Join("references", "style.txt"), filepath.Join("scripts", "prepare.sh")} {
if _, err := os.Stat(filepath.Join(destination, "release-notes", file)); err != nil {
t.Fatalf("imported fixture file %q: %v", file, err)
}
}
result, err = importSkillsFromDir("codex", source, destination)
if err != nil {
t.Fatal(err)
}
if got, want := strings.Join(result.Existing, ","), "release-notes"; got != want {
t.Fatalf("existing = %q, want %q", got, want)
}
if len(result.Imported) != 0 {
t.Fatalf("repeated import copied skills: %#v", result.Imported)
}
}
func TestImportSkillsLeavesConflictsAndUnsafeSourcesUntouched(t *testing.T) {
source := t.TempDir()
destination := t.TempDir()
writeCatalogSkill(t, source, "release-notes", "source instructions")
writeCatalogSkill(t, destination, "release-notes", "existing instructions")
writeCatalogSkill(t, source, "nested-link", "safe manifest")
if err := os.Symlink(filepath.Join(source, "release-notes", skillFilename), filepath.Join(source, "nested-link", "reference")); err != nil {
t.Skipf("symlink not supported: %v", err)
}
if err := os.Symlink(filepath.Join(source, "release-notes"), filepath.Join(source, "linked-skill")); err != nil {
t.Skipf("symlink not supported: %v", err)
}
result, err := importSkillsFromDir("codex", source, destination)
if err != nil {
t.Fatal(err)
}
if len(result.Imported) != 0 || len(result.Existing) != 0 {
t.Fatalf("unexpected successful import: %#v", result)
}
if got, err := os.ReadFile(filepath.Join(destination, "release-notes", skillFilename)); err != nil || !strings.Contains(string(got), "existing instructions") {
t.Fatalf("conflicting destination changed: %q, %v", got, err)
}
failed := make(map[string]bool)
for _, failure := range result.Failures {
failed[failure.Name] = true
}
for _, name := range []string{"release-notes", "nested-link", "linked-skill"} {
if !failed[name] {
t.Fatalf("missing failure for %q: %#v", name, result.Failures)
}
}
}
func TestImportSkillsRejectsSymlinkedRoot(t *testing.T) {
root := t.TempDir()
source := filepath.Join(t.TempDir(), "codex-skills")
if err := os.Symlink(root, source); err != nil {
t.Skipf("symlink not supported: %v", err)
}
result, err := importSkillsFromDir("codex", source, t.TempDir())
if err == nil || !strings.Contains(err.Error(), "symlinks are not supported") {
t.Fatalf("symlinked root error = %v", err)
}
if len(result.Imported) != 0 || len(result.Existing) != 0 || len(result.Failures) != 0 {
t.Fatalf("symlinked root result = %#v", result)
}
}
func TestImportSkillsMissingRootAndConfiguredRoots(t *testing.T) {
result, err := importSkillsFromDir("codex", filepath.Join(t.TempDir(), "missing"), t.TempDir())
if err != nil {
t.Fatal(err)
}
if len(result.Imported) != 0 || len(result.Existing) != 0 || len(result.Failures) != 0 {
t.Fatalf("missing root result = %#v", result)
}
destination := t.TempDir()
rootBase := t.TempDir()
roots := map[string]string{
"codex": filepath.Join(rootBase, "codex"),
"claude": filepath.Join(rootBase, "claude"),
"pi": filepath.Join(rootBase, "pi"),
}
for _, test := range []struct {
source string
root string
name string
}{
{source: "codex", root: roots["codex"], name: "from-codex"},
{source: "claude", root: roots["claude"], name: "from-claude"},
{source: "pi", root: roots["pi"], name: "from-pi"},
} {
t.Run(test.source, func(t *testing.T) {
writeCatalogSkill(t, test.root, test.name, "from "+test.source)
result, err = importSkillsFromRoots(test.source, roots, destination)
if err != nil {
t.Fatal(err)
}
if result.SourceDir != test.root {
t.Fatalf("source dir = %q, want %q", result.SourceDir, test.root)
}
if _, err := os.Stat(filepath.Join(destination, test.name, skillFilename)); err != nil {
t.Fatalf("conventional source was not imported: %v", err)
}
})
}
if _, err := importSkillsFromRoots("unknown", roots, destination); err == nil || !strings.Contains(err.Error(), "unknown skill source") {
t.Fatalf("unknown source error = %v", err)
}
}
func TestConventionalSkillImportRoots(t *testing.T) {
home := t.TempDir()
roots := conventionalSkillImportRoots(home)
for source, want := range map[string]string{
"codex": filepath.Join(home, ".codex", "skills"),
"claude": filepath.Join(home, ".claude", "skills"),
"pi": filepath.Join(home, ".pi", "agent", "skills"),
} {
if got := roots[source]; got != want {
t.Fatalf("%s root = %q, want %q", source, got, want)
}
}
}
func TestImportSkillsRejectsUnreadableManifest(t *testing.T) {
source := t.TempDir()
writeCatalogSkill(t, source, "private", "do not read")
manifest := filepath.Join(source, "private", skillFilename)
if err := os.Chmod(manifest, 0); err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = os.Chmod(manifest, 0o644) })
if _, err := os.ReadFile(manifest); err == nil {
t.Skip("test user can read a mode-000 file")
}
result, err := importSkillsFromDir("codex", source, t.TempDir())
if err != nil {
t.Fatal(err)
}
if len(result.Failures) != 1 || result.Failures[0].Name != "private" {
t.Fatalf("failures = %#v", result.Failures)
}
}
-8
View File
@@ -1,8 +0,0 @@
---
name: release-notes
description: Draft concise release notes.
---
# Release notes
Use short bullets.
-450
View File
@@ -1,450 +0,0 @@
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)
}
-258
View File
@@ -1,258 +0,0 @@
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)
}
}
-49
View File
@@ -1,49 +0,0 @@
//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
}
-40
View File
@@ -1,40 +0,0 @@
//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)
}
}
-134
View File
@@ -1,134 +0,0 @@
//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)
}
}
-15
View File
@@ -1,15 +0,0 @@
//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)
}
}
-711
View File
@@ -1,711 +0,0 @@
package tools
import (
"bufio"
"cmp"
"context"
"encoding/json"
"fmt"
"io"
"os"
"path/filepath"
"slices"
"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. Pass multiple edits to change separate parts of the file in one call."
}
func (e *Edit) Schema() api.ToolFunction {
editProps := api.NewToolPropertiesMap()
editProps.Set("old_text", api.ToolProperty{
Type: api.PropertyType{"string"},
Description: "Exact text for one targeted replacement. Must match the original file exactly once and must not overlap with any other edit's old_text.",
})
editProps.Set("new_text", api.ToolProperty{
Type: api.PropertyType{"string"},
Description: "Replacement text for this targeted edit.",
})
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("edits", api.ToolProperty{
Type: api.PropertyType{"array"},
Items: api.ToolProperty{
Type: api.PropertyType{"object"},
Properties: editProps,
Required: []string{"old_text", "new_text"},
},
Description: "One or more exact-text replacements. Each is matched against the original file, not against the output of earlier edits. Keep old_text as small as possible while still unique in the file; merge changes to the same or adjacent lines into a single edit.",
})
props.Set("replace_all", api.ToolProperty{
Type: api.PropertyType{"boolean"},
Description: "Replace every occurrence. Defaults to false; only applies when a single edit is provided.",
})
return api.ToolFunction{
Name: e.Name(),
Description: e.Description(),
Parameters: api.ToolFunctionParameters{
Type: "object",
Properties: props,
Required: []string{"path", "edits"},
},
}
}
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")
}
edits, replaceAll, err := parseEditArgs(args)
if err != nil {
return agent.ToolResult{}, err
}
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)
var updated string
replacements := 0
if replaceAll {
matches := strings.Count(content, edits[0].OldText)
if matches == 0 {
return agent.ToolResult{}, fmt.Errorf("old_text was not found in %s", path)
}
updated = strings.ReplaceAll(content, edits[0].OldText, edits[0].NewText)
replacements = matches
} else {
// Every edit is matched against the original file content rather
// than the output of earlier edits, so each edit must match exactly
// once and edits must target disjoint regions.
matched := make([]editMatch, 0, len(edits))
for i, edit := range edits {
count := strings.Count(content, edit.OldText)
if count == 0 {
return agent.ToolResult{}, editNotFoundError(path, i, len(edits))
}
if count > 1 {
return agent.ToolResult{}, editAmbiguousError(path, i, len(edits), count)
}
matched = append(matched, editMatch{
editIndex: i,
offset: strings.Index(content, edit.OldText),
length: len(edit.OldText),
newText: edit.NewText,
})
replacements++
}
slices.SortFunc(matched, func(a, b editMatch) int { return cmp.Compare(a.offset, b.offset) })
for i := 1; i < len(matched); i++ {
prev, cur := matched[i-1], matched[i]
if prev.offset+prev.length > cur.offset {
return agent.ToolResult{}, fmt.Errorf("edits[%d] and edits[%d] overlap in %s; merge them into one edit or target disjoint text", prev.editIndex, cur.editIndex, path)
}
}
// Apply from the end of the file backwards so earlier offsets stay valid.
updated = content
for i := len(matched) - 1; i >= 0; i-- {
m := matched[i]
updated = updated[:m.offset] + m.newText + updated[m.offset+m.length:]
}
}
if updated == content {
return agent.ToolResult{}, fmt.Errorf("edit produced no changes in %s; replacement text is identical to the original", path)
}
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 edit%s, %d replacement%s).", path, len(edits), plural(len(edits)), replacements, plural(replacements))}, nil
}
// editReplacement is one targeted replacement within an edit call.
type editReplacement struct {
OldText string
NewText string
}
// editMatch locates one editReplacement within the original file content.
type editMatch struct {
editIndex int
offset int
length int
newText string
}
// parseEditArgs normalizes edit arguments from a tool call into a list of
// replacements. It accepts the `edits` array form and tolerates legacy
// top-level old_text/new_text args as well as stringified JSON, mirroring
// the pi coding agent's argument handling.
func parseEditArgs(args map[string]any) ([]editReplacement, bool, error) {
replaceAll, _ := args["replace_all"].(bool)
var edits []editReplacement
if raw, ok := args["edits"]; ok {
parsed, err := parseEditArray(raw)
if err != nil {
return nil, false, err
}
edits = parsed
}
// Fold a legacy top-level old_text/new_text pair into edits.
if oldText, ok := args["old_text"].(string); ok {
newText, ok := args["new_text"].(string)
if !ok {
return nil, false, fmt.Errorf("new_text parameter is required")
}
edits = append(edits, editReplacement{OldText: oldText, NewText: newText})
}
if len(edits) == 0 {
return nil, false, fmt.Errorf("edits parameter is required")
}
for i, edit := range edits {
if edit.OldText == "" {
if len(edits) == 1 {
return nil, false, fmt.Errorf("old_text parameter is required")
}
return nil, false, fmt.Errorf("edits[%d].old_text must not be empty", i)
}
}
if replaceAll && len(edits) != 1 {
return nil, false, fmt.Errorf("replace_all only applies to a single edit")
}
return edits, replaceAll, nil
}
func parseEditArray(raw any) ([]editReplacement, error) {
if s, ok := raw.(string); ok {
// Some models serialize array arguments as a JSON string.
if err := json.Unmarshal([]byte(s), &raw); err != nil {
return nil, fmt.Errorf("edits must be an array of {old_text, new_text} objects")
}
}
items, ok := raw.([]any)
if !ok {
return nil, fmt.Errorf("edits must be an array of {old_text, new_text} objects")
}
edits := make([]editReplacement, 0, len(items))
for i, item := range items {
entry, ok := item.(map[string]any)
if !ok {
return nil, fmt.Errorf("edits[%d] must be an object with old_text and new_text", i)
}
oldText, oldOK := editTextArg(entry, "old_text", "oldText")
newText, newOK := editTextArg(entry, "new_text", "newText")
if !oldOK || !newOK {
return nil, fmt.Errorf("edits[%d] must be an object with old_text and new_text", i)
}
edits = append(edits, editReplacement{OldText: oldText, NewText: newText})
}
return edits, nil
}
// editTextArg reads the first present string key, tolerating both snake_case
// and camelCase spellings that models emit.
func editTextArg(entry map[string]any, keys ...string) (string, bool) {
for _, key := range keys {
if value, ok := entry[key].(string); ok {
return value, true
}
}
return "", false
}
func editNotFoundError(path string, editIndex, totalEdits int) error {
if totalEdits == 1 {
return fmt.Errorf("old_text was not found in %s", path)
}
return fmt.Errorf("edits[%d].old_text was not found in %s", editIndex, path)
}
func editAmbiguousError(path string, editIndex, totalEdits, occurrences int) error {
if totalEdits == 1 {
return fmt.Errorf("old_text matched %d times in %s; set replace_all to true to replace every match", occurrences, path)
}
return fmt.Errorf("edits[%d].old_text matched %d times in %s; each edit must match exactly once, so provide more surrounding context", editIndex, occurrences, path)
}
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"
}
-571
View File
@@ -1,571 +0,0 @@
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 TestEditAppliesMultipleEdits(t *testing.T) {
dir := t.TempDir()
path := filepath.Join(dir, "note.txt")
if err := os.WriteFile(path, []byte("alpha beta gamma delta\n"), 0o644); err != nil {
t.Fatal(err)
}
result, err := (&Edit{}).Execute(context.Background(), agent.ToolContext{WorkingDir: dir}, map[string]any{
"path": "note.txt",
"edits": []any{
map[string]any{"old_text": "beta", "new_text": "BETA"},
map[string]any{"old_text": "delta", "new_text": "DELTA"},
},
})
if err != nil {
t.Fatal(err)
}
if !strings.Contains(result.Content, "2 edits, 2 replacements") {
t.Fatalf("result = %q", result.Content)
}
content, err := os.ReadFile(path)
if err != nil {
t.Fatal(err)
}
if string(content) != "alpha BETA gamma DELTA\n" {
t.Fatalf("content = %q", content)
}
}
func TestEditMatchesEditsAgainstOriginalContent(t *testing.T) {
dir := t.TempDir()
path := filepath.Join(dir, "note.txt")
if err := os.WriteFile(path, []byte("abc def\n"), 0o644); err != nil {
t.Fatal(err)
}
// edits[1] must target the original "def", not the one introduced by edits[0].
_, err := (&Edit{}).Execute(context.Background(), agent.ToolContext{WorkingDir: dir}, map[string]any{
"path": "note.txt",
"edits": []any{
map[string]any{"old_text": "abc", "new_text": "def"},
map[string]any{"old_text": "def", "new_text": "ghi"},
},
})
if err != nil {
t.Fatal(err)
}
content, err := os.ReadFile(path)
if err != nil {
t.Fatal(err)
}
if string(content) != "def ghi\n" {
t.Fatalf("content = %q", content)
}
}
func TestEditRejectsOverlappingEdits(t *testing.T) {
dir := t.TempDir()
path := filepath.Join(dir, "note.txt")
if err := os.WriteFile(path, []byte("abc\n"), 0o644); err != nil {
t.Fatal(err)
}
_, err := (&Edit{}).Execute(context.Background(), agent.ToolContext{WorkingDir: dir}, map[string]any{
"path": "note.txt",
"edits": []any{
map[string]any{"old_text": "ab", "new_text": "x"},
map[string]any{"old_text": "bc", "new_text": "y"},
},
})
if err == nil {
t.Fatal("expected overlapping edits to fail")
}
if !strings.Contains(err.Error(), "overlap") {
t.Fatalf("err = %v", err)
}
}
func TestEditMultipleEditsNotFoundIndexed(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)
}
_, err := (&Edit{}).Execute(context.Background(), agent.ToolContext{WorkingDir: dir}, map[string]any{
"path": "note.txt",
"edits": []any{
map[string]any{"old_text": "hello", "new_text": "hi"},
map[string]any{"old_text": "missing", "new_text": "x"},
},
})
if err == nil {
t.Fatal("expected missing edit to fail")
}
if !strings.Contains(err.Error(), "edits[1]") {
t.Fatalf("err = %v", err)
}
}
func TestEditMultipleEditsAmbiguousIndexed(t *testing.T) {
dir := t.TempDir()
path := filepath.Join(dir, "note.txt")
if err := os.WriteFile(path, []byte("hello same same\n"), 0o644); err != nil {
t.Fatal(err)
}
_, err := (&Edit{}).Execute(context.Background(), agent.ToolContext{WorkingDir: dir}, map[string]any{
"path": "note.txt",
"edits": []any{
map[string]any{"old_text": "hello", "new_text": "hi"},
map[string]any{"old_text": "same", "new_text": "x"},
},
})
if err == nil {
t.Fatal("expected ambiguous edit to fail")
}
if !strings.Contains(err.Error(), "edits[1]") || !strings.Contains(err.Error(), "matched 2 times") {
t.Fatalf("err = %v", err)
}
}
func TestEditRejectsEmptyEdits(t *testing.T) {
dir := t.TempDir()
path := filepath.Join(dir, "note.txt")
if err := os.WriteFile(path, []byte("hello\n"), 0o644); err != nil {
t.Fatal(err)
}
for name, args := range map[string]map[string]any{
"missing edits": {"path": "note.txt"},
"empty edits": {"path": "note.txt", "edits": []any{}},
} {
if _, err := (&Edit{}).Execute(context.Background(), agent.ToolContext{WorkingDir: dir}, args); err == nil {
t.Fatalf("%s: expected error", name)
} else if !strings.Contains(err.Error(), "edits parameter is required") {
t.Fatalf("%s: err = %v", name, err)
}
}
}
func TestEditRejectsEmptyOldTextInArray(t *testing.T) {
dir := t.TempDir()
path := filepath.Join(dir, "note.txt")
if err := os.WriteFile(path, []byte("hello\n"), 0o644); err != nil {
t.Fatal(err)
}
_, err := (&Edit{}).Execute(context.Background(), agent.ToolContext{WorkingDir: dir}, map[string]any{
"path": "note.txt",
"edits": []any{
map[string]any{"old_text": "hello", "new_text": "hi"},
map[string]any{"old_text": "", "new_text": "x"},
},
})
if err == nil {
t.Fatal("expected empty old_text to fail")
}
if !strings.Contains(err.Error(), "edits[1].old_text must not be empty") {
t.Fatalf("err = %v", err)
}
}
func TestEditAcceptsJSONStringEdits(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)
}
// Some models serialize array arguments as a JSON string.
_, err := (&Edit{}).Execute(context.Background(), agent.ToolContext{WorkingDir: dir}, map[string]any{
"path": "note.txt",
"edits": `[{"oldText": "hello", "newText": "hi"}, {"oldText": "world", "newText": "earth"}]`,
})
if err != nil {
t.Fatal(err)
}
content, err := os.ReadFile(path)
if err != nil {
t.Fatal(err)
}
if string(content) != "hi earth\n" {
t.Fatalf("content = %q", content)
}
}
func TestEditRejectsReplaceAllWithMultipleEdits(t *testing.T) {
dir := t.TempDir()
path := filepath.Join(dir, "note.txt")
if err := os.WriteFile(path, []byte("a b c\n"), 0o644); err != nil {
t.Fatal(err)
}
_, err := (&Edit{}).Execute(context.Background(), agent.ToolContext{WorkingDir: dir}, map[string]any{
"path": "note.txt",
"replace_all": true,
"edits": []any{
map[string]any{"old_text": "a", "new_text": "x"},
map[string]any{"old_text": "b", "new_text": "y"},
},
})
if err == nil {
t.Fatal("expected replace_all with multiple edits to fail")
}
if !strings.Contains(err.Error(), "replace_all") {
t.Fatalf("err = %v", err)
}
}
func TestEditRejectsNoChange(t *testing.T) {
dir := t.TempDir()
path := filepath.Join(dir, "note.txt")
if err := os.WriteFile(path, []byte("hello\n"), 0o644); 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": "hello",
})
if err == nil {
t.Fatal("expected no-change edit to fail")
}
if !strings.Contains(err.Error(), "no changes") {
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)
}
}
-121
View File
@@ -1,121 +0,0 @@
//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)
}
}
-41
View File
@@ -1,41 +0,0 @@
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.
// Model-initiated loads require approval because a skill's instructions can
// influence the rest of the run. Explicit user activation is handled by the
// session's synthetic skill call and bypasses this adapter.
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) RequiresApproval(map[string]any) bool { return true }
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
}
-163
View File
@@ -1,163 +0,0 @@
package tools
import (
"context"
"os"
"path/filepath"
"strings"
"testing"
"github.com/ollama/ollama/agent"
"github.com/ollama/ollama/api"
)
func TestSkillLoadsCoreCatalogWithApproval(t *testing.T) {
catalog := testSkillCatalog(t)
tool := &Skill{Catalog: catalog}
if !agent.ToolRequiresApproval(tool, map[string]any{"name": "release-notes"}) {
t.Fatal("model-initiated skill loading should require approval")
}
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)
}
}
func TestModelSkillLoadRequiresApproval(t *testing.T) {
for _, tt := range []struct {
name string
approval agent.Approval
prompt bool
wantCalls int
wantPrompts int
wantResult string
}{
{name: "rejected", approval: agent.Approval{Reason: "Skill loading denied."}, prompt: true, wantCalls: 1, wantPrompts: 1, wantResult: "Skill loading denied."},
{name: "approved", approval: agent.Approval{Allow: true}, prompt: true, wantCalls: 2, wantPrompts: 1, wantResult: "Use concise bullets."},
{name: "headless denied", wantCalls: 1, wantResult: "Tool execution requires approval"},
} {
t.Run(tt.name, func(t *testing.T) {
catalog := testSkillCatalog(t)
args := api.NewToolCallFunctionArguments()
args.Set("name", "release-notes")
client := &skillTestClient{responses: [][]api.ChatResponse{
{{Message: api.Message{Role: "assistant", ToolCalls: []api.ToolCall{{
ID: "call_skill_1",
Function: api.ToolCallFunction{Name: "skill", Arguments: args},
}}}}},
{{Message: api.Message{Role: "assistant", Content: "done"}}},
}}
var prompter *skillApprovalPrompter
var approvalPrompter agent.ApprovalPrompter
if tt.prompt {
prompter = &skillApprovalPrompter{result: tt.approval}
approvalPrompter = prompter
}
registry := &agent.Registry{}
registry.Register(&Skill{Catalog: catalog})
result, err := (&agent.Session{
Client: client,
Tools: registry,
ApprovalPrompter: approvalPrompter,
}).Run(context.Background(), agent.RunOptions{
Model: "test",
NewMessages: []api.Message{{Role: "user", Content: "load the release-notes skill"}},
})
if err != nil {
t.Fatal(err)
}
if tt.prompt {
if got := len(prompter.requests); got != tt.wantPrompts {
t.Fatalf("approval prompts = %d, want %d", got, tt.wantPrompts)
}
request := prompter.requests[0]
if len(request.Calls) != 1 || request.Calls[0].ToolName != "skill" || request.Calls[0].ApprovalScope != "skill" || request.Calls[0].Args["name"] != "release-notes" {
t.Fatalf("approval request = %#v", request)
}
}
if got := client.calls; got != tt.wantCalls {
t.Fatalf("model calls = %d, want %d", got, tt.wantCalls)
}
var toolResult string
for _, message := range result.Messages {
if message.Role == "tool" && message.ToolCallID == "call_skill_1" {
toolResult = message.Content
break
}
}
if !strings.Contains(toolResult, tt.wantResult) {
t.Fatalf("skill tool result = %q, want it to contain %q", toolResult, tt.wantResult)
}
})
}
}
func TestExplicitSkillActivationBypassesApproval(t *testing.T) {
catalog := testSkillCatalog(t)
client := &skillTestClient{responses: [][]api.ChatResponse{{{Message: api.Message{Role: "assistant", Content: "done"}}}}}
prompter := &skillApprovalPrompter{result: agent.Approval{}}
result, err := (&agent.Session{
Client: client,
Skills: catalog,
ApprovalPrompter: prompter,
}).Run(context.Background(), agent.RunOptions{
Model: "test",
NewMessages: []api.Message{{Role: "user", Content: "draft release notes"}},
SkillName: "release-notes",
})
if err != nil {
t.Fatal(err)
}
if len(prompter.requests) != 0 {
t.Fatalf("explicit activation prompted for approval: %#v", prompter.requests)
}
if len(result.Messages) != 4 || result.Messages[2].ToolName != "skill" || !strings.Contains(result.Messages[2].Content, "Use concise bullets.") {
t.Fatalf("synthetic skill activation = %#v", result.Messages)
}
}
func testSkillCatalog(t *testing.T) *agent.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 := agent.DiscoverSkills(dir)
if err != nil {
t.Fatal(err)
}
return catalog
}
type skillTestClient struct {
responses [][]api.ChatResponse
calls int
}
func (c *skillTestClient) Chat(_ context.Context, _ *api.ChatRequest, fn api.ChatResponseFunc) error {
if c.calls >= len(c.responses) {
return nil
}
for _, response := range c.responses[c.calls] {
if err := fn(response); err != nil {
return err
}
}
c.calls++
return nil
}
type skillApprovalPrompter struct {
requests []agent.ApprovalRequest
result agent.Approval
}
func (p *skillApprovalPrompter) PromptApproval(_ context.Context, request agent.ApprovalRequest) (agent.Approval, error) {
p.requests = append(p.requests, request)
return p.result, nil
}
-186
View File
@@ -1,186 +0,0 @@
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.",
})
}
-214
View File
@@ -1,214 +0,0 @@
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)
}
})
}
}
-654
View File
@@ -1,654 +0,0 @@
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
System string
Format string
Options map[string]any
Think *api.ThinkValue
KeepAlive *api.Duration
ContextWindowTokens int
AllowAllTools bool
ToolsDisabled bool
MultiModal bool
}
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) {
// Unlike `ollama run`, the bare `ollama` root command doesn't define
// --insecure, so GetBool would error; treat it as false.
insecure, _ := cmd.Flags().GetBool("insecure")
info, resolved, err := showOrPullModel(cmd, client, opts.Model, insecure, "run")
if err != nil {
return nil, err
}
// The model may have been resolved to a different name (e.g. its
// ":cloud" variant).
opts.Model = resolved
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)
}
var skillCatalog *coreagent.SkillCatalog
reloadSkills := func() (*coreagent.SkillCatalog, error) {
catalog, err := coreagent.LoadDefaultSkills(cwd)
if err != nil {
return nil, err
}
if ignored := catalog.ExcludeNames(agentchat.BuiltinSlashCommandNames()); len(ignored) > 0 {
fmt.Fprintf(os.Stderr, "\033[1mwarning:\033[0m ignoring agent skill(s): %s\n", strings.Join(ignored, ", "))
}
for _, diagnostic := range catalog.Diagnostics() {
fmt.Fprintf(os.Stderr, "\033[1mwarning:\033[0m ignored invalid agent skill: %v\n", diagnostic)
}
skillCatalog = catalog
return catalog, nil
}
if _, err := reloadSkills(); err != nil {
return fmt.Errorf("load agent skills: %w", err)
}
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,
ImportSkills: coreagent.ImportSkills,
ReloadSkills: reloadSkills,
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 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: strings.TrimSpace(model.Description),
Recommended: model.Recommended,
AvailabilityBadge: model.AvailabilityBadge,
})
}
return items
}
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 := launch.LoadedContextWindow(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 launch.SameModelRef(modelName, rec.Model) {
return rec.ContextLength
}
}
return 0
}
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 launch.LoadedContextWindow(ctx, client, modelName), nil
}
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)
}
-189
View File
@@ -1,189 +0,0 @@
package cmd
import (
"errors"
"os"
"path/filepath"
"strconv"
"strings"
"testing"
"time"
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 TestAgentSkillCommandCollisionsAreIgnored(t *testing.T) {
dir := t.TempDir()
for _, name := range []string{"release-notes", "system", "exit"} {
if err := os.Mkdir(filepath.Join(dir, name), 0o755); err != nil {
t.Fatal(err)
}
content := "---\nname: " + name + "\ndescription: Test skill.\n---\nInstructions."
if err := os.WriteFile(filepath.Join(dir, name, "SKILL.md"), []byte(content), 0o644); err != nil {
t.Fatal(err)
}
}
catalog, err := coreagent.DiscoverSkills(dir)
if err != nil {
t.Fatal(err)
}
ignored := catalog.ExcludeNames(agentchat.BuiltinSlashCommandNames())
if got, want := strings.Join(ignored, ", "), "exit, system"; got != want {
t.Fatalf("ignored skills = %q, want %q", got, want)
}
if _, err := catalog.Load("release-notes"); err != nil {
t.Fatalf("non-conflicting skill should remain available: %v", err)
}
if context := catalog.SystemContext(); !strings.Contains(context, "release-notes: Test skill.") || strings.Contains(context, "system: Test skill.") || strings.Contains(context, "exit: Test skill.") {
t.Fatalf("skill context = %q", context)
}
for _, name := range []string{"system", "exit"} {
if _, err := catalog.Load(name); err == nil {
t.Fatalf("conflicting skill %q should be ignored", name)
}
}
}
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 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)
}
}
+31 -11
View File
@@ -2128,32 +2128,52 @@ Environment Variables:
}
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
}
opts := agentTUIOptions{
Model: modelName,
Options: map[string]any{},
}
info, err := prepareAgentModel(cmd, client, &opts, false)
info, resolvedModel, err := showOrPullModel(cmd, client, modelName, false, "run")
if err != nil {
if handleCloudAuthorizationError(err) {
return nil
}
return err
}
opts.System = info.System
opts.Model = resolvedModel
ensureCloudStub(cmd.Context(), client, opts.Model)
if err := saveLastAgentModel(opts.Model); err != nil {
opts.Think, err = inferThinkingOption(&info.Capabilities, &opts, false)
if err != nil {
return err
}
if err := GenerateAgentTUI(cmd, client, opts); err != nil {
if handleCloudAuthorizationError(err) {
return nil
audioCapable := slices.Contains(info.Capabilities, model.CapabilityAudio)
opts.MultiModal = slices.Contains(info.Capabilities, model.CapabilityVision) || audioCapable
if len(info.ProjectorInfo) != 0 {
opts.MultiModal = true
}
for key := range info.ModelInfo {
if strings.Contains(key, ".vision.") {
opts.MultiModal = true
break
}
return fmt.Errorf("error running agent: %w", err)
}
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 nil
}
+41
View File
@@ -2390,3 +2390,44 @@ func TestIsLocalhost(t *testing.T) {
})
}
}
func TestRunCommandHasNoAgentFlags(t *testing.T) {
root := NewCLI()
run, _, err := root.Find([]string{"run"})
if err != nil {
t.Fatal(err)
}
for _, name := range []string{"resume", "headless", "auto-approve-tools", "skill", "experimental", "experimental-yolo", "experimental-websearch"} {
if flag := run.Flags().Lookup(name); flag != nil {
t.Errorf("run command still exposes former agent flag --%s", name)
}
}
}
func TestFormerAgentEntryPointsAreRejected(t *testing.T) {
tests := [][]string{
{"run", "llama3", "--resume"},
{"run", "llama3", "--headless"},
{"run", "llama3", "--auto-approve-tools"},
{"run", "llama3", "--skill", "release-notes"},
{"run", "llama3", "--experimental"},
{"run", "llama3", "--experimental-yolo"},
{"run", "llama3", "--experimental-websearch"},
{"agent"},
}
for _, args := range tests {
t.Run(strings.Join(args, " "), func(t *testing.T) {
root := NewCLI()
root.SetArgs(args)
err := root.Execute()
if err == nil {
t.Fatalf("former agent entry point %q succeeded", args)
}
if !strings.Contains(err.Error(), "unknown") {
t.Fatalf("former agent entry point %q returned %v, want unknown command or flag", args, err)
}
})
}
}
-189
View File
@@ -1,189 +0,0 @@
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')
}
-223
View File
@@ -1,223 +0,0 @@
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)
}
}
}
-475
View File
@@ -1,475 +0,0 @@
package chat
import (
"context"
"encoding/json"
"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 summary := editApprovalSummary(call.Args); summary != "" {
lines = append(lines, strings.Split(summary, "\n")...)
}
return chatMetaStyle.Render(strings.Join(lines, "\n"))
default:
if len(call.Args) == 0 {
return ""
}
return strings.Join(renderToolCallArgs(call.Args, width), "\n")
}
}
// editApprovalSummary summarizes edit tool arguments for the approval dialog:
// the edit count and total sizes for the edits array form, or the legacy
// top-level old_text/new_text sizes.
func editApprovalSummary(args map[string]any) string {
edits := rawEditEntries(args)
if len(edits) > 0 {
oldChars, newChars := 0, 0
for _, entry := range edits {
if oldText, ok := editEntryText(entry, "old_text"); ok {
oldChars += len([]rune(oldText))
}
if newText, ok := editEntryText(entry, "new_text"); ok {
newChars += len([]rune(newText))
}
}
return fmt.Sprintf("edits: %d (old: %d chars, new: %d chars)", len(edits), oldChars, newChars)
}
if raw, ok := args["edits"]; ok {
if s, ok := raw.(string); ok {
return fmt.Sprintf("edits: %d chars", len([]rune(s)))
}
}
var lines []string
if oldText, ok := rawStringArg(args, "old_text"); ok {
lines = append(lines, fmt.Sprintf("old_text: %d chars", len([]rune(oldText))))
}
if newText, ok := rawStringArg(args, "new_text"); ok {
lines = append(lines, fmt.Sprintf("new_text: %d chars", len([]rune(newText))))
}
return strings.Join(lines, "\n")
}
// rawEditEntries extracts the edits array entries from tool arguments,
// tolerating a JSON string encoding and camelCase keys.
func rawEditEntries(args map[string]any) []map[string]any {
raw, ok := args["edits"]
if !ok {
return nil
}
if s, ok := raw.(string); ok {
var decoded []map[string]any
if err := json.Unmarshal([]byte(s), &decoded); err != nil {
return nil
}
return decoded
}
items, ok := raw.([]any)
if !ok {
return nil
}
entries := make([]map[string]any, 0, len(items))
for _, item := range items {
if entry, ok := item.(map[string]any); ok {
entries = append(entries, entry)
}
}
return entries
}
func editEntryText(entry map[string]any, snake string) (string, bool) {
if value, ok := entry[snake].(string); ok {
return value, true
}
camel := strings.TrimSuffix(snake, "_text") + "Text"
if value, ok := entry[camel].(string); ok {
return value, true
}
return "", false
}
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
}
-520
View File
@@ -1,520 +0,0 @@
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 TestChatApprovalRendersSkillLoad(t *testing.T) {
request := coreagent.ApprovalRequest{
WorkingDir: "/repo",
Calls: []coreagent.ApprovalToolCall{{
ToolCallID: "call-skill-1",
ToolName: "skill",
Args: map[string]any{"name": "release-notes"},
ApprovalScope: "skill",
}},
}
lines := stripANSI(strings.Join((&chatModel{approvalPrompt: &chatApprovalPrompt{request: request}}).renderApprovalPromptLines(80), "\n"))
for _, want := range []string{"name: release-notes", "2. Always allow skill"} {
if !strings.Contains(lines, want) {
t.Fatalf("skill approval prompt missing %q:\n%s", want, lines)
}
}
m := chatModel{}
m.upsertApprovalToolEntries(request)
if len(m.entries) != 1 || !strings.Contains(stripANSI(toolStatusLine(m.entries[0])), `skill("release-notes") needs approval`) {
t.Fatalf("skill approval entry = %#v", m.entries)
}
}
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")
}
}
-1246
View File
File diff suppressed because it is too large Load Diff
-66
View File
@@ -1,66 +0,0 @@
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
}
-434
View File
@@ -1,434 +0,0 @@
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()
}
-247
View File
@@ -1,247 +0,0 @@
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)
}
}
-101
View File
@@ -1,101 +0,0 @@
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)
}
-565
View File
@@ -1,565 +0,0 @@
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
lines []string
linesWidth 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
}
if m.promptDebug.lines != nil && m.promptDebug.linesWidth == width {
return m.promptDebug.lines
}
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"))
} else {
for i, tool := range req.Tools {
if i > 0 {
lines = append(lines, "")
}
lines = append(lines, promptDebugToolLines(i+1, tool, innerWidth)...)
}
}
m.promptDebug.lines = lines
m.promptDebug.linesWidth = width
return m.promptDebug.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()
}
-374
View File
@@ -1,374 +0,0 @@
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 != "" {
if event.Tokens > 0 {
m.thinkingTokens = max(m.thinkingTokens, event.Tokens)
} else {
m.thinkingTokens += approximateTokenCount(event.Thinking)
}
idx := m.ensureLiveAssistantMessage()
if !m.thinking {
m.thinkingPhaseStart = len(m.liveMessages[idx].Thinking)
}
m.thinking = true
m.liveMessages[idx].Thinking += event.Thinking
m.syncThinkingEntry(m.liveMessages[idx].Thinking[m.thinkingPhaseStart:])
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()
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()
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()
}
-454
View File
@@ -1,454 +0,0 @@
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 TestApplyAgentEventStreamsThinkingThenCollapsesOnAssistantOrTool(t *testing.T) {
m := chatModel{running: true}
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventThinkingDelta, Thinking: "first "})
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventThinkingDelta, Thinking: "second", Tokens: 7})
if len(m.entries) != 1 || m.entries[0].role != "thinking" || !m.entries[0].expanded || m.entries[0].content != "first second" {
t.Fatalf("live thinking entry = %#v", m.entries)
}
if got := m.liveMessages[0].Thinking; got != "first second" {
t.Fatalf("live message thinking = %q, want full streamed value", got)
}
if view := stripANSI(m.renderTranscript(100)); !strings.Contains(view, "Thinking ↓ 7 tokens") || !strings.Contains(view, "first second") {
t.Fatalf("live thinking trace missing from transcript:\n%s", view)
}
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventMessageDelta, Content: "answer"})
if m.entries[0].status != "done" || m.entries[0].expanded {
t.Fatalf("assistant content should collapse thinking: %#v", m.entries[0])
}
collapsed := stripANSI(m.renderTranscript(100))
if !strings.Contains(collapsed, "Thought") || strings.Contains(collapsed, "7 tokens") || strings.Contains(collapsed, "first second") {
t.Fatalf("collapsed thinking should remain as a thought row without trace content:\n%s", collapsed)
}
if got := m.liveMessages[0].Thinking; got != "first second" {
t.Fatalf("collapsing display must not change request history: %q", got)
}
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventThinkingDelta, Thinking: "tool plan"})
if entry := m.entries[len(m.entries)-1]; entry.role != "thinking" || entry.content != "tool plan" {
t.Fatalf("second thinking phase should contain only its own deltas: %#v", entry)
}
if got := m.liveMessages[0].Thinking; got != "first secondtool plan" {
t.Fatalf("message history should retain both thinking phases exactly: %q", got)
}
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventToolStarted, ToolCallID: "call-1", ToolName: "bash"})
if entry := m.entries[len(m.entries)-2]; entry.role != "thinking" || entry.status != "done" || entry.expanded {
t.Fatalf("tool transition should collapse thinking: %#v", entry)
}
}
func TestApplyAgentEventDoesNotCreateThinkingEntryWithoutThinking(t *testing.T) {
m := chatModel{running: true}
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventThinkingDelta, Tokens: 12})
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventMessageDelta, Content: "answer"})
if len(m.entries) != 1 || m.entries[0].role != "assistant" {
t.Fatalf("empty thinking event should not create a trace: %#v", m.entries)
}
if len(m.liveMessages) != 1 || m.liveMessages[0].Thinking != "" {
t.Fatalf("empty thinking event should not alter message history: %#v", m.liveMessages)
}
}
func TestApplyAgentEventPreservesCollapsedThoughtsAcrossToolGrouping(t *testing.T) {
m := chatModel{running: true}
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventThinkingDelta, Thinking: "first plan", Tokens: 1})
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventToolStarted, ToolCallID: "call-1", ToolName: "bash"})
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventToolFinished, ToolCallID: "call-1", ToolName: "bash", Content: "one"})
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventThinkingDelta, Thinking: "second plan", Tokens: 1})
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventToolStarted, ToolCallID: "call-2", ToolName: "bash"})
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventToolFinished, ToolCallID: "call-2", ToolName: "bash", Content: "two"})
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventToolStarted, ToolCallID: "call-3", ToolName: "bash"})
if len(m.entries) != 5 {
t.Fatalf("entries = %#v, want two thought rows and three tool rows", m.entries)
}
for _, index := range []int{0, 2} {
entry := m.entries[index]
if entry.role != "thinking" || entry.status != "done" || entry.expanded {
t.Fatalf("collapsed thought %d = %#v", index, entry)
}
}
if transcript := stripANSI(m.renderTranscript(100)); strings.Count(transcript, "Thought") != 2 || strings.Contains(transcript, "1 token") {
t.Fatalf("transcript should retain both thought rows:\n%s", transcript)
}
}
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
-425
View File
@@ -1,425 +0,0 @@
package chat
import (
"strings"
"unicode"
"unicode/utf8"
"github.com/charmbracelet/lipgloss"
"github.com/mattn/go-runewidth"
)
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(renderMarkdownRunes(parseMarkdownInline(heading))))
continue
}
if trimmed == "" {
rendered = append(rendered, "")
continue
}
rendered = append(rendered, wrapMarkdownInline(line, width)...)
}
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
}
type markdownInlineStyle uint8
const (
markdownPlain markdownInlineStyle = iota
markdownStrong
markdownCode
)
type markdownInlineRune struct {
r rune
style markdownInlineStyle
}
// wrapMarkdownInline parses a complete source line before wrapping it. That
// keeps emphasis intact when its opening and closing delimiters land on
// different visual lines.
func wrapMarkdownInline(line string, width int) []string {
return wrapInlineRunes(parseMarkdownInline(line), width)
}
func wrapInlineRunes(runes []markdownInlineRune, width int) []string {
if len(runes) == 0 {
return []string{""}
}
var rendered []string
for len(runes) > 0 {
hardCut, spaceCut, currentWidth := 0, 0, 0
for i, item := range runes {
nextWidth := currentWidth + runewidth.RuneWidth(item.r)
if nextWidth > width {
break
}
currentWidth = nextWidth
hardCut = i + 1
if unicode.IsSpace(item.r) && currentWidth > width/2 {
spaceCut = i
}
}
cut := hardCut
if spaceCut > 0 {
cut = spaceCut
}
if cut == 0 {
cut = 1
}
lineRunes := trimMarkdownSpace(runes[:cut])
rendered = append(rendered, renderMarkdownRunes(lineRunes))
runes = trimMarkdownSpace(runes[cut:])
}
return rendered
}
func parseMarkdownInline(line string) []markdownInlineRune {
var out []markdownInlineRune
for len(line) > 0 {
if strings.HasPrefix(line, "`") {
if end := strings.Index(line[1:], "`"); end >= 0 {
out = appendMarkdownRunes(out, line[1:end+1], markdownCode)
line = line[end+2:]
continue
}
}
if (strings.HasPrefix(line, "**") || strings.HasPrefix(line, "__")) && canOpenMarkdownStrong(out) {
delimiter := line[:2]
if end := strings.Index(line[2:], delimiter); end >= 0 {
out = appendMarkdownRunes(out, line[2:end+2], markdownStrong)
line = line[end+4:]
continue
}
}
r, size := utf8.DecodeRuneInString(line)
out = append(out, markdownInlineRune{r: r, style: markdownPlain})
line = line[size:]
}
return out
}
// canOpenMarkdownStrong keeps delimiter-like text in bare URLs and identifiers
// literal, only treating ** / __ as strong emphasis at the common
// whitespace- or punctuation-delimited form.
func canOpenMarkdownStrong(out []markdownInlineRune) bool {
if len(out) == 0 {
return true
}
previous := out[len(out)-1].r
return (unicode.IsSpace(previous) || unicode.IsPunct(previous)) && !markdownStrongInURL(out)
}
func markdownStrongInURL(out []markdownInlineRune) bool {
start := len(out)
for start > 0 && !unicode.IsSpace(out[start-1].r) {
start--
}
var token strings.Builder
for _, item := range out[start:] {
token.WriteRune(item.r)
}
return strings.Contains(token.String(), "://")
}
func appendMarkdownRunes(out []markdownInlineRune, text string, style markdownInlineStyle) []markdownInlineRune {
for _, r := range text {
out = append(out, markdownInlineRune{r: r, style: style})
}
return out
}
func trimMarkdownSpace(runes []markdownInlineRune) []markdownInlineRune {
start, end := 0, len(runes)
for start < end && unicode.IsSpace(runes[start].r) {
start++
}
for end > start && unicode.IsSpace(runes[end-1].r) {
end--
}
return runes[start:end]
}
func renderMarkdownRunes(runes []markdownInlineRune) string {
var b strings.Builder
for start := 0; start < len(runes); {
end := start + 1
for end < len(runes) && runes[end].style == runes[start].style {
end++
}
var text strings.Builder
for _, item := range runes[start:end] {
text.WriteRune(item.r)
}
switch runes[start].style {
case markdownStrong:
b.WriteString(chatStrongStyle.Render(text.String()))
case markdownCode:
b.WriteString(chatInlineCodeStyle.Render(text.String()))
default:
b.WriteString(text.String())
}
start = end
}
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], markdownInlineWidth(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 {
lines := wrapInlineRunes(parseMarkdownInline(cell), max(1, width))
if len(lines) == 0 {
return []string{""}
}
return lines
}
// markdownInlineWidth reports the visible width of a cell once Markdown
// delimiters are parsed away, so columns size to rendered content.
func markdownInlineWidth(cell string) int {
width := 0
for _, item := range parseMarkdownInline(cell) {
width += runewidth.RuneWidth(item.r)
}
return width
}
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.TrimSpace(cell)
cell = strings.TrimPrefix(cell, ":")
cell = strings.TrimSuffix(cell, ":")
if cell == "" || strings.Trim(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)
}
-263
View File
@@ -1,263 +0,0 @@
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)
}
-441
View File
@@ -1,441 +0,0 @@
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])
}
}
-307
View File
@@ -1,307 +0,0 @@
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
-113
View File
@@ -1,113 +0,0 @@
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, "")
}
-129
View File
@@ -1,129 +0,0 @@
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)
chatStrongStyle = 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()
)
-166
View File
@@ -1,166 +0,0 @@
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
View File
@@ -46,8 +46,8 @@ type menuItem struct {
}
var runModelMenuItem = menuItem{
title: "Chat, Code, & Work",
description: "Chat with models, code, search the web, and delegate real work",
title: "Chat with a model",
description: "Start an interactive chat with a model",
isRunModel: true,
}
+2 -2
View File
@@ -114,8 +114,8 @@ func TestMenuRendersRootLaunchChoices(t *testing.T) {
view := menu.View()
for _, want := range []string{
"Chat, Code, & Work",
"Chat with models, code, search the web, and delegate real work",
"Chat with a model",
"Start an interactive chat with a model",
"Launch Claude Code",
"Launch OpenCode",
"Launch Hermes Agent",
+1 -1
View File
@@ -23,7 +23,6 @@ require (
github.com/agnivade/levenshtein v1.1.1
github.com/charmbracelet/bubbletea v1.3.10
github.com/charmbracelet/lipgloss v1.1.0
github.com/charmbracelet/x/ansi v0.10.1
github.com/d4l3k/go-bfloat16 v0.0.0-20211005043715-690c3bdd05f1
github.com/dlclark/regexp2 v1.11.5
github.com/emirpasic/gods/v2 v2.0.0-alpha
@@ -49,6 +48,7 @@ require (
github.com/buger/jsonparser v1.1.1 // indirect
github.com/bytedance/sonic/loader v0.1.1 // indirect
github.com/charmbracelet/colorprofile v0.2.3-0.20250311203215-f60798e515dc // indirect
github.com/charmbracelet/x/ansi v0.10.1 // indirect
github.com/charmbracelet/x/cellbuf v0.0.13-0.20250311204145-2c3ea96c31dd // indirect
github.com/charmbracelet/x/term v0.2.1 // indirect
github.com/chewxy/hm v1.0.0 // indirect