Compare commits

...
Author SHA1 Message Date
ParthSareen b0b7bf26f3 refactor: simplify usage command 2026-07-27 12:14:37 -04:00
ParthSareen 90dd5b3a70 feat(cli): add usage command 2026-07-23 12:58:14 -07:00
ParthSareen 4855c61358 feat(api): add account usage endpoint 2026-07-23 12:58:08 -07:00
Parth Sareen 573386c35e agent: skills system (#17203) 2026-07-17 10:32:22 -07:00
Eva H 794a254111 anthropic: close text block before starting thinking block (#17225) 2026-07-17 10:24:38 -04:00
Parth Sareen 714b6fc2a4 agent: allow unlimited tool rounds for cloud models by default (#17217) 2026-07-16 19:07:01 -07:00
Parth Sareen 61e1b1ba5e agent: clean up semantics, UX, DX, and procedural code (#17212) 2026-07-16 19:06:06 -07:00
Parth Sareen 5865a01e48 agent: reorder working directory instruction (#17228) 2026-07-16 19:05:30 -07:00
Parth Sareen e61c1c73fe cmd: remove dead agent prompt wrappers (#17227) 2026-07-16 17:16:14 -07:00
Eva H 03d61e1925 launch: keep Claude Code channels available (#17210) 2026-07-16 13:08:50 -04:00
Parth Sareen 30c390384e cmd: put current working dir in the system prompt (#17188) 2026-07-15 12:10:28 -07:00
Parth Sareen d590830091 agent/tools: isolate web tests from cloud policy (#17208) 2026-07-15 12:07:05 -07:00
Parth Sareen fdcf9efafd fix launch model picker recovery (#17170) 2026-07-15 11:31:22 -07:00
Parth Sareen 76188f60cd docs: add VS Code extension setup (#17158) 2026-07-15 11:30:43 -07:00
Daniel Hiltgen 8a0016f826 model: align gemma4 chat template handling (#17182)
Incorporate the upstream Gemma4 chat template refinements for tool-calling stability, turn closure, and multi-turn reasoning. This updates the native renderer and checked-in HF template fixtures to keep adjacent assistant/tool continuations in the same model turn, add the post-tool thought-channel cue when thinking is enabled, and match Google's default of not replaying historical thinking before a later user turn.

Also preserve null tool arguments through Gemma4 rendering/parsing and extend the Jinja2 parity coverage for these upstream behaviors.
2026-07-14 15:42:04 -07:00
Michael Yang d49b96d9ab docs: collapsed previous retirements (#17167) 2026-07-14 14:00:10 -07:00
Parth Sareen 3bd506bd1c agent/tools: surface actionable web auth error (#17169) 2026-07-14 13:35:16 -07:00
Jesse Gross 123b1f2479 mlxrunner: raise the MTP pending-flush cap to 256 tokens
Per-token cost of the batched head forward keeps falling until the
flush is large enough to reach the fastest kernels: NAX matmul tiles
for dense heads, and the segmented gather path for MoE heads, which
needs tokens*topK/experts >= 4. Measured across the qwen3.6 heads,
256 is the smallest cap past every threshold and within a few percent
of each head's per-token floor. The cost is bounded: up to 2.5 MiB of
pinned hiddens per request and a flush stall under one decode step.
2026-07-14 10:32:04 -07:00
Jesse Gross 556245843a mlxrunner: key the cache trie by token pairs for draft caches
A draft cache pairs each slot with the token that follows it, so the
deepest stored pair always names one token past what a prefix match
can verify - at generation end, the sampled-but-never-committed
final token. Restoring at the match point reuses that pair blind: a
stop token stripped from the next prompt, or any divergence at the
boundary, leaves it stale, and pairing never rewrites below the
resume position, quietly lowering draft acceptance.

Key the trie by token pairs instead: the key for offset i packs
(token i, token i+1), so matching k keys verifies k+1 tokens and
every match is a valid restore point. A pair is reused only if the
token it names matched, and prefill re-evaluates the boundary token,
rebuilding its pair with the token that actually follows. A token
gets a key only once its successor is recorded, so endings record the
final sampled token - never forwarded - and the trie stays level with
the caches. Without a look-ahead the keys are the tokens and behavior
is unchanged.

The recorded tokens' slice bounds used to reject state past them for
free; close now checks the invariant against the stored keys
directly. The test harness rests requests the way the pipeline does -
the deepest recorded token never enters the caches.
2026-07-14 10:32:04 -07:00
Jesse Gross c963822dca mlxrunner: construct per-model state at load
The caches, the speculation binding, and the drafter were each built
lazily inside the first request: begin constructed the caches, and open
bound the cache partition and made a fresh drafter every time. All of
it is a property of the loaded model, so build it once at load.
newPrefixCache replaces the lazy construction in begin, speculation
binds when it is created, and the drafter splits the way speculation
does: a persistent mtpDrafter constructed at load opens each request's
mtpDraftSession, whose constructor syncs the pairing cursor to the
draft caches' restored offset.
2026-07-14 10:32:04 -07:00
Jesse Gross dd49563d55 mlxrunner: rename kvCache to prefixCache
The type coordinates every cache kind — KV, sliding-window, recurrent —
around prefix matching over the trie, so kv was a misnomer.
2026-07-14 10:32:04 -07:00
Jesse Gross 4e96f4dbf2 cache: stop recurrent conv state from pinning the forward buffer
Keeping the recurrent conv state small was handled unevenly: the committed
live state was recopied on every commit — wasted work on single-token
decode, where the window is already tiny — while boundary states captured as
snapshots could still be plain slices of the forward-sized convolution
buffer. A cached slice pins that whole buffer even though the trie's eviction
accounting only counts the slice's bytes, so recurrent cache memory piled up
across requests and eviction could never reclaim it.

Compact each boundary state to its real size once, where it is produced in
the conv wrapper, so live state and snapshots own only their own bytes and
eviction sees the true cost. Single-token decode leaves the already-tiny
window as a slice.

Fixes #16698
2026-07-14 10:32:04 -07:00
Jesse Gross d573a2367b nn/recurrent: derive conv boundary states from a single conv pass
The MTP validation forward schedules a snapshot at every drafted token, which
made CausalConv1D re-run the depthwise conv once per segment to recover each
boundary's conv tail. A conv boundary state is just the trailing convTail
input positions, so run the conv once over the whole window and slice each
boundary tail from the shared buffer, removing the per-token conv launches.
2026-07-14 10:32:04 -07:00
frob 4f7786d0ba mlx: configurable model load timeout (#14796) 2026-07-13 16:06:29 -07:00
Parth Sareen f1a0ffd621 launch: rename Codex App integration to ChatGPT (#17161) 2026-07-13 14:34:17 -07:00
Parth Sareen cd600e19a3 cmd/tui: simplify integration selection and update menu description (#17159) 2026-07-13 12:52:40 -07:00
Daniel Hiltgen 59bd0b49bb mlx: restore NAX in Metal v4 builds (#17160)
MLX now requires a macOS 26.2 deployment target for NAX kernels. Ollama's Metal v4 build still targeted 26.0, so recent MLX bumps silently built mlx_metal_v4 without NAX kernels.
2026-07-13 12:50:38 -07:00
Parth Sareen 82f905cd9c cmd: agent UI (#17017) 2026-07-09 17:27:31 -07:00
Parth Sareen cb3d98ccb2 launch: warn before launching old agent models (#17063) 2026-07-09 16:55:10 -07:00
Jesse Gross d47859ce49 create: select the qwen3.5 parser and renderer for Qwen3.5/Next
Qwen3.5/Qwen3-Next architecture strings contain the substring "qwen3", so the
broad qwen3 match claimed them for the generic parser and qwen3-coder
renderer, whose template doesn't frame the thinking block — an empty
<think></think> leaked into content and think=false was ignored. Match the
family first via isQwen35Family so the parser, renderer, and
thinking-capability checks share one variant list.
2026-07-08 11:12:48 -07:00
Daniel Hiltgen a6293eb516 llm: allow iGPU mmproj offload with fit padding (#16996)
* llm: allow iGPU mmproj offload with fit padding

llama.cpp's fit pass sizes text-model placement before the multimodal projector is loaded. Ollama had been avoiding that risk on non-Metal iGPUs by disabling projector offload entirely, which forces CLIP onto CPU on GB10 and Strix Halo even when the projector has ample memory available.

Let integrated GPUs use the same projector-memory check as other GPUs. When projector offload is enabled, add the estimated projector memory plus the existing 1 GiB headroom to Ollama-owned LLAMA_ARG_FIT_TARGET so fit leaves space for the later projector allocation. If Ollama/device setup already supplied a fit target, add the projector pad to it. If the user set LLAMA_ARG_FIT_TARGET explicitly, leave it exactly as provided.

Fixes #16419

* review comments
2026-07-07 15:28:42 -07:00
Arkadeep Dutta 892e7f6be6 server: apply format constraint for all thinking parsers when think=false (#15901) 2026-07-07 11:54:50 -07:00
Patrick Devine f3d69a3dee server: remove unused internal/ code (#17071) 2026-07-07 11:44:38 -07:00
Daniel Hiltgen 67b6a1c2d4 create: harden GGUF create flows (#17062)
* create: harden GGUF create flows

* lint
2026-07-06 16:20:20 -07:00
Parth Sareen 87b64213b4 launch: disable claude code telemetry by default (#17061) 2026-07-06 15:24:11 -07:00
Daniel Hiltgen f2d069f6df mlx: update to de7b4ed9 (#17056) 2026-07-06 13:31:22 -07:00
Michael Yang 5208ae7500 server: remove OLLAMA_EXPERIMENT=client2 (#16962) 2026-07-06 13:15:39 -07:00
Daniel Hiltgen 9d779572a7 llama.cpp update (#17055)
Bump to b9888.
2026-07-06 12:52:15 -07:00
Patrick Devine 964ea42c09 mlx: x/create rewrite (#16919)
This is a rewrite of the create functionality for the MLX engine.

The core idea behind the create functionality is to break the import/convert into a pipeline of distinct phases:

* Read (scan the safetensors directory for the various bits of metadata)
* Classify (determine what the import type)
* Plan (determine any transforms that need to be done)
* Write (transform any data as necessary and write out the blobs)
* Create the manifest

Each architecture has a "policy" which determines how to convert the model correctly. A number of different formats for safetensors are supported including:

* nvfp4 (two formats: model optimized, torch)
* fp8 datatypes (convert to mxfp8)
* standard bf16 based weights

A number of cleanups/simplifications have been done including:

* using the baked in names for the tensors instead of munging them into something else
* unified 3d expert tensors (instead of separate per expert tensors)
* fewer unnecessary transforms to the various tensors in a model (keep a model as close to the source as possible)
* unified capability checking
* draft model handling (for MTP) is done on the same path

Image generation has been intentionally removed.
2026-07-03 18:30:45 -07:00
Daniel Hiltgen dba1e27fa8 llama: enable FA on CUDA CC 6.x GPUs (#16994)
Recent upstream Pascal kernel fixes let us compile native SM60/SM61 kernels again instead of relying on PTX JIT, so allow Flash Attention auto at runtime for CC 6.x devices.

Fixes #16591

Fixes #16754
2026-07-02 17:11:39 -07:00
Daniel Hiltgen e436db25ff compat: use UTF-8-safe file open (#16999)
Use ggml_fopen for compat tensor reads so Windows paths with Unicode characters are converted through the same UTF-8-to-wide path as llama.cpp model loading.

Fixes #16493
2026-07-02 16:59:23 -07:00
Daniel Hiltgen 26acfa42b5 rocm: remove no longer supported devices (#17010)
The presets and docs had fallen out of sync with what our current ROCm versions on Linux and Windows actually support.  We rely on Vulkan now to cover these older unsupported devices.
2026-07-02 16:59:01 -07:00
Daniel Hiltgen 7b22ac9683 llama: clean up dead code from llama-server work (#17007)
These pieces were missed in the final merge of llama-server and are dead code.
2026-07-02 12:51:54 -07:00
Parth Sareen a2b3a5e9a3 agent: harness core (#16963) 2026-07-02 11:44:31 -07:00
Kevin Park 624cada952 discover: fall back to standard CUDA when the JetPack runner is absent (#16949)
* discover: use the SBSA CUDA build on JetPack 7 (L4T r38+)

JetPack 7 supports SBSA-based CUDA, so the standard cuda_v13 build — shipped
in the base linux-arm64 package, and given the Orin arch (CC 8.7) in #16628
— runs on these devices.

JETSON_JETPACK=7 previously selected a nonexistent jetpack7 runner, so
runner.go skipped every CUDA library and discovery fell back to CPU. The L4T
releases JetPack 7 uses (r38 on Thor, r39 on Orin) also hit the unrecognized
branch, and install.sh warned the version was unsupported. Map JetPack 7+
(L4T r38 and newer) to cuda_v13 (returned as "" from cudaJetpack); no
Jetson-specific download is needed, so install.sh no longer warns.

Fixes #16602

* discover: fall back to standard CUDA when the JetPack runner is absent

Per review, drop the L4T-version mapping (in cudaJetpack and install.sh) and
instead clear the jetpack override in runner.go when the detected cuda_jetpack
runner isn't installed. Normal discovery then selects the standard cuda_v13
build, which supports Orin (CC 8.7) on JetPack 7.
2026-07-02 08:34:35 -07:00
Michael Yang cecd265d3a docs(cloud): update retirement list (#17000) 2026-07-01 19:43:14 -07:00
Mark Ward 2ea95fb059 fix cuda toolkit lookup and parallel (#16613)
* fix cuda toolkit lookup and parallel

* support user override first

* enable control over the nested parallel count
2026-06-30 10:56:54 -07:00
Daniel Hiltgen 8e7be3aed1 ci: avoid unbounded parallelism (#16966)
build-darwin has gotten very slow in the past few releases, most likely due to unbounded parallelism in the MLX build causing the builder to thrash
2026-06-30 10:49:55 -07:00
Patrick Devine 710292ff4f mlx: tighten up gemma4 moe loading code (#16964)
This change allows .experts.gate_proj / .up_proj / .down_proj tensor names to each
be used for both quantized (i.e. nvfp4 and mxfp8) and non-quantized (bf16) models.
Previous to this only non-quantized models used that tensor naming scheme.
2026-06-29 21:15:08 -07:00
Bruce MacDonald ada1eb5163 launch: check for min version for hermes desktop (#16912) 2026-06-29 11:50:11 -07:00
Daniel Hiltgen 1c5ebbf5f4 llama.cpp update (#16960) 2026-06-29 09:43:41 -07:00
Daniel Hiltgen 7926b99e0e mlx: bump dependency (#16935)
Update MLX to 548dd80.

Fix direct MLX tests to run on pinned MLX threads so test execution matches the runner's MLX thread-affinity model.
2026-06-29 09:39:11 -07:00
Aditya Aggarwal 32a97b7493 tools: ignore braces inside JSON strings when detecting tool call end (#16937)
Parser.done() counted the tag's open/close characters ({}, []) without
tracking JSON string context, so a streamed tool call whose string
argument value contained a closing brace or bracket (e.g.
{"code": "if (x) { y }"}) was treated as complete too early and flushed to
the user as plain text instead of being parsed as a tool call.

findArguments() in the same file already tracks string context; apply the
same handling in done() so open/close characters inside string values are
ignored.
2026-06-27 12:00:55 -07:00
Daniel Hiltgen d26a58557d MLX: wire up scheduler selected context size for ps (#16918)
In the PS output, expose the scheduler selected size (clamped by model context size) instead of always reporting the model max context.  This will help provide a hint to clients to keep the context size below this value to avoid paging and poor performance on smaller VRAM systems.
2026-06-26 08:47:03 -07:00
Parth Sareen 2e474c98f9 parser/renderer: add Ornith 9B renderer/parser support (#16920) 2026-06-25 23:18:47 -07:00
Bruce MacDonald 2cb2c5381f launch: update hermes install urls to official (#16913) 2026-06-25 16:22:19 -07:00
Eva H 2a6b50421a fix capability grid dark mode style (#16907) 2026-06-25 13:55:39 -04:00
Daniel Hiltgen f22ec2ec49 CUDA: require driver 550 or newer for v12 (#16895)
Our cuda_v12 build requires nvcc fatbin compression, which in turn requires driver 550 or newer.  This change filters incompatible CUDA devices based on the runtime and driver version.  This allows users to build from source with older toolkits to support older drivers.

Fixes #16449
2026-06-25 08:46:00 -07:00
Eva H d9075caf1a docs: redesign coding integration docs (#16808) 2026-06-25 10:03:59 -04:00
Daniel Hiltgen e11eeb3ba0 llama.cpp version update (#16548) 2026-06-24 14:03:12 -07:00
Daniel Hiltgen 0a408b2225 jetson: add CC 87 for CUDA v13 (#16628)
The new Jetpack 7.2 supports SBSA based CUDA, so we can add the architecture now.
2026-06-24 14:02:41 -07:00
Daniel Hiltgen 16739dee60 server: align generate with native chat templates (#16878)
* server: align generate with native chat templates

/api/generate rebuilt chat-like prompts through the Go template path even when the model selected its native GGUF Jinja chat template, so the same model rendered differently between generate and chat.

Route chat-like generate requests through the shared native chat preparation path, keep deprecated context and image handling working there, and keep explicit OLLAMA_GO_TEMPLATE overrides intact.

Fixes #16792

* review comments

Fall back to "{{ .Prompt }}" when lacking templates
2026-06-24 13:43:56 -07:00
Eva HandParth Sareen d48d790baf docs: redesign docs landing and integrations overview (#16807)
Co-authored-by: Parth Sareen <parth.sareen@ollama.com>
2026-06-24 16:28:28 -04:00
Philip Sinitsin 0463940334 llm: fix ollama ps double-counting mmap'd weights on partial offload (#16709)
* llm: fix ollama ps double-counting mmap'd weights on partial offload

With mmap enabled, llama-server reports each CPU_Mapped model buffer as the
file-offset span of its CPU-resident tensors. During partial offload that span
covers nearly the whole file because the first and last tensors stay on CPU, so
the parsed buffer sizes count the offloaded weights twice and ollama ps shows
roughly 2x the real size with a false CPU/GPU split. Model weights can never
exceed the model file on disk, so trim the excess over the file size from the
mmap-backed portion when computing MemorySize. This makes the reported size
independent of use_mmap; VRAM accounting and scheduler placement are unchanged.

* llm: exclude repacked model buffers from the mmap overlap trim

The trim that corrects mmap double-counting computed the overlap from all
model buffers, including real copies such as CPU_REPACK. On a CPU-only
repacked model that inflated the excess and trimmed the repack out,
undercounting by the repack size (llama3.2 reported ~1918 MiB instead of
~3218 MiB).

Compute the overlap from file-backed buffers only: mmap views and direct
device copies, whose spans can overlap the file on partial offload.
Repacked or host-pinned CPU copies are separate allocations that never
overlap the on-disk weights, so leave them intact. Adds a CPU_Mapped +
CPU_REPACK regression test and corrects the Metal case to the real total.
2026-06-24 11:43:20 -07:00
Daniel Hiltgen 570679c9e0 mlx: update and fix CUDA JIT packaging (#16871)
Bump MLX to the latest selected upstream ref and update the MLX/imagegen
wrappers and tests for the new API behavior.

Fix the CUDA MLX archive so runtime NVRTC kernels work after deployment:
package CUTE/CUTLASS headers, include the CUDA runtime header closure, and
stage a coherent CUDA-toolkit-matched CCCL tree instead of MLX's fetched CCCL
for CUDA payloads. The previous archive could build successfully but crash at
runtime due to missing or incompatible JIT headers.
2026-06-24 10:36:02 -07:00
Daniel Hiltgen 89a171cc70 llm: use host Vulkan loader on Windows (#16869)
Stop bundling the Vulkan loader and resolve the host runtime for Windows Vulkan discovery and backend dependency loading.

Fixes #16677
2026-06-24 10:35:48 -07:00
Daniel Hiltgen 33878e671a llama: default qwen2.5vl window attention metadata (#16868)
Existing qwen2.5vl GGUFs can contain an empty qwen25vl.vision.fullatt_block_indexes array. The compat layer translated the projector metadata but left clip.vision.n_wa_pattern unset, causing llama-server to fail loading the CLIP model.

Default the runtime compat value to the standard Qwen2.5-VL pattern when the key cannot be derived, and make the converter emit the same default for nil or empty fullatt block metadata.

Fixes #16540
2026-06-24 10:35:29 -07:00
Parth SareenandDaniel Hiltgen c191a145bb llm: preserve generation headroom for shifted prompts (#16856)
---------

Co-authored-by: Daniel Hiltgen <daniel@ollama.com>
2026-06-23 15:29:40 -07:00
Parth Sareen 479e1cf94e docs: document max think level (#16877) 2026-06-23 15:29:15 -07:00
Daniel Hiltgen 836507378b llm: size mmproj offload by projector memory (#16866)
* llm: size mmproj offload by projector memory

Replace the blanket 10 GiB VRAM cutoff with a projector tensor-size estimate plus backend headroom, while preserving the existing CPU-only, partial text offload, shared-memory GPU, and startup OOM retry gates.

This is a stopgap until fit accounts for mmproj memory directly.

The same limited-vram path appears in the qwen3.5 vision hang report: the logs show --no-mmproj-offload on a 7.5 GiB RTX 5050 with about 6.4 GiB free while llama-server estimates the inline mmproj at about 962 MiB.

Fixes #16496

Fixes #16570

* review comments
2026-06-23 13:04:02 -07:00
anishandanish 46bc1bcb4c llama: add sm_86 architecture to cuda_v13_windows preset (#16834)
The llama_cuda_v13_windows preset in llama/server/CMakePresets.json was missing sm_86 and sm_80 architectures, causing RTX 3060 laptop and similar mobile RTX 30-series GPUs to be skipped during runtime GPU detection on Windows with CUDA 13. The Linux preset (llama_cuda_v13_linux) included these architectures as "86-virtual" and "80-virtual", but the Windows preset only had "75-virtual;89-virtual;100-virtual;120-virtual", excluding Ampere mobile GPUs.

Signed-off-by: anish <anishesg@users.noreply.github.com>
Co-authored-by: anish <anishesg@users.noreply.github.com>
2026-06-23 07:35:21 -07:00
Bruce MacDonald 2a8b31531e launch/codex: detect model drift when Codex App UI switches away from Ollama (#16864)
ollama launch codex-app sets root-level model_provider = "ollama-launch-codex-app"
in ~/.codex/config.toml to route requests through the local Ollama server.
In Codex, model_provider is a global config key, there is no per-model provider
in the catalog schema (ModelInfo has no model_provider field), so it applies to
every model, not just Ollama ones.

When a user switches to a built-in OpenAI model (e.g. gpt-5.5) in the Codex App
UI, the UI writes model = "gpt-5.5" to config.toml but does NOT update
model_provider. The root model_provider stays "ollama-launch-codex-app", so the
OpenAI model request goes to http://localhost:11434/v1/responses instead of
OpenAI API, resulting in a 404 ("model gpt-5.5 not found"). The user is
stuck: OpenAI models silently route to localhost until they know to run
"ollama launch codex-app --restore".

Fix: CurrentModel() now verifies the configured model appears as a slug in the
Ollama-managed catalog before reporting the integration as active. When the
model has drifted (user selected a non-Ollama model in the UI), CurrentModel()
returns empty, so the launcher accurately shows the integration as inactive and
the user is directed to restore or re-launch.
2026-06-22 15:38:19 -07:00
Jesse Gross 505e35f2b9 mlxrunner: choose the speculative draft length to maximize throughput
The heuristic schedule grew the draft toward a fixed cap on acceptance alone,
maximizing accepted-tokens-per-step rather than throughput, and on a
steep-forward target it regressed below no speculation. Replace it with an
engine-level controller that drafts the depth maximizing
committed-tokens-per-wallclock from live per-position acceptance and persisted
per-width forward cost, with no draft-length cap; the heuristic schedule and
the OLLAMA_MLX_MTP_* env vars go with it.
2026-06-22 15:25:45 -07:00
Jesse Gross 114875133b mlxrunner: resolve each speculative round in one host sync
Acceptance took two blocking evals per round: one to read the accepted mask,
then a second for the bonus or residual token whose graph needed the
host-known rejection point. Sample the residual at every rejection point in
one batched draw alongside the bonus row, so a single eval covers acceptance
and the next token.
2026-06-22 15:25:45 -07:00
Jesse Gross 42c330283b mlxrunner: run one target forward per MTP decode step
Each speculative round ran the target stack twice — once for the current
token's hidden and base logits, once to validate the drafts — capping
throughput below plain decode. Fuse them into one forward over [current,
draft_0..draft_{N-1}], whose hidden rows already line up with the acceptance
math, so the separate base-logits unembed disappears from the drafted path.
2026-06-22 15:25:45 -07:00
Jesse Gross f93efe2809 mlxrunner: apply in-flight drafts to proposal penalty history
Sampler.Distribution built row i as if draftTokens[:i] were appended, leaving
a single-row proposal call with no draft history, so a drafter skipped the
repeat/presence penalties the target's validation applies and re-proposed
penalized tokens. Align rows with the end of the draft chain instead: the
final row sees every draft token, each earlier row one fewer.
2026-06-22 15:25:45 -07:00
Jesse Gross 28fbbb06d5 mlxrunner: support draft heads that maintain draft caches
Generalize the draft path so a head that maintains a KV cache (EAGLE-style)
and Gemma's read-only single-position assistant both fit one drafter
interface with no per-model branches, and make the committed stream the
drafter's maintenance mechanism — every committed run is reported, the
drafter pairs each draft slot with its look-ahead token and flushes completed
pairs to the draft caches. The draft KV thus stays prefix-cached alongside
the target in every session, drafting or not.
2026-06-22 15:25:45 -07:00
Jesse Gross 340c51bbb7 mlxrunner: host speculative decoding in the text generation pipeline
The pipeline and the MTP decoder each owned a decode loop with duplicated
prefill, budget, and emission handling. Split the pipeline into prefill and
decode phases behind a decoder interface, with the decode loop the sole
emitter enforcing the NumPredict budget, and split speculation into a generic
engine that returns the accepted run and a drafter interface that owns only
how proposals are made.
2026-06-22 15:25:45 -07:00
Jesse Gross 2e9d68dc38 mlxrunner: unify the MTP decode paths
Greedy is a special case of sampled decoding — at temperature 0 the sampler
yields a point mass, so rejection-sampling acceptance reduces to argmax-match
— so collapse the separate greedy, sampled, and serial paths into one. MTP
now honors any temperature, penalty, and top-k/p/min-p setting; logprobs
remain the only gated feature.
2026-06-22 15:25:45 -07:00
Sahil Kadadekar fc58544422 discover: fix inverted iGPU/dGPU Vulkan classification on Windows hybrid graphics (#16669)
On Windows hybrid-graphics systems (Intel iGPU + NVIDIA dGPU), discovery
could classify the integrated GPU as discrete and the discrete GPU as
integrated, dropping the dGPU's Vulkan device and scheduling models onto
the iGPU's shared system RAM (#16667). Two index-keyed correlations
between independently-ordered device enumerations caused this:

1. The native probe's stderr was concatenated into the output passed to
   parseVulkanUMA. The probe enumerates Vulkan devices in its own order,
   so its ggml_vulkan uma lines overwrote llama-server's index-keyed UMA
   map with inverted values. Parse UMA metadata only from llama-server's
   own output.

2. applyWindowsVulkanRefinement required the raw vkEnumeratePhysicalDevices
   count to equal llama-server's Vulkan device count. The raw enumeration
   is a superset on real systems (D3D12 mapping-layer devices, Microsoft
   Basic Render Driver), so the refinement that reads the authoritative
   VkPhysicalDeviceType was always skipped. Match devices by name against
   the probed superset instead, bailing only when a device has no match or
   matches conflicting device types.

Verified on the hardware from #16667 (Intel RaptorLake-S + RTX 4080
Laptop): the raw probe returns 5 devices vs llama-server's 2; with this
change the iGPU is dropped as integrated, the dGPU's Vulkan device
dedupes against CUDA0, and the model loads on the dGPU with no
environment overrides.

Fixes #16667
2026-06-22 14:52:03 -07:00
Eva H e434a93884 launch: auto-install opencode when missing (#16806) 2026-06-19 10:12:11 -07:00
Eva H 9c02d8e69d launch: auto-install Claude Code (#16802) 2026-06-19 10:11:50 -07:00
Eva H 07ed752353 launch: add thinking capability detection to opencode (#15434) 2026-06-18 13:45:16 -04:00
Parth Sareen e1f7f9cbdb ci: pin darwin release xcode (#16788) 2026-06-17 13:01:10 -07:00
Patrick Devine 8c432fc88a llama: update llama.cpp to b9672 (#16775) 2026-06-16 23:15:52 -07:00
Jeffrey Morgan acfb50d9af models: add cohere2_moe (Command A / North) to the MLX engine (#16670)
Implements Cohere2MoeForCausalLM (e.g. CohereLabs/North-Mini-Code-1.0)
2026-06-16 23:15:21 -07:00
Jeffrey Morgan 0f047feef5 llm: context shift allow shiftable prompts (#16764) 2026-06-16 12:55:52 -07:00
Patrick Devine 9e4ed74efe integration: look for the "hf" tool in integration tests (#16765)
The "huggingface-cli" tool is deprecated, so only try to use the "hf" tool.
2026-06-16 11:04:54 -07:00
Jeffrey Morgan bbb40a0a6c server: context shift for context windows larger than 8k, add error when hitting context limit (#16712) 2026-06-15 11:36:50 -07:00
Jeffrey Morgan 993acc7504 model: update lfm2 parser/renderer for optional thinking (#16359) 2026-06-14 20:37:08 -07:00
Jeffrey Morgan 7ea692cb2b llama: update llama.cpp to b9637 (#16609) 2026-06-14 20:05:08 -07:00
Parth Sareen 12e04379cd launch: Fix launch provider drift (#16683) 2026-06-11 17:21:46 -07:00
Parafee41 f8a48df24d llm: decouple prompt caching from context shift (#16639)
This PR separates prompt caching from the public shift request option for native llama-server requests.

Previously, shift controlled two different mechanisms:

context shifting / overflow behavior
per-request llama-server cache_prompt
That meant callers could not request shift: false without also disabling prompt caching.

Fixes #16635
2026-06-11 16:05:24 -07:00
Patrick Devine 82e0ddb6fe mlxrunner: harden linear/embedding layers against over-promotion (#16682)
Adding/Multiplying a tensor by a scalar w/ a different data type
can cause the tensor to be promoted and cause performance issues.

This change adds several guards against over-promotion.
2026-06-11 13:56:25 -07:00
Jesse Gross 1abd56b6e6 mlxrunner: record committed MTP drafts before streaming them
The batched MTP accept paths advance the cache by the whole accepted run
before streaming it to the client. If the stream was cancelled partway
(e.g. the caller disconnects), the loop returned before recording the
remaining accepted tokens, leaving the cache offset ahead of
session.outputs. close() then indexed the token log past its end and
panicked with a slice-bounds error.

Record the whole run to session.outputs before streaming any of it, so a
cancelled stream can no longer desync the cache from the token log.

The same bug is present on main, with identical mechanics: the accept
paths there commit the cache to before+accepted and then stream in a loop
that returns on cancellation before recording the rest.
2026-06-09 00:39:19 -07:00
Jesse Gross ded2db7d86 mlxrunner: capture prefill snapshots across the forward
Prefill no longer splits its batch at each requested snapshot offset. The
session schedules the pending offsets on every cache before prefill, runs the
forward in full-size chunks, and attaches the captured snapshots to the trie
afterward. Offsets the prefill never crosses (it leaves one token for decode
seeding) are dropped instead of materializing a node for tokens never written,
and snapshots from an abandoned prefill are released on session close.
2026-06-09 00:39:19 -07:00
Jesse Gross d00622060f mlxrunner: drive MTP speculation through cache snapshots
Speculation used a parallel hierarchy of wrapper cache types that shadowed
the live caches and reconciled against them on commit. Replace it with
snapshot/restore on the live caches themselves: a cache snapshots itself as
a write crosses each offset, and the runner commits a batched draft by
restoring to the accepted count. The wrappers and the comparison plumbing
around them are gone.

Snapshots are lazy. A KV or rotating capture indexes into the live buffer and
owns no memory until a destructive write forces a copy-out, so rejecting a
draft is free.

Recurrent layers now validate in the same batched pass rather than falling
back to serial. A gated-delta layer reports its interior split offsets and
hands back the recurrent state at each one, which the cache records as a
snapshot.
2026-06-09 00:39:19 -07:00
Jesse Gross 177aefb8a9 nn/recurrent: return per-boundary states from the gated-delta kernels
CausalConv1D and GatedDelta now run their scan in segments cut at optional
WithSnapshotSplits offsets and return the recurrent state at each boundary
instead of just the final state. The output is identical to the unsegmented
scan; segmenting only adds a few kernel launches, not extra recurrence compute.

This lets a batched forward capture interior recurrent state without re-running
the scan, which the cache will use for speculative validation rollback points.
RecurrentCache.Put and the Qwen3.5 layer now thread the boundary-state slices,
committing the final entry as the live state.
2026-06-09 00:39:19 -07:00
Jesse Gross 07588c64ee mlxrunner/cache: split KVCache and RotatingKVCache into their own files
cache.go had grown to hold every cache kind. Move KVCache (and its
speculative wrappers) to kvcache.go and RotatingKVCache (and its
sliding-window mask applier) to rotating.go, leaving cache.go with the
shared interfaces and the Speculation transaction. Pure relocation;
no behavior change.
2026-06-09 00:39:19 -07:00
Jesse Gross 4c97a940ca mlxthread: preserve the original stack when worker work panics
Work that panics on the locked MLX worker goroutine was recovered and
re-raised on the caller, so the printed trace pointed at the re-panic
site in this package rather than the code that actually panicked.

Capture the worker stack at recovery and carry it through a value that
implements error, so the runtime prints the original location in the
fatal trace.
2026-06-09 00:39:19 -07:00
Bruce MacDonald 74cbf1d2c2 docs: omp (#16552)
Add docs for explaining and setting up "oh my pi" (omp)
2026-06-08 11:43:51 -07:00
Bruce MacDonald 5c1e37eb67 docs: hermes desktop (#16549) 2026-06-08 11:43:11 -07:00
Jeffrey Morgan f0078ae476 docs: update docs examples to use Gemma 4 instead of Gemma 3 (#16607) 2026-06-07 12:43:13 -07:00
Jeffrey Morgan 96201a623a Add AGENTS.md and CLAUDE.md to root repository (#16604) 2026-06-07 10:57:59 -07:00
Daniel Hiltgen 9c94c2b11e docs: describe llama.cpp update process (#16603) 2026-06-07 10:27:47 -07:00
Parth Sareen e09b3f9fb5 openai: align models list with tags (#16556) 2026-06-05 17:59:05 -07:00
Bruce MacDonald a0099da2d1 launch: use native Windows Hermes config path (#16558) 2026-06-05 17:29:19 -07:00
Chris Chenandfuleinist 25e0e81e12 docs: update Zod example to use native toJSONSchema (#14746)
Co-authored-by: fuleinist <fuleinist@gmail.com>
2026-06-05 16:21:07 -07:00
Bruce MacDonald 87cff95af8 launch: oh-my-pi (#16410) 2026-06-04 17:49:49 -07:00
Patrick Devine 3ef69ef784 mlx: allow the embedding layer to use the nvfp4 global scale (#16527) 2026-06-04 17:40:01 -07:00
Michael Yang 1a7786be14 docs: add cloud model retirement (#16528) 2026-06-04 15:18:38 -07:00
Bruce MacDonald 3370ff8b1c launch: hermes-desktop app (#16516)
Add support to launch the hermes-desktop app alongside the hermes agent from ollama launch. It will go through the install on first run if hermes-desktop is not already installed.
2026-06-04 11:51:36 -07:00
Daniel Hiltgen 455f57457d llama.cpp version update (#16511)
Bump llama.cpp to b9509, which includes the upstream Gemma 4 12B multimodal projector fixes for the n_head=0 divide-by-zero crash seen on x86/CUDA/Linux/Windows.

Fixes #16479
Fixes #16489
Fixes #16491
Fixes #16492
Fixes #16495
2026-06-04 08:20:57 -07:00
Bruce MacDonald 1d955ed990 integrations: hermes windows install (#16487) 2026-06-03 17:40:45 -07:00
Eva H d071237131 docs: add Cline CLI integration doc (#16341) 2026-06-03 20:30:01 -04:00
364 changed files with 45778 additions and 26693 deletions

No files matched your search

+16
View File
@@ -39,11 +39,27 @@ jobs:
APPLE_ID: ${{ vars.APPLE_ID }}
MACOS_SIGNING_KEY: ${{ secrets.MACOS_SIGNING_KEY }}
MACOS_SIGNING_KEY_PASSWORD: ${{ secrets.MACOS_SIGNING_KEY_PASSWORD }}
DEVELOPER_DIR: /Applications/Xcode_26.4.1.app/Contents/Developer
CGO_CFLAGS: '-mmacosx-version-min=14.0 -O3'
CGO_CXXFLAGS: '-mmacosx-version-min=14.0 -O3'
CGO_LDFLAGS: '-mmacosx-version-min=14.0 -O3'
steps:
- uses: actions/checkout@v4
- name: Select Xcode 26.4.1
shell: bash
run: |
set -euo pipefail
if [ ! -d "${DEVELOPER_DIR}" ]; then
echo "Missing ${DEVELOPER_DIR}"
ls -1 /Applications | grep '^Xcode' || true
exit 1
fi
sudo xcode-select -s "${DEVELOPER_DIR}"
sw_vers
xcodebuild -version
xcrun --sdk macosx --show-sdk-version
xcrun --find metal
- run: |
echo $MACOS_SIGNING_KEY | base64 --decode > certificate.p12
security create-keychain -p password build.keychain
+21
View File
@@ -0,0 +1,21 @@
# AGENTS.md
## Building
For a full build from the repository root:
```sh
cmake -B build .
cmake --build build --parallel 8
./ollama serve
```
For quick Go-only iteration against an existing native payload:
```sh
go build .
go run . serve
```
See `docs/development.md` for prerequisites, platform notes, GPU backends, and
the full development workflow.
+3
View File
@@ -0,0 +1,3 @@
# CLAUDE.md
See `AGENTS.md` for the shared agent instructions for this repository.
+1 -1
View File
@@ -1 +1 @@
b9493
b9888
+1 -1
View File
@@ -1 +1 @@
2165dc08d7b33258260aa849d39f087d50e62962
de7b4ed986b6d6f55b8ace5e73c24d1ca0bea89b
+5 -5
View File
@@ -77,10 +77,10 @@ ollama launch openclaw
### Chat with a model
Run and chat with [Gemma 3](https://ollama.com/library/gemma3):
Run and chat with [Gemma 4](https://ollama.com/library/gemma4):
```
ollama run gemma3
ollama run gemma4
```
See [ollama.com/library](https://ollama.com/library) for the full list.
@@ -93,7 +93,7 @@ Ollama has a REST API for running and managing models.
```
curl http://localhost:11434/api/chat -d '{
"model": "gemma3",
"model": "gemma4",
"messages": [{
"role": "user",
"content": "Why is the sky blue?"
@@ -113,7 +113,7 @@ pip install ollama
```python
from ollama import chat
response = chat(model='gemma3', messages=[
response = chat(model='gemma4', messages=[
{
'role': 'user',
'content': 'Why is the sky blue?',
@@ -132,7 +132,7 @@ npm i ollama
import ollama from "ollama";
const response = await ollama.chat({
model: "gemma3",
model: "gemma4",
messages: [{ role: "user", content: "Why is the sky blue?" }],
});
console.log(response.message.content);
+198
View File
@@ -0,0 +1,198 @@
package agent
import (
"context"
"strings"
"sync"
)
type ApprovalRequest struct {
WorkingDir string
Calls []ApprovalToolCall
}
func (r *ApprovalRequest) AddToolCall(id, name, scope string, args map[string]any) {
r.Calls = append(r.Calls, ApprovalToolCall{
ToolCallID: id,
ToolName: name,
Args: args,
ApprovalScope: scope,
})
}
type ApprovalToolCall struct {
ToolCallID string
ToolName string
Args map[string]any
ApprovalScope string
}
type Approval struct {
Allow bool
AllowAll bool
AllowScopes []string
Reason string
}
type ApprovalPrompter interface {
PromptApproval(context.Context, ApprovalRequest) (Approval, error)
}
type ApprovalState struct {
mu sync.RWMutex
allowAll bool
scopes map[string]bool
}
func (s *ApprovalState) Set(allowAll bool, scopes map[string]bool) {
if s == nil {
return
}
s.mu.Lock()
defer s.mu.Unlock()
s.allowAll = allowAll
s.scopes = cloneApprovalScopes(scopes)
}
// GrantAll grants blanket approval for all future tool calls.
func (s *ApprovalState) GrantAll() {
if s == nil {
return
}
s.mu.Lock()
defer s.mu.Unlock()
s.allowAll = true
}
// AllGranted reports whether blanket approval has been granted.
func (s *ApprovalState) AllGranted() bool {
if s == nil {
return false
}
s.mu.RLock()
defer s.mu.RUnlock()
return s.allowAll
}
func (s *ApprovalState) Allows(scope string) bool {
if s == nil {
return false
}
s.mu.RLock()
defer s.mu.RUnlock()
return s.allowAll || s.scopes[scope]
}
// Apply merges an approval's scopes and allow-all flag into the state. It
// returns true if the approval grants permission (allow-all or at least one
// scope). It does not mutate the approval; the caller sets Allow based on the
// returned value.
func (s *ApprovalState) Apply(result *Approval) bool {
if s == nil || result == nil {
return false
}
s.mu.Lock()
defer s.mu.Unlock()
granted := false
if result.AllowAll {
s.allowAll = true
granted = true
}
if len(result.AllowScopes) > 0 {
granted = true
s.grantScopesLocked(result.AllowScopes)
}
return granted
}
// GrantScopes merges the given scopes into the state.
func (s *ApprovalState) GrantScopes(scopes []string) {
if s == nil {
return
}
s.mu.Lock()
defer s.mu.Unlock()
s.grantScopesLocked(scopes)
}
// grantScopesLocked adds trimmed, non-empty scopes to the state. Caller must
// hold s.mu.
func (s *ApprovalState) grantScopesLocked(scopes []string) {
if s.scopes == nil {
s.scopes = make(map[string]bool, len(scopes))
}
for _, scope := range scopes {
scope = strings.TrimSpace(scope)
if scope != "" {
s.scopes[scope] = true
}
}
}
func cloneApprovalScopes(src map[string]bool) map[string]bool {
if len(src) == 0 {
return nil
}
dst := make(map[string]bool, len(src))
for scope, allowed := range src {
if allowed {
dst[scope] = true
}
}
return dst
}
func (s *Session) needsApproval(tool Tool, name string, args map[string]any) bool {
return ToolRequiresApproval(tool, args) && !s.allows(toolApprovalScope(tool, name, args))
}
// allows reports whether scope is permitted by the session's accumulated approval state.
func (s *Session) allows(scope string) bool {
if s == nil || s.ApprovalState == nil {
return false
}
return s.ApprovalState.Allows(scope)
}
// applyApproval merges an approval result into the session's state and marks
// the result as allowed when scopes or allow-all were granted.
func (s *Session) applyApproval(result *Approval) {
if s == nil || result == nil {
return
}
if s.ApprovalState == nil {
s.ApprovalState = &ApprovalState{}
}
if s.ApprovalState.Apply(result) {
result.Allow = true
}
}
func (s *Session) authorizeToolCalls(ctx context.Context, req ApprovalRequest) (Approval, error) {
if s == nil || len(req.Calls) == 0 || (s.ApprovalState != nil && s.ApprovalState.AllGranted()) {
return Approval{Allow: true}, nil
}
if s.ApprovalPrompter == nil {
return Approval{
Reason: "Tool execution requires approval, but no approval prompter is available.",
}, nil
}
result, err := s.ApprovalPrompter.PromptApproval(ctx, req)
if err != nil {
return Approval{}, err
}
s.applyApproval(&result)
return result, nil
}
// toolApprovalScope returns the approval scope key for a tool invocation.
// If the tool implements ScopedTool, its ApprovalScope method determines the
// scope (e.g. shell tools scope to "<tool>\x00<command>"). Otherwise the scope
// is the trimmed tool name.
func toolApprovalScope(tool Tool, toolName string, args map[string]any) string {
if scoped, ok := tool.(ScopedTool); ok {
return scoped.ApprovalScope(args)
}
return strings.TrimSpace(toolName)
}
+95
View File
@@ -0,0 +1,95 @@
package agent
import (
"context"
"strings"
"testing"
"github.com/ollama/ollama/api"
)
type mockTool struct {
name string
}
func (m mockTool) Name() string { return m.name }
func (m mockTool) Description() string { return "" }
func (m mockTool) Schema() api.ToolFunction {
return api.ToolFunction{Name: m.name}
}
func (m mockTool) Execute(context.Context, ToolContext, map[string]any) (ToolResult, error) {
return ToolResult{}, nil
}
func TestToolApprovalScopeUsesScopedTool(t *testing.T) {
shellTool := mockScopedTool{
mockTool: mockTool{name: "bash"},
scope: func(args map[string]any) string {
if cmd, ok := args["command"].(string); ok {
cmd = strings.TrimSpace(cmd)
if cmd != "" {
return "bash\x00" + cmd
}
}
return "bash"
},
}
plainTool := mockTool{name: "edit"}
tests := []struct {
tool Tool
name string
args map[string]any
want string
}{
{shellTool, "bash", map[string]any{"command": " pwd "}, "bash\x00pwd"},
{shellTool, "bash", map[string]any{"command": "Get-ChildItem"}, "bash\x00Get-ChildItem"},
{plainTool, "edit", map[string]any{"path": "README.md"}, "edit"},
}
for _, tt := range tests {
if got := toolApprovalScope(tt.tool, tt.name, tt.args); got != tt.want {
t.Fatalf("toolApprovalScope(%q) = %q, want %q", tt.name, got, tt.want)
}
}
}
type mockScopedTool struct {
mockTool
scope func(args map[string]any) string
}
func (m mockScopedTool) ApprovalScope(args map[string]any) string {
return m.scope(args)
}
func TestSessionApplyApprovalScopes(t *testing.T) {
session := &Session{}
result := Approval{AllowScopes: []string{"edit", "bash\x00pwd", " "}}
session.applyApproval(&result)
if !result.Allow {
t.Fatal("scoped approval should allow the current request")
}
if !session.allows("edit") || !session.allows("bash\x00pwd") {
t.Fatal("scoped approval was not saved")
}
if session.allows("bash") || session.allows("bash\x00ls") {
t.Fatal("shell approval was too broad")
}
if session.ApprovalState.AllGranted() {
t.Fatal("allow all = true, want false for scoped approval")
}
}
func TestSessionApplyApprovalAllowAll(t *testing.T) {
session := &Session{}
result := Approval{AllowAll: true}
session.applyApproval(&result)
if !result.Allow || !session.allows("anything") {
t.Fatalf("allow all = %v result = %#v, want allow all", session.ApprovalState.AllGranted(), result)
}
}
+667
View File
@@ -0,0 +1,667 @@
package agent
import (
"context"
"encoding/json"
"errors"
"fmt"
"net/http"
"strings"
"github.com/ollama/ollama/api"
)
// Compaction wire-format. These constants and helpers are the single canonical
// definition of how a compacted turn is represented in message history.
const (
CompactionSummaryMessagePrefix = "Conversation summary:\n"
CompactionToolName = "summary"
CompactionToolCallID = "ollama_compaction"
CompactionContinueInstruction = "continue the task in progress. the history has been compacted, do not mention compaction to the user"
)
const (
defaultCompactionContextWindowTokens = 32768
defaultCompactionKeepUserTurns = 3
defaultCompactionThreshold = 0.8
compactOnlySummaryContextTokens = 16000
maxCompactionSummaryRunes = 16 * 1024
compactionSystemPrompt = "Summarize the archived part of an Ollama agent conversation. Preserve user goals, decisions, files, commands, tool results, and unresolved tasks needed to continue. Omit private reasoning and return only the summary."
)
type Compactor interface {
MaybeCompact(context.Context, CompactionRequest) (CompactionResult, error)
// ContextWindowTokens returns the effective context window size in
// tokens, resolving runtime options against configured defaults.
ContextWindowTokens(options map[string]any) int
// Threshold returns the compaction threshold as a fraction of the
// context window (e.g. 0.8 means compact at 80% capacity).
Threshold() float64
// ShouldCompact reports whether a compaction should run and returns the
// trigger reason. An empty trigger means compaction is not needed.
ShouldCompact(req CompactionRequest) (trigger string, should bool)
}
type CompactionOptions struct {
ContextWindowTokens int
KeepUserTurns int
Threshold float64
}
type CompactionRequest struct {
ChatID string
Model string
SystemPrompt string
Messages []api.Message
Tools api.Tools
Format string
Latest api.ChatResponse
Options map[string]any
KeepAlive *api.Duration
Think *api.ThinkValue
Force bool
ContinueTask bool
KeepUserTurns *int
Progress func(CompactionProgress)
}
type CompactionProgress struct {
Tokens int
}
type CompactionResult struct {
Messages []api.Message
Compacted bool
Due bool
Summary string
Reason string
}
type SimpleCompactor struct {
Client ChatClient
Options CompactionOptions
}
func (c *SimpleCompactor) MaybeCompact(ctx context.Context, req CompactionRequest) (CompactionResult, error) {
result := CompactionResult{Messages: req.Messages}
if c == nil {
return result, nil
}
result.Due = req.Force || c.shouldCompact(req)
if !result.Due {
return result, nil
}
if c.Client == nil {
result.Reason = "compaction is unavailable"
return result, nil
}
keepUserTurns := c.keepUserTurns(req.Options)
if req.KeepUserTurns != nil {
keepUserTurns = *req.KeepUserTurns
}
prefix, previousSummary, archive, suffix, _, ok := splitCompactionMessages(req.Messages, keepUserTurns)
if !ok || len(archive) == 0 {
result.Reason = "nothing to compact"
return result, nil
}
summary, err := c.summarize(ctx, req, previousSummary, archive)
if err != nil {
result.Reason = err.Error()
return result, err
}
summary = truncateCompactionSummary(strings.TrimSpace(summary))
if summary == "" {
summary, err = c.summarizeEmptyFallback(ctx, req, previousSummary, archive)
if err != nil {
result.Reason = err.Error()
return result, err
}
summary = truncateCompactionSummary(strings.TrimSpace(summary))
}
if summary == "" {
result.Reason = "summary was empty"
return result, nil
}
compacted := make([]api.Message, 0, len(prefix)+len(suffix)+2)
compacted = append(compacted, prefix...)
compacted = append(compacted, CompactionSummaryMessages(summary, req.ContinueTask)...)
compacted = append(compacted, suffix...)
result.Messages = compacted
result.Compacted = true
result.Summary = summary
return result, nil
}
func (c *SimpleCompactor) shouldCompact(req CompactionRequest) bool {
contextWindow := c.contextWindowTokens(req.Options)
threshold := int(float64(contextWindow) * c.threshold())
if threshold <= 0 {
return false
}
if req.Latest.PromptEvalCount > 0 && req.Latest.PromptEvalCount >= threshold {
return true
}
return estimateCompactionRequestTokens(req) >= threshold
}
func (c *SimpleCompactor) contextWindowTokens(options map[string]any) int {
return ResolveContextWindowTokens(options, c.Options.ContextWindowTokens)
}
// ContextWindowTokens resolves the effective context window from runtime
// options or configured defaults. Satisfies the Compactor interface.
func (c *SimpleCompactor) ContextWindowTokens(options map[string]any) int {
if c == nil {
return 0
}
return c.contextWindowTokens(options)
}
func (c *SimpleCompactor) threshold() float64 {
return ResolveCompactionThreshold(c.Options.Threshold)
}
// Threshold returns the configured compaction threshold fraction. Satisfies
// the Compactor interface.
func (c *SimpleCompactor) Threshold() float64 {
if c == nil {
return 0
}
return c.threshold()
}
// ShouldCompact reports whether compaction is due and the trigger reason.
// Satisfies the Compactor interface.
func (c *SimpleCompactor) ShouldCompact(req CompactionRequest) (string, bool) {
if c == nil {
return "", false
}
if req.Force {
return "force", true
}
if c.shouldCompact(req) {
contextWindow := c.contextWindowTokens(req.Options)
threshold := int(float64(contextWindow) * c.threshold())
if req.Latest.PromptEvalCount > 0 && req.Latest.PromptEvalCount >= threshold {
return "prompt_eval", true
}
return "estimate", true
}
return "", false
}
func (c *SimpleCompactor) keepUserTurns(options map[string]any) int {
contextWindow := c.contextWindowTokens(options)
if contextWindow > 0 && contextWindow < compactOnlySummaryContextTokens {
return 0
}
if c.Options.KeepUserTurns > 0 {
return c.Options.KeepUserTurns
}
return defaultCompactionKeepUserTurns
}
func ResolveContextWindowTokens(options map[string]any, configured int) int {
if n := intOption(options, "num_ctx"); n > 0 {
return n
}
if configured > 0 {
return configured
}
return defaultCompactionContextWindowTokens
}
func ResolveCompactionThreshold(configured float64) float64 {
if configured > 0 {
return configured
}
return defaultCompactionThreshold
}
func (c *SimpleCompactor) summarize(ctx context.Context, req CompactionRequest, previousSummary string, archive []api.Message) (string, error) {
body, err := compactionPrompt(previousSummary, archive, c.compactionPromptBodyBudgetTokens(req.Options))
if err != nil {
return "", err
}
chatReq := &api.ChatRequest{
Model: req.Model,
Messages: []api.Message{
{
Role: "system",
Content: compactionSystemPrompt,
},
{
Role: "user",
Content: body,
},
},
Options: req.Options,
Think: req.Think,
}
if req.KeepAlive != nil {
chatReq.KeepAlive = req.KeepAlive
}
var summary strings.Builder
if err := c.Client.Chat(ctx, chatReq, func(response api.ChatResponse) error {
summary.WriteString(response.Message.Content)
if req.Progress != nil {
tokens := response.EvalCount
if tokens <= 0 {
tokens = estimateCompactionTokens(summary.String())
}
req.Progress(CompactionProgress{Tokens: tokens})
}
return nil
}); err != nil {
return "", err
}
return summary.String(), nil
}
func (c *SimpleCompactor) summarizeEmptyFallback(ctx context.Context, req CompactionRequest, previousSummary string, archive []api.Message) (string, error) {
retry := req
retry.Think = &api.ThinkValue{Value: false}
summary, err := c.summarize(ctx, retry, previousSummary, archive)
if err == nil {
return summary, nil
}
if !isUnsupportedCompactionThinkError(err) {
return "", err
}
if req.Think == nil {
return "", nil
}
retry.Think = nil
return c.summarize(ctx, retry, previousSummary, archive)
}
func isUnsupportedCompactionThinkError(err error) bool {
if err == nil {
return false
}
text := strings.ToLower(err.Error())
if !strings.Contains(text, "think") {
return false
}
var statusErr api.StatusError
if errors.As(err, &statusErr) && statusErr.StatusCode != 0 {
return statusErr.StatusCode == http.StatusBadRequest
}
return strings.Contains(text, "does not support") || strings.Contains(text, "not supported") || strings.Contains(text, "unsupported")
}
// compactionSummaryMessageForTask renders a compaction summary as the content
// string stored on the synthetic tool-result message.
func compactionSummaryMessageForTask(summary string, continueTask bool) string {
content := CompactionSummaryMessagePrefix + strings.TrimSpace(summary)
if continueTask {
content = strings.TrimSpace(content) + "\n\n" + CompactionContinueInstruction
}
return content
}
// CompactionSummaryMessages renders a compaction summary as the assistant
// tool-call plus tool-result pair that represents a compacted turn in the
// message history.
func CompactionSummaryMessages(summary string, continueTask bool) []api.Message {
return []api.Message{
{
Role: "assistant",
ToolCalls: []api.ToolCall{{
ID: CompactionToolCallID,
Function: api.ToolCallFunction{
Name: CompactionToolName,
},
}},
},
{
Role: "tool",
ToolName: CompactionToolName,
ToolCallID: CompactionToolCallID,
Content: compactionSummaryMessageForTask(summary, continueTask),
},
}
}
func (c *SimpleCompactor) compactionPromptBodyBudgetTokens(options map[string]any) int {
contextWindow := c.contextWindowTokens(options)
threshold := int(float64(contextWindow) * c.threshold())
if threshold <= 0 {
return 0
}
systemTokens := estimateCompactionTokens("system") + estimateCompactionTokens(compactionSystemPrompt)
userRoleTokens := estimateCompactionTokens("user")
budget := threshold - systemTokens - userRoleTokens
if budget <= 0 {
return 0
}
return budget
}
func truncateCompactionSummary(summary string) string {
return Truncate(summary, TruncateConfig{
MaxRunes: maxCompactionSummaryRunes,
Label: "summary",
})
}
func estimateCompactionTokens(text string) int {
text = strings.TrimSpace(text)
if text == "" {
return 0
}
return ApproximateTokens(len([]rune(text)))
}
func estimateMessagesTokens(messages []api.Message) int {
var total int
for _, msg := range messages {
total += estimateCompactionTokens(msg.Role)
total += estimateCompactionTokens(msg.Content)
total += estimateCompactionTokens(msg.Thinking)
total += estimateCompactionTokens(msg.ToolName)
total += estimateCompactionTokens(msg.ToolCallID)
for _, call := range msg.ToolCalls {
total += estimateCompactionTokens(call.Function.Name)
total += estimateCompactionTokens(call.Function.Arguments.String())
}
}
return total
}
func estimateCompactionRequestTokens(req CompactionRequest) int {
requestMessages := sanitizeMessagesForEstimate(req.Messages)
if strings.TrimSpace(req.SystemPrompt) != "" {
requestMessages = make([]api.Message, 0, len(req.Messages)+1)
requestMessages = append(requestMessages, api.Message{Role: "system", Content: strings.TrimSpace(req.SystemPrompt)})
requestMessages = append(requestMessages, sanitizeMessagesForEstimate(req.Messages)...)
}
payload := struct {
Messages []api.Message `json:"messages,omitempty"`
Tools api.Tools `json:"tools,omitempty"`
Format json.RawMessage `json:"format,omitempty"`
}{
Messages: requestMessages,
Tools: req.Tools,
}
if rawFormat, ok := compactionFormatForEstimate(req.Format); ok {
payload.Format = rawFormat
}
if data, err := json.Marshal(payload); err == nil {
return estimateCompactionTokens(string(data))
}
total := estimateMessagesTokens(requestMessages)
total += estimateCompactionTokens(req.Tools.String())
total += estimateCompactionTokens(req.Format)
return total
}
func (s *Session) estimateRunPromptTokens(opts RunOptions, messages []api.Message) int {
return estimateCompactionRequestTokens(CompactionRequest{
SystemPrompt: opts.SystemPrompt,
Messages: messages,
Tools: s.availableTools(),
Format: opts.Format,
Options: opts.Options,
})
}
func (s *Session) checkPreflightPromptBudget(opts RunOptions, messages []api.Message) error {
contextWindow := s.contextWindowTokens(opts)
if contextWindow <= 0 {
return nil
}
estimated := s.estimateRunPromptTokens(opts, messages)
if estimated < contextWindow {
return nil
}
return fmt.Errorf("prompt is too large for the current context (~%d/%d tokens). Reduce the system prompt or message history, compact the conversation, or use a model with a larger context", estimated, contextWindow)
}
func (s *Session) checkPostCompactionPromptBudget(opts RunOptions, messages []api.Message) error {
contextWindow := s.contextWindowTokens(opts)
if contextWindow <= 0 {
return nil
}
estimated := s.estimateRunPromptTokens(opts, messages)
if estimated < contextWindow {
return nil
}
return fmt.Errorf("history is still too large after compaction (~%d/%d tokens). Start a fresh request, reduce the system prompt or history, or use a model with a larger context", estimated, contextWindow)
}
func sanitizeMessagesForEstimate(messages []api.Message) []api.Message {
requestMessages := sanitizeMessagesForRequest(messages)
for i := range requestMessages {
// Image token accounting is model-specific. Without the active model's
// tokenizer and vision accounting, raw image bytes/base64 make the
// estimate look much larger than the prompt the model actually sees.
requestMessages[i].Images = nil
}
return requestMessages
}
func compactionFormatForEstimate(format string) (json.RawMessage, bool) {
format = strings.TrimSpace(format)
if format == "" {
return nil, false
}
if format == "json" {
return json.RawMessage(`"json"`), true
}
if !json.Valid([]byte(format)) {
return nil, false
}
return json.RawMessage(format), true
}
func compactionPrompt(previousSummary string, archive []api.Message, maxTokens int) (string, error) {
messages := make([]api.Message, 0, len(archive))
for _, msg := range archive {
msg.Thinking = ""
msg.Images = nil
messages = append(messages, msg)
}
return renderCompactionPrompt(previousSummary, fitCompactionMessagesToBudget(previousSummary, messages, maxTokens))
}
func renderCompactionPrompt(previousSummary string, messages []api.Message) (string, error) {
payload, err := json.MarshalIndent(messages, "", " ")
if err != nil {
return "", fmt.Errorf("marshal compaction messages: %w", err)
}
var b strings.Builder
if strings.TrimSpace(previousSummary) != "" {
b.WriteString("Previous summary:\n")
b.WriteString(strings.TrimSpace(previousSummary))
b.WriteString("\n\n")
}
b.WriteString("Messages to archive as JSON:\n")
b.Write(payload)
return b.String(), nil
}
func fitCompactionMessagesToBudget(previousSummary string, messages []api.Message, maxTokens int) []api.Message {
if maxTokens <= 0 {
return messages
}
fitted := append([]api.Message(nil), messages...)
for range 16 {
body, err := renderCompactionPrompt(previousSummary, fitted)
if err != nil || estimateCompactionTokens(body) <= maxTokens {
return fitted
}
idx := largestCompactionContentMessage(fitted)
if idx < 0 {
return fitted
}
overageTokens := estimateCompactionTokens(body) - maxTokens
currentRunes := len([]rune(fitted[idx].Content))
nextRunes := currentRunes - overageTokens*4 - 256
if nextRunes >= currentRunes {
nextRunes = currentRunes / 2
}
fitted[idx].Content = truncateToolResultContentTo(fitted[idx].Content, nextRunes)
}
return fitted
}
func largestCompactionContentMessage(messages []api.Message) int {
idx := -1
size := 0
for i, msg := range messages {
n := len([]rune(msg.Content))
if n > size {
idx = i
size = n
}
}
return idx
}
func splitCompactionMessages(messages []api.Message, keepUserTurns int) (prefix []api.Message, previousSummary string, archive []api.Message, suffix []api.Message, keptUserTurns int, ok bool) {
if keepUserTurns < 0 {
keepUserTurns = defaultCompactionKeepUserTurns
}
start := 0
for start < len(messages) && messages[start].Role == "system" && !isCompactionSummary(messages[start]) {
prefix = append(prefix, messages[start])
start++
}
candidates := make([]api.Message, 0, len(messages)-start)
for i := start; i < len(messages); i++ {
msg := messages[i]
if isCompactionSummary(msg) {
previousSummary = CompactionSummaryText(msg.Content)
continue
}
if isCompactionToolCall(msg) {
if i+1 < len(messages) && isCompactionSummary(messages[i+1]) {
previousSummary = CompactionSummaryText(messages[i+1].Content)
i++
}
continue
}
candidates = append(candidates, msg)
}
userTurnIndexes := make([]int, 0, keepUserTurns)
for i := len(candidates) - 1; i >= 0; i-- {
if candidates[i].Role == "user" {
userTurnIndexes = append(userTurnIndexes, i)
}
}
keptUserTurns = keepUserTurns
if len(userTurnIndexes) <= keptUserTurns {
keptUserTurns = len(userTurnIndexes) - 1
}
if keptUserTurns < 0 {
keptUserTurns = 0
}
suffixStart := len(candidates)
if keptUserTurns > 0 {
suffixStart = userTurnIndexes[keptUserTurns-1]
}
if suffixStart <= 0 || len(candidates[:suffixStart]) == 0 {
return prefix, previousSummary, nil, nil, keptUserTurns, false
}
return prefix, previousSummary, candidates[:suffixStart], candidates[suffixStart:], keptUserTurns, true
}
func isCompactionToolName(name string) bool {
return name == CompactionToolName
}
func isCompactionSummary(msg api.Message) bool {
return (msg.Role == "user" || msg.Role == "system" || (msg.Role == "tool" && isCompactionToolName(msg.ToolName))) &&
strings.HasPrefix(msg.Content, CompactionSummaryMessagePrefix)
}
// IsCompactionSummary reports whether msg uses the canonical compaction
// summary message representation.
func IsCompactionSummary(msg api.Message) bool {
return isCompactionSummary(msg)
}
// CompactionSummaryContent returns the user-visible summary from msg when it
// is a canonical compaction summary.
func CompactionSummaryContent(msg api.Message) (string, bool) {
if !isCompactionSummary(msg) {
return "", false
}
return CompactionSummaryText(msg.Content), true
}
// IsCompactionToolResult reports whether msg is the synthetic tool result used
// to represent compaction in message history.
func IsCompactionToolResult(msg api.Message) bool {
return msg.Role == "tool" && (isCompactionToolName(msg.ToolName) || msg.ToolCallID == CompactionToolCallID)
}
// IsCompactionToolCall reports whether msg is the synthetic assistant tool
// call paired with a compaction summary result.
func IsCompactionToolCall(msg api.Message) bool {
return isCompactionToolCall(msg)
}
func isCompactionToolCall(msg api.Message) bool {
if msg.Role != "assistant" {
return false
}
for _, call := range msg.ToolCalls {
if isCompactionToolName(call.Function.Name) {
return true
}
}
return false
}
// CompactionSummaryText reverses CompactionSummaryMessages, returning the
// user-visible summary text with the prefix and any continuation instruction
// removed.
func CompactionSummaryText(content string) string {
return strings.TrimSpace(strings.TrimSuffix(
strings.TrimSpace(strings.TrimPrefix(content, CompactionSummaryMessagePrefix)),
CompactionContinueInstruction,
))
}
func intOption(options map[string]any, key string) int {
if options == nil {
return 0
}
switch v := options[key].(type) {
case int:
return v
case int64:
return int(v)
case float64:
return int(v)
case float32:
return int(v)
case json.Number:
n, _ := v.Int64()
return int(n)
default:
return 0
}
}
+773
View File
@@ -0,0 +1,773 @@
package agent
import (
"context"
"net/http"
"strings"
"testing"
"github.com/ollama/ollama/api"
)
type scriptedCompactionClient struct {
responses [][]api.ChatResponse
errs []error
requests []*api.ChatRequest
}
func (c *scriptedCompactionClient) Chat(_ context.Context, req *api.ChatRequest, fn api.ChatResponseFunc) error {
c.requests = append(c.requests, req)
i := len(c.requests) - 1
if i < len(c.responses) {
for _, response := range c.responses[i] {
if err := fn(response); err != nil {
return err
}
}
}
if i < len(c.errs) {
return c.errs[i]
}
return nil
}
func assertCompactionSummaryPair(t *testing.T, messages []api.Message) {
t.Helper()
if len(messages) != 2 {
t.Fatalf("compaction summary pair len = %d, want 2: %#v", len(messages), messages)
}
if messages[0].Role != "assistant" || len(messages[0].ToolCalls) != 1 || messages[0].ToolCalls[0].Function.Name != CompactionToolName {
t.Fatalf("compaction assistant message = %#v", messages[0])
}
if messages[0].ToolCalls[0].Function.Arguments.Len() != 0 {
t.Fatalf("compaction summary tool call should not have arguments: %#v", messages[0].ToolCalls[0].Function.Arguments.ToMap())
}
if messages[1].Role != "tool" || messages[1].ToolName != CompactionToolName || messages[1].ToolCallID != messages[0].ToolCalls[0].ID {
t.Fatalf("compaction tool result = %#v", messages[1])
}
if !strings.HasPrefix(messages[1].Content, CompactionSummaryMessagePrefix) {
t.Fatalf("compaction tool result missing summary prefix: %#v", messages[1])
}
}
func TestSimpleCompactorSummarizesOldMessages(t *testing.T) {
client := &fakeClient{
responses: [][]api.ChatResponse{{
{Message: api.Message{Role: "assistant", Content: "summary"}},
}},
}
compactor := &SimpleCompactor{Client: client, Options: CompactionOptions{
ContextWindowTokens: 16000,
KeepUserTurns: 2,
Threshold: 0.5,
}}
messages := []api.Message{
{Role: "system", Content: "stay pinned"},
{Role: "user", Content: "old request"},
{Role: "assistant", Content: "old answer", Thinking: "hidden"},
{Role: "user", Content: "recent one"},
{Role: "assistant", Content: "recent answer"},
{Role: "user", Content: "recent two"},
}
result, err := compactor.MaybeCompact(context.Background(), CompactionRequest{
ChatID: "chat-1",
Model: "model",
Messages: messages,
Latest: api.ChatResponse{Metrics: api.Metrics{PromptEvalCount: 12000}},
})
if err != nil {
t.Fatal(err)
}
if !result.Compacted {
t.Fatal("expected compaction")
}
compacted := result.Messages
if len(compacted) != 6 {
t.Fatalf("compacted messages = %d, want 6", len(compacted))
}
if compacted[0].Content != "stay pinned" {
t.Fatalf("first message = %#v", compacted[0])
}
if result.Summary != "summary" {
t.Fatalf("result summary = %q", result.Summary)
}
assertCompactionSummaryPair(t, compacted[1:3])
if compacted[3].Content != "recent one" || compacted[5].Content != "recent two" {
t.Fatalf("recent turns were not kept: %#v", compacted)
}
if len(client.requests) != 1 {
t.Fatalf("summary requests = %d, want 1", len(client.requests))
}
if strings.Contains(client.requests[0].Messages[1].Content, "hidden") {
t.Fatal("compaction prompt should omit thinking")
}
}
func TestSimpleCompactorKeepsOnlySummaryForSmallContext(t *testing.T) {
client := &fakeClient{
responses: [][]api.ChatResponse{{
{Message: api.Message{Role: "assistant", Content: "small context summary"}},
}},
}
compactor := &SimpleCompactor{Client: client, Options: CompactionOptions{
ContextWindowTokens: compactOnlySummaryContextTokens - 1,
KeepUserTurns: 3,
Threshold: 0.5,
}}
result, err := compactor.MaybeCompact(context.Background(), CompactionRequest{
ChatID: "chat-1",
Model: "model",
ContinueTask: true,
Messages: []api.Message{
{Role: "system", Content: "pinned"},
{Role: "user", Content: "old request"},
{Role: "assistant", Content: "old answer"},
{Role: "user", Content: "latest request"},
},
Latest: api.ChatResponse{Metrics: api.Metrics{PromptEvalCount: 12000}},
})
if err != nil {
t.Fatal(err)
}
if !result.Compacted {
t.Fatal("expected compaction")
}
if len(result.Messages) != 3 {
t.Fatalf("messages = %#v, want system plus compaction summary pair", result.Messages)
}
if result.Messages[0].Content != "pinned" {
t.Fatalf("leading system message not kept: %#v", result.Messages)
}
assertCompactionSummaryPair(t, result.Messages[1:])
if !strings.Contains(result.Messages[2].Content, CompactionContinueInstruction) {
t.Fatalf("tool result missing continue instruction: %q", result.Messages[2].Content)
}
}
func TestSimpleCompactorAddsContinueTaskInstructionOnlyToToolResult(t *testing.T) {
client := &fakeClient{
responses: [][]api.ChatResponse{{
{Message: api.Message{Role: "assistant", Content: "summary"}},
}},
}
compactor := &SimpleCompactor{Client: client, Options: CompactionOptions{
ContextWindowTokens: 100,
KeepUserTurns: 1,
Threshold: 0.5,
}}
result, err := compactor.MaybeCompact(context.Background(), CompactionRequest{
ChatID: "chat-1",
Model: "model",
ContinueTask: true,
Messages: []api.Message{
{Role: "user", Content: "old request"},
{Role: "assistant", Content: "old answer"},
{Role: "user", Content: "recent request"},
},
Latest: api.ChatResponse{Metrics: api.Metrics{PromptEvalCount: 75}},
})
if err != nil {
t.Fatal(err)
}
if result.Summary != "summary" {
t.Fatalf("result summary = %q", result.Summary)
}
content := result.Messages[1].Content
if !strings.Contains(content, CompactionContinueInstruction) {
t.Fatalf("tool result missing continue instruction: %q", content)
}
if got := CompactionSummaryText(content); got != "summary" {
t.Fatalf("visible summary text = %q", got)
}
}
func TestSimpleCompactorTruncatesOversizedSummary(t *testing.T) {
longSummary := strings.Repeat("x", maxCompactionSummaryRunes+1024)
client := &fakeClient{
responses: [][]api.ChatResponse{{
{Message: api.Message{Role: "assistant", Content: longSummary}},
}},
}
compactor := &SimpleCompactor{Client: client, Options: CompactionOptions{
ContextWindowTokens: 100,
KeepUserTurns: 1,
Threshold: 0.5,
}}
result, err := compactor.MaybeCompact(context.Background(), CompactionRequest{
ChatID: "chat-1",
Model: "model",
Messages: []api.Message{
{Role: "user", Content: "old one"},
{Role: "assistant", Content: "old answer"},
{Role: "user", Content: "recent one"},
},
Latest: api.ChatResponse{Metrics: api.Metrics{PromptEvalCount: 75}},
})
if err != nil {
t.Fatal(err)
}
if !result.Compacted {
t.Fatal("expected compaction")
}
if runeCount := len([]rune(result.Summary)); runeCount > maxCompactionSummaryRunes+200 {
t.Fatalf("summary runes = %d, want <= %d (plus marker)", runeCount, maxCompactionSummaryRunes)
}
if !strings.Contains(result.Summary, "[summary truncated:") {
t.Fatalf("summary missing truncation marker: %q", result.Summary)
}
if !strings.Contains(result.Messages[1].Content, "[summary truncated:") {
t.Fatalf("compacted message missing truncation marker: %#v", result.Messages)
}
}
func TestSimpleCompactorRetriesEmptySummaryWithThinkFalse(t *testing.T) {
client := &scriptedCompactionClient{
responses: [][]api.ChatResponse{
{{Message: api.Message{Role: "assistant", Thinking: "internal summary plan"}}},
{{Message: api.Message{Role: "assistant", Content: "fallback summary"}}},
},
}
compactor := &SimpleCompactor{Client: client, Options: CompactionOptions{
ContextWindowTokens: 100,
KeepUserTurns: 1,
Threshold: 0.5,
}}
result, err := compactor.MaybeCompact(context.Background(), CompactionRequest{
Model: "model",
Messages: []api.Message{
{Role: "user", Content: "old request"},
{Role: "assistant", Content: "old answer"},
{Role: "user", Content: "recent request"},
},
Force: true,
})
if err != nil {
t.Fatal(err)
}
if !result.Compacted || result.Summary != "fallback summary" {
t.Fatalf("compaction result = %#v", result)
}
if len(client.requests) != 2 {
t.Fatalf("summary requests = %d, want 2", len(client.requests))
}
if client.requests[0].Think != nil {
t.Fatalf("first summary request think = %#v, want nil", client.requests[0].Think)
}
if client.requests[1].Think == nil || client.requests[1].Think.Value != false {
t.Fatalf("fallback summary request think = %#v, want false", client.requests[1].Think)
}
}
func TestSimpleCompactorIgnoresUnsupportedThinkFalseFallback(t *testing.T) {
client := &scriptedCompactionClient{
responses: [][]api.ChatResponse{
{{Message: api.Message{Role: "assistant", Thinking: "internal summary plan"}}},
nil,
},
errs: []error{
nil,
api.StatusError{StatusCode: http.StatusBadRequest, ErrorMessage: "model does not support thinking"},
},
}
compactor := &SimpleCompactor{Client: client, Options: CompactionOptions{
ContextWindowTokens: 100,
KeepUserTurns: 1,
Threshold: 0.5,
}}
result, err := compactor.MaybeCompact(context.Background(), CompactionRequest{
Model: "model",
Messages: []api.Message{
{Role: "user", Content: "old request"},
{Role: "assistant", Content: "old answer"},
{Role: "user", Content: "recent request"},
},
Force: true,
})
if err != nil {
t.Fatal(err)
}
if result.Compacted || result.Reason != "summary was empty" {
t.Fatalf("compaction result = %#v", result)
}
if len(client.requests) != 2 {
t.Fatalf("summary requests = %d, want 2", len(client.requests))
}
if client.requests[1].Think == nil || client.requests[1].Think.Value != false {
t.Fatalf("fallback summary request think = %#v, want false", client.requests[1].Think)
}
}
func TestSimpleCompactorFallsBackToUnsetThinkWhenThinkFalseUnsupported(t *testing.T) {
client := &scriptedCompactionClient{
responses: [][]api.ChatResponse{
{{Message: api.Message{Role: "assistant", Thinking: "internal summary plan"}}},
nil,
{{Message: api.Message{Role: "assistant", Content: "unset think summary"}}},
},
errs: []error{
nil,
api.StatusError{StatusCode: http.StatusBadRequest, ErrorMessage: "think level is not supported"},
nil,
},
}
compactor := &SimpleCompactor{Client: client, Options: CompactionOptions{
ContextWindowTokens: 100,
KeepUserTurns: 1,
Threshold: 0.5,
}}
thinkHigh := &api.ThinkValue{Value: "high"}
result, err := compactor.MaybeCompact(context.Background(), CompactionRequest{
Model: "model",
Messages: []api.Message{
{Role: "user", Content: "old request"},
{Role: "assistant", Content: "old answer"},
{Role: "user", Content: "recent request"},
},
Think: thinkHigh,
Force: true,
})
if err != nil {
t.Fatal(err)
}
if !result.Compacted || result.Summary != "unset think summary" {
t.Fatalf("compaction result = %#v", result)
}
if len(client.requests) != 3 {
t.Fatalf("summary requests = %d, want 3", len(client.requests))
}
if client.requests[0].Think != thinkHigh {
t.Fatalf("first summary request think = %#v, want original", client.requests[0].Think)
}
if client.requests[1].Think == nil || client.requests[1].Think.Value != false {
t.Fatalf("fallback summary request think = %#v, want false", client.requests[1].Think)
}
if client.requests[2].Think != nil {
t.Fatalf("unsupported fallback retry think = %#v, want nil", client.requests[2].Think)
}
}
func TestSimpleCompactorKeepsFewerTurnsForShortChats(t *testing.T) {
client := &fakeClient{
responses: [][]api.ChatResponse{{
{Message: api.Message{Role: "assistant", Content: "short summary"}},
}},
}
compactor := &SimpleCompactor{Client: client, Options: CompactionOptions{
ContextWindowTokens: 16000,
KeepUserTurns: 3,
Threshold: 0.5,
}}
result, err := compactor.MaybeCompact(context.Background(), CompactionRequest{
ChatID: "chat-1",
Model: "model",
Messages: []api.Message{
{Role: "user", Content: "old request"},
{Role: "assistant", Content: "old answer"},
{Role: "user", Content: "latest request"},
},
Latest: api.ChatResponse{Metrics: api.Metrics{PromptEvalCount: 12000}},
})
if err != nil {
t.Fatal(err)
}
if !result.Compacted {
t.Fatal("expected compaction")
}
if len(result.Messages) != 3 {
t.Fatalf("messages = %#v, want compaction tool pair plus latest request", result.Messages)
}
assertCompactionSummaryPair(t, result.Messages[:2])
if result.Messages[2].Content != "latest request" {
t.Fatalf("latest turn was not kept: %#v", result.Messages)
}
}
func TestSimpleCompactorCanArchiveWholeShortChat(t *testing.T) {
client := &fakeClient{
responses: [][]api.ChatResponse{{
{Message: api.Message{Role: "assistant", Content: "whole summary"}},
}},
}
compactor := &SimpleCompactor{Client: client, Options: CompactionOptions{
ContextWindowTokens: 100,
KeepUserTurns: 3,
Threshold: 0.5,
}}
result, err := compactor.MaybeCompact(context.Background(), CompactionRequest{
ChatID: "chat-1",
Model: "model",
Messages: []api.Message{
{Role: "user", Content: "only request"},
{Role: "assistant", Content: "only answer"},
},
Latest: api.ChatResponse{Metrics: api.Metrics{PromptEvalCount: 75}},
})
if err != nil {
t.Fatal(err)
}
if !result.Compacted {
t.Fatal("expected compaction")
}
if len(result.Messages) != 2 {
t.Fatalf("messages = %#v, want only compaction tool pair", result.Messages)
}
assertCompactionSummaryPair(t, result.Messages)
}
func TestSimpleCompactorSkipsBelowThreshold(t *testing.T) {
client := &fakeClient{}
compactor := &SimpleCompactor{Client: client, Options: CompactionOptions{
ContextWindowTokens: 100,
Threshold: 0.8,
}}
messages := []api.Message{
{Role: "user", Content: "one"},
{Role: "user", Content: "two"},
{Role: "user", Content: "three"},
{Role: "user", Content: "four"},
{Role: "user", Content: "five"},
}
result, err := compactor.MaybeCompact(context.Background(), CompactionRequest{
Model: "model",
Messages: messages,
Latest: api.ChatResponse{Metrics: api.Metrics{PromptEvalCount: 50}},
})
if err != nil {
t.Fatal(err)
}
if result.Compacted {
t.Fatal("did not expect compaction")
}
if result.Due {
t.Fatal("below-threshold compaction should not be due")
}
if len(result.Messages) != len(messages) {
t.Fatalf("messages changed below threshold: %#v", result.Messages)
}
if len(client.requests) != 0 {
t.Fatalf("summary requests = %d, want 0", len(client.requests))
}
}
func TestSimpleCompactorUsesEstimatedMessagesWhenPromptEvalMissing(t *testing.T) {
client := &fakeClient{
responses: [][]api.ChatResponse{{
{Message: api.Message{Role: "assistant", Content: "estimated summary"}},
}},
}
compactor := &SimpleCompactor{Client: client, Options: CompactionOptions{
ContextWindowTokens: 100,
KeepUserTurns: 1,
Threshold: 0.8,
}}
result, err := compactor.MaybeCompact(context.Background(), CompactionRequest{
Model: "model",
Messages: []api.Message{
{Role: "user", Content: "old request"},
{Role: "assistant", Content: "old answer"},
{Role: "user", Content: "read large output"},
{Role: "assistant", ToolCalls: []api.ToolCall{{
ID: "call-1",
Function: api.ToolCallFunction{
Name: "read",
},
}}},
{Role: "tool", ToolName: "read", ToolCallID: "call-1", Content: strings.Repeat("x", 360)},
},
})
if err != nil {
t.Fatal(err)
}
if !result.Due || !result.Compacted {
t.Fatalf("expected estimate-driven compaction, got %#v", result)
}
if result.Summary != "estimated summary" {
t.Fatalf("summary = %q", result.Summary)
}
}
func TestSimpleCompactorEstimateIncludesRequestPreamble(t *testing.T) {
compactor := &SimpleCompactor{Client: nil, Options: CompactionOptions{
ContextWindowTokens: 100,
Threshold: 0.8,
}}
if !compactor.shouldCompact(CompactionRequest{
SystemPrompt: strings.Repeat("system ", 360),
Messages: []api.Message{{Role: "user", Content: "tiny"}},
}) {
t.Fatal("system prompt should count toward compaction estimate")
}
if !compactor.shouldCompact(CompactionRequest{
Messages: []api.Message{{Role: "user", Content: "tiny"}},
Tools: api.Tools{{
Type: "function",
Function: api.ToolFunction{
Name: "verbose_tool",
Description: strings.Repeat("description ", 360),
},
}},
}) {
t.Fatal("tool definitions should count toward compaction estimate")
}
}
func TestCompactionPromptFitsBudgetByTruncatingLargeToolOutput(t *testing.T) {
largeToolOutput := strings.Repeat("x", 10_000)
body, err := compactionPrompt("", []api.Message{
{Role: "user", Content: "what changed?"},
{Role: "assistant", ToolCalls: []api.ToolCall{{
ID: "call-1",
Function: api.ToolCallFunction{
Name: "bash",
},
}}},
{Role: "tool", ToolName: "bash", ToolCallID: "call-1", Content: largeToolOutput},
}, 300)
if err != nil {
t.Fatal(err)
}
if estimateCompactionTokens(body) > 300 {
t.Fatalf("compaction prompt tokens = %d, want <= 300", estimateCompactionTokens(body))
}
if strings.Count(body, "x") >= len(largeToolOutput) {
t.Fatal("large tool output was not truncated")
}
if !strings.Contains(body, "[tool output truncated: showing first ~") {
t.Fatalf("truncation marker missing from compaction prompt: %q", body)
}
}
func TestCompactionPromptRetruncatesAlreadyTruncatedToolOutput(t *testing.T) {
alreadyTruncated := strings.Repeat("x", 7000) + "\n\n[tool output truncated: showing first ~100 tokens and last ~100 tokens; omitted ~99999 tokens. Use a narrower command, line range, or search query if more detail is needed.]\n\n" + strings.Repeat("y", 7000)
body, err := compactionPrompt("", []api.Message{
{Role: "user", Content: "what changed?"},
{Role: "assistant", ToolCalls: []api.ToolCall{{
ID: "call-1",
Function: api.ToolCallFunction{
Name: "bash",
},
}}},
{Role: "tool", ToolName: "bash", ToolCallID: "call-1", Content: alreadyTruncated},
}, 300)
if err != nil {
t.Fatal(err)
}
if estimateCompactionTokens(body) > 300 {
t.Fatalf("compaction prompt tokens = %d, want <= 300", estimateCompactionTokens(body))
}
if strings.Count(body, "x")+strings.Count(body, "y") >= 14_000 {
t.Fatal("already-truncated tool output was not truncated again")
}
if !strings.Contains(body, "[tool output truncated: showing first ~") {
t.Fatalf("truncation marker missing from compaction prompt: %q", body)
}
}
func TestCompactionSummaryTextStripsPrefix(t *testing.T) {
content := compactionSummaryMessageForTask("worked on branch changes", false)
if got := CompactionSummaryText(content); got != "worked on branch changes" {
t.Fatalf("summary text = %q", got)
}
}
func TestCompactionSummaryCanTellModelToContinueTask(t *testing.T) {
content := compactionSummaryMessageForTask("worked on branch changes", true)
if !strings.Contains(content, CompactionContinueInstruction) {
t.Fatalf("summary message missing continue instruction: %q", content)
}
if got := CompactionSummaryText(content); got != "worked on branch changes" {
t.Fatalf("summary text = %q", got)
}
}
func TestResolveContextWindowTokensPrefersExplicitNumCtx(t *testing.T) {
tests := []struct {
name string
options map[string]any
configured int
want int
}{
{
name: "explicit smaller num ctx",
options: map[string]any{"num_ctx": 4096},
configured: 8192,
want: 4096,
},
{
name: "explicit num ctx can exceed configured metadata",
options: map[string]any{"num_ctx": 131072},
configured: 8192,
want: 131072,
},
{
name: "metadata without explicit num ctx",
configured: 32768,
want: 32768,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := ResolveContextWindowTokens(tt.options, tt.configured); got != tt.want {
t.Fatalf("ResolveContextWindowTokens() = %d, want %d", got, tt.want)
}
})
}
}
func TestSimpleCompactorForceCompactsWithoutPromptEvalCount(t *testing.T) {
client := &fakeClient{
responses: [][]api.ChatResponse{{
{Message: api.Message{Role: "assistant", Content: "forced summary"}},
}},
}
compactor := &SimpleCompactor{Client: client, Options: CompactionOptions{
ContextWindowTokens: 100,
KeepUserTurns: 1,
Threshold: 0.8,
}}
result, err := compactor.MaybeCompact(context.Background(), CompactionRequest{
Model: "model",
Messages: []api.Message{
{Role: "user", Content: "old"},
{Role: "assistant", Content: "old answer"},
{Role: "user", Content: "recent"},
},
Force: true,
})
if err != nil {
t.Fatal(err)
}
if !result.Due || !result.Compacted {
t.Fatalf("forced compaction result = %#v", result)
}
if result.Summary != "forced summary" {
t.Fatalf("summary = %q", result.Summary)
}
}
func TestSimpleCompactorDefaultsToKeepingThreeUserTurns(t *testing.T) {
client := &fakeClient{
responses: [][]api.ChatResponse{{
{Message: api.Message{Role: "assistant", Content: "summary"}},
}},
}
compactor := &SimpleCompactor{Client: client, Options: CompactionOptions{
ContextWindowTokens: 16000,
Threshold: 0.5,
}}
result, err := compactor.MaybeCompact(context.Background(), CompactionRequest{
ChatID: "chat-1",
Model: "model",
Messages: []api.Message{
{Role: "user", Content: "old"},
{Role: "assistant", Content: "old answer"},
{Role: "user", Content: "one"},
{Role: "assistant", Content: "one answer"},
{Role: "user", Content: "two"},
{Role: "assistant", Content: "two answer"},
{Role: "user", Content: "three"},
},
Latest: api.ChatResponse{Metrics: api.Metrics{PromptEvalCount: 12000}},
})
if err != nil {
t.Fatal(err)
}
if !result.Compacted {
t.Fatal("expected compaction")
}
assertCompactionSummaryPair(t, result.Messages[:2])
if got := result.Messages[2].Content; got != "one" {
t.Fatalf("first kept turn = %q, want one", got)
}
}
func TestSimpleCompactorCarriesPreviousSummary(t *testing.T) {
client := &fakeClient{
responses: [][]api.ChatResponse{{
{Message: api.Message{Role: "assistant", Content: "new summary"}},
}},
}
compactor := &SimpleCompactor{Client: client, Options: CompactionOptions{
ContextWindowTokens: 16000,
KeepUserTurns: 1,
Threshold: 0.5,
}}
result, err := compactor.MaybeCompact(context.Background(), CompactionRequest{
Model: "model",
Messages: []api.Message{
{Role: "system", Content: CompactionSummaryMessagePrefix + "old summary"},
{Role: "user", Content: "old"},
{Role: "assistant", Content: "old answer"},
{Role: "user", Content: "recent"},
},
Latest: api.ChatResponse{Metrics: api.Metrics{PromptEvalCount: 12000}},
})
if err != nil {
t.Fatal(err)
}
if !result.Compacted {
t.Fatal("expected compaction")
}
if !strings.Contains(client.requests[0].Messages[1].Content, "Previous summary:\nold summary") {
t.Fatalf("previous summary missing from request: %q", client.requests[0].Messages[1].Content)
}
}
func TestSimpleCompactorCarriesPreviousToolSummaryAndPlacesNewSummaryBeforeKeptSuffix(t *testing.T) {
client := &fakeClient{
responses: [][]api.ChatResponse{{
{Message: api.Message{Role: "assistant", Content: "new summary"}},
}},
}
compactor := &SimpleCompactor{Client: client, Options: CompactionOptions{
ContextWindowTokens: 16000,
KeepUserTurns: 1,
Threshold: 0.5,
}}
messages := []api.Message{
{Role: "user", Content: "kept before old summary"},
CompactionSummaryMessages("old summary", false)[0],
CompactionSummaryMessages("old summary", false)[1],
{Role: "user", Content: "latest request"},
}
result, err := compactor.MaybeCompact(context.Background(), CompactionRequest{
Model: "model",
Messages: messages,
Latest: api.ChatResponse{Metrics: api.Metrics{PromptEvalCount: 12000}},
})
if err != nil {
t.Fatal(err)
}
if !result.Compacted {
t.Fatal("expected compaction")
}
if !strings.Contains(client.requests[0].Messages[1].Content, "Previous summary:\nold summary") {
t.Fatalf("previous summary missing from request: %q", client.requests[0].Messages[1].Content)
}
if len(result.Messages) != 3 {
t.Fatalf("messages = %#v, want compaction pair plus latest request", result.Messages)
}
assertCompactionSummaryPair(t, result.Messages[:2])
if result.Messages[2].Content != "latest request" {
t.Fatalf("kept suffix = %#v", result.Messages)
}
}
+177
View File
@@ -0,0 +1,177 @@
package agent
import (
"context"
"errors"
"github.com/ollama/ollama/api"
)
type EventType string
const (
EventMessageDelta EventType = "message_delta"
EventThinkingDelta EventType = "thinking_delta"
EventToolCallDetected EventType = "tool_call_detected"
EventToolStarted EventType = "tool_started"
EventToolFinished EventType = "tool_finished"
EventCompactionStarted EventType = "compaction_started"
EventCompactionProgress EventType = "compaction_progress"
EventCompacted EventType = "compacted"
EventCompactionSkipped EventType = "compaction_skipped"
EventRunFinished EventType = "run_finished"
EventError EventType = "error"
)
// ToolStatus is the typed lifecycle state for a tool call, carried on
// Event.ToolStatus for tool events.
type ToolStatus string
const (
ToolStatusRunning ToolStatus = "running"
ToolStatusDone ToolStatus = "done"
ToolStatusFailed ToolStatus = "failed"
ToolStatusDenied ToolStatus = "denied"
ToolStatusDisabled ToolStatus = "disabled"
ToolStatusSkipped ToolStatus = "skipped"
)
// RunStatus is the typed terminal outcome of a run, carried on Event.Status for
// run_finished events.
type RunStatus string
const (
RunStatusDone RunStatus = "done"
RunStatusDenied RunStatus = "denied"
RunStatusCanceled RunStatus = "canceled"
)
// CompactionTrigger is the typed reason a compaction ran or was attempted,
// carried on Event.CompactionTrigger for compaction events.
type CompactionTrigger string
const (
CompactionTriggerForce CompactionTrigger = "force"
CompactionTriggerPromptEval CompactionTrigger = "prompt_eval"
CompactionTriggerEstimate CompactionTrigger = "estimate"
CompactionTriggerToolOutput CompactionTrigger = "tool_output"
CompactionTriggerError CompactionTrigger = "error"
CompactionTriggerDue CompactionTrigger = "due"
)
type Event struct {
Type EventType `json:"type"`
RunID string `json:"runId,omitempty"`
ChatID string `json:"chatId,omitempty"`
Model string `json:"model,omitempty"`
Status RunStatus `json:"status,omitempty"`
ToolStatus ToolStatus `json:"toolStatus,omitempty"`
CompactionTrigger CompactionTrigger `json:"compactionTrigger,omitempty"`
ToolCallID string `json:"toolCallId,omitempty"`
ToolName string `json:"toolName,omitempty"`
WorkingDir string `json:"workingDir,omitempty"`
Content string `json:"content,omitempty"`
Thinking string `json:"thinking,omitempty"`
ToolCalls []api.ToolCall `json:"toolCalls,omitempty"`
Messages []api.Message `json:"messages,omitempty"`
Args map[string]any `json:"args,omitempty"`
Tokens int `json:"tokens,omitempty"`
Error string `json:"error,omitempty"`
}
type EventSink interface {
Emit(Event) error
}
type EventSinkFunc func(Event) error
func (fn EventSinkFunc) Emit(event Event) error {
if fn == nil {
return nil
}
return fn(event)
}
// eventMetadata carries the run identification fields shared by all events.
type eventMetadata struct {
runID string
chatID string
model string
}
func newEventMetadata(runID string, opts RunOptions) eventMetadata {
return eventMetadata{runID: runID, chatID: opts.ChatID, model: opts.Model}
}
func newMessageDelta(m eventMetadata, content string) Event {
return Event{Type: EventMessageDelta, RunID: m.runID, ChatID: m.chatID, Model: m.model, Content: content}
}
func newThinkingDelta(m eventMetadata, thinking string) Event {
return Event{Type: EventThinkingDelta, RunID: m.runID, ChatID: m.chatID, Model: m.model, Thinking: thinking}
}
func newToolCallDetected(m eventMetadata, calls []api.ToolCall) Event {
return Event{Type: EventToolCallDetected, RunID: m.runID, ChatID: m.chatID, Model: m.model, ToolCalls: calls}
}
func newToolStarted(m eventMetadata, callID, toolName, workingDir string, args map[string]any) Event {
return Event{Type: EventToolStarted, RunID: m.runID, ChatID: m.chatID, Model: m.model, ToolStatus: ToolStatusRunning, ToolCallID: callID, ToolName: toolName, WorkingDir: workingDir, Args: args}
}
func newToolFinished(m eventMetadata, status ToolStatus, callID, toolName, workingDir string, args map[string]any, content, errMsg string) Event {
ev := Event{Type: EventToolFinished, RunID: m.runID, ChatID: m.chatID, Model: m.model, ToolStatus: status, ToolCallID: callID, ToolName: toolName, WorkingDir: workingDir, Args: args, Content: content}
if errMsg != "" {
ev.Error = errMsg
}
return ev
}
func newRunFinished(m eventMetadata, status RunStatus) Event {
return Event{Type: EventRunFinished, RunID: m.runID, ChatID: m.chatID, Model: m.model, Status: status}
}
func newErrorEvent(m eventMetadata, errMsg string) Event {
return Event{Type: EventError, RunID: m.runID, ChatID: m.chatID, Model: m.model, Error: errMsg}
}
func newCompactionProgress(m eventMetadata, tokens int) Event {
return Event{Type: EventCompactionProgress, RunID: m.runID, ChatID: m.chatID, Model: m.model, Tokens: tokens}
}
func newCompactionStarted(m eventMetadata, trigger CompactionTrigger) Event {
return Event{Type: EventCompactionStarted, RunID: m.runID, ChatID: m.chatID, Model: m.model, CompactionTrigger: trigger}
}
func newCompactionSkipped(m eventMetadata, trigger CompactionTrigger, content string) Event {
return Event{Type: EventCompactionSkipped, RunID: m.runID, ChatID: m.chatID, Model: m.model, CompactionTrigger: trigger, Content: content}
}
func newCompacted(m eventMetadata, messages []api.Message, trigger CompactionTrigger, content string) Event {
return Event{Type: EventCompacted, RunID: m.runID, ChatID: m.chatID, Model: m.model, CompactionTrigger: trigger, Content: content, Messages: messages}
}
func (s *Session) emit(event Event) error {
if s == nil {
return nil
}
var errs []error
for _, sink := range s.EventSinks {
if sink == nil {
continue
}
if err := sink.Emit(event); err != nil {
errs = append(errs, err)
}
}
return errors.Join(errs...)
}
func (s *Session) emitIgnoringCanceled(ctx context.Context, event Event) error {
err := s.emit(event)
if err != nil && ctx != nil && ctx.Err() != nil {
//nolint:nilerr // Event sinks may close during cancellation; cancellation is not a user-facing emit failure.
return nil
}
return err
}
+104
View File
@@ -0,0 +1,104 @@
package agent
import (
"context"
"fmt"
"sort"
"github.com/ollama/ollama/api"
)
type ToolContext struct {
WorkingDir string
}
type ToolResult struct {
Content string
WorkingDir string
}
type Tool interface {
Name() string
Description() string
Schema() api.ToolFunction
Execute(context.Context, ToolContext, map[string]any) (ToolResult, error)
}
type ApprovalRequired interface {
RequiresApproval(map[string]any) bool
}
// ScopedTool is implemented by tools that need per-invocation approval
// scoping beyond the tool name (e.g. shell commands scoped to the exact
// command string). Tools that don't implement this are scoped by name only.
type ScopedTool interface {
ApprovalScope(args map[string]any) string
}
type Registry struct {
tools map[string]Tool
}
func (r *Registry) Register(tool Tool) {
if r == nil || tool == nil {
return
}
if r.tools == nil {
r.tools = make(map[string]Tool)
}
r.tools[tool.Name()] = tool
}
func (r *Registry) Get(name string) (Tool, bool) {
if r == nil {
return nil, false
}
tool, ok := r.tools[name]
return tool, ok
}
func (r *Registry) Names() []string {
if r == nil {
return nil
}
names := make([]string, 0, len(r.tools))
for name := range r.tools {
names = append(names, name)
}
sort.Strings(names)
return names
}
func (r *Registry) Tools() api.Tools {
if r == nil {
return nil
}
names := r.Names()
apiTools := make(api.Tools, 0, len(names))
for _, name := range names {
tool := r.tools[name]
apiTools = append(apiTools, api.Tool{
Type: "function",
Function: tool.Schema(),
})
}
return apiTools
}
func (r *Registry) Execute(ctx context.Context, toolCtx ToolContext, call api.ToolCall) (ToolResult, error) {
tool, ok := r.Get(call.Function.Name)
if !ok {
return ToolResult{}, fmt.Errorf("unknown tool: %s", call.Function.Name)
}
return tool.Execute(ctx, toolCtx, call.Function.Arguments.ToMap())
}
func ToolRequiresApproval(tool Tool, args map[string]any) bool {
if tool == nil {
return false
}
if t, ok := tool.(ApprovalRequired); ok {
return t.RequiresApproval(args)
}
return false
}
+1099
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
@@ -0,0 +1,57 @@
package agent
import (
"context"
"strings"
"github.com/google/uuid"
"github.com/ollama/ollama/api"
)
// activateSkill loads opts.SkillName from the catalog and injects a synthetic
// assistant tool call plus tool result before the first model request, so the
// transcript looks like a real skill tool invocation. It emits the same
// tool_call_detected -> tool_started -> tool_finished lifecycle the model path
// uses, and returns the messages to prepend. A blank SkillName is a no-op.
func (s *Session) activateSkill(ctx context.Context, runID string, opts RunOptions) ([]api.Message, error) {
name := strings.TrimSpace(opts.SkillName)
if name == "" {
return nil, nil
}
select {
case <-ctx.Done():
return nil, ctx.Err()
default:
}
skill, err := s.Skills.Load(name)
if err != nil {
return nil, err
}
args := api.NewToolCallFunctionArguments()
args.Set("name", skill.Name)
call := api.ToolCall{
ID: "call_skill_" + uuid.NewString(),
Function: api.ToolCallFunction{Name: "skill", Arguments: args},
}
result := api.Message{
Role: "tool",
ToolName: "skill",
ToolCallID: call.ID,
Content: skill.Content(),
}
meta := newEventMetadata(runID, opts)
if err := s.emit(newToolCallDetected(meta, []api.ToolCall{call})); err != nil {
return nil, err
}
if err := s.emit(newToolStarted(meta, call.ID, "skill", s.currentWorkingDir(), args.ToMap())); err != nil {
return nil, err
}
if err := s.emitIgnoringCanceled(ctx, newToolFinished(meta, ToolStatusDone, call.ID, "skill", s.currentWorkingDir(), args.ToMap(), result.Content, "")); err != nil {
return nil, err
}
return []api.Message{
{Role: "assistant", ToolCalls: []api.ToolCall{call}},
result,
}, nil
}
+74
View File
@@ -0,0 +1,74 @@
package agent
import (
"context"
"os"
"path/filepath"
"strings"
"testing"
"github.com/ollama/ollama/api"
)
type skillTestClient struct{ requests []*api.ChatRequest }
func (c *skillTestClient) Chat(_ context.Context, req *api.ChatRequest, fn api.ChatResponseFunc) error {
c.requests = append(c.requests, req)
return fn(api.ChatResponse{Message: api.Message{Role: "assistant", Content: "Done."}})
}
func testSkillCatalog(t *testing.T) *SkillCatalog {
t.Helper()
dir := t.TempDir()
path := filepath.Join(dir, "release-notes")
if err := os.Mkdir(path, 0o755); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(filepath.Join(path, "SKILL.md"), []byte("---\nname: release-notes\ndescription: Draft release notes.\n---\nUse concise bullets."), 0o644); err != nil {
t.Fatal(err)
}
catalog, err := DiscoverSkills(dir)
if err != nil {
t.Fatal(err)
}
return catalog
}
func TestSessionSkillActivationPreservesCallAndResultOrder(t *testing.T) {
catalog := testSkillCatalog(t)
client := &skillTestClient{}
events := &recordingEventSink{}
result, err := (&Session{Client: client, Skills: catalog, EventSinks: []EventSink{events}}).Run(context.Background(), RunOptions{
Model: "test",
NewMessages: []api.Message{{Role: "user", Content: "draft release notes"}},
SkillName: "release-notes",
})
if err != nil {
t.Fatal(err)
}
if len(result.Messages) != 4 {
t.Fatalf("transcript = %#v", result.Messages)
}
call, toolTranscript := result.Messages[1], result.Messages[2]
if call.Role != "assistant" || len(call.ToolCalls) != 1 || call.ToolCalls[0].Function.Name != "skill" || !strings.HasPrefix(call.ToolCalls[0].ID, "call_skill_") {
t.Fatalf("call message = %#v", call)
}
if toolTranscript.Role != "tool" || toolTranscript.ToolName != "skill" || toolTranscript.ToolCallID != call.ToolCalls[0].ID || !strings.Contains(toolTranscript.Content, "Use concise bullets.") {
t.Fatalf("tool result = %#v", toolTranscript)
}
if len(client.requests) != 1 || len(client.requests[0].Messages) != 3 || client.requests[0].Messages[2].ToolCallID != call.ToolCalls[0].ID {
t.Fatalf("model request did not preserve transcript: %#v", client.requests)
}
var skillEvents []EventType
for _, event := range events.events {
if event.ToolName == "skill" || event.Type == EventToolCallDetected {
skillEvents = append(skillEvents, event.Type)
}
}
if len(skillEvents) < 3 {
t.Fatalf("skill event order = %#v, want tool_call_detected,tool_started,tool_finished", skillEvents)
}
if got, want := strings.Join([]string{string(skillEvents[0]), string(skillEvents[1]), string(skillEvents[2])}, ","), "tool_call_detected,tool_started,tool_finished"; got != want {
t.Fatalf("skill event order = %#v, want %s", skillEvents, want)
}
}
+438
View File
@@ -0,0 +1,438 @@
package agent
import (
"errors"
"fmt"
"io/fs"
"os"
"path/filepath"
"regexp"
"sort"
"strings"
"gopkg.in/yaml.v3"
)
const (
// SkillsDirEnv overrides the user-level Ollama-owned skills directory. The
// cross-client .agents/skills/ convention and project-level .ollama/skills/
// are also scanned (see LoadDefaultSkills); on a name collision, Ollama-owned
// directories take precedence over .agents/skills/, and project-level takes
// precedence over user-level.
SkillsDirEnv = "OLLAMA_SKILLS"
skillFilename = "SKILL.md"
maxSkillBytes = 1 << 20
bundledSkillCreatorName = "skill-creator"
bundledSkillCreatorContent = `---
name: skill-creator
description: Create or improve reusable skills. Use when the user wants a reusable skill, asks how to author SKILL.md, or needs help installing a skill.
---
# Create a skill
Create a focused, reusable instruction package. Treat a skill as guidance for the model, not as a way to gain new permissions or bypass safety controls.
## Choose the location
Create user skills beside this one. The skill directory shown in the loaded skill context is this skill's location; its parent is the user skill root. This bundled skill normally lives at ~/.ollama/skills/skill-creator, so new user skills normally go at ~/.ollama/skills/<skill-name>/SKILL.md.
Use a project-local skill directory only when the user asks to keep the skill with that project. Do not overwrite an existing skill without the user's approval. New and changed skills are discovered when the agent starts, so tell the user to begin a new agent session afterward.
## Follow the required shape
Use the directory name as the skill name. Use lowercase letters, numbers, and single hyphens only. Keep the name short and no longer than 64 characters.
Every skill needs a SKILL.md with YAML frontmatter followed by Markdown instructions:
~~~md
---
name: release-notes
description: Draft concise release notes from completed changes. Use when the user asks for a changelog, release notes, or GitHub release copy.
---
# Draft release notes
Write the workflow here.
~~~
Require a non-empty description that says both what the skill does and when to use it. Keep the body procedural and concise. Put detailed schemas, long examples, and variant-specific guidance in references/ only when the skill needs them.
Use scripts/ for repeatable or fragile operations that benefit from deterministic execution. Use assets/ for files that belong in generated output. Do not add README files, changelogs, or setup notes that do not help the model perform the task.
## Create safely
1. Identify the repeated task, expected inputs, and useful output.
2. Choose the smallest name and description that reliably trigger the skill.
3. Create the folder and SKILL.md; add resources only when they remove real repeated work.
4. Re-read the completed file and verify its frontmatter, directory-name match, and relative resource paths.
5. Tell the user where it was created and that a new agent session will discover it.
Skills provide instructions only. They do not grant filesystem, network, shell, or approval privileges, and they do not make a tool available. Use only the tools that are actually available, follow their normal approval rules, and ask before actions that need user authorization.
`
)
var skillName = regexp.MustCompile(`^[a-z0-9]+(?:-[a-z0-9]+)*$`)
// SkillsDir returns the canonical runtime-owned skill directory.
func SkillsDir() (string, error) {
if path := strings.TrimSpace(os.Getenv(SkillsDirEnv)); path != "" {
return filepath.Abs(path)
}
if xdg := strings.TrimSpace(os.Getenv("XDG_CONFIG_HOME")); xdg != "" {
return filepath.Join(xdg, "ollama", "skills"), nil
}
home, err := os.UserHomeDir()
if err != nil {
return "", err
}
return filepath.Join(home, ".ollama", "skills"), nil
}
// Skill is a validated, loadable instruction set. It never grants tool
// permissions; it is supplied to the model as ordinary tool-result content.
type Skill struct {
Name string
Description string
Instructions string
Path string
}
func (s Skill) Content() string {
var b strings.Builder
fmt.Fprintf(&b, "<skill name=%q>\n%s\n", s.Name, strings.TrimSpace(s.Instructions))
if s.Path != "" {
dir := filepath.Dir(s.Path)
fmt.Fprintf(&b, "Skill directory: %s\n", dir)
b.WriteString("Relative paths in this skill are relative to the skill directory.\n")
}
if resources := s.resources(); len(resources) > 0 {
b.WriteString("<skill_resources>\n")
for _, r := range resources {
fmt.Fprintf(&b, " <file>%s</file>\n", r)
}
b.WriteString("</skill_resources>\n")
}
b.WriteString("</skill>")
return b.String()
}
// resources lists bundled files one level deep under scripts/, references/,
// and assets/ without reading them, so the model can load them on demand.
func (s Skill) resources() []string {
if s.Path == "" {
return nil
}
dir := filepath.Dir(s.Path)
var resources []string
for _, sub := range []string{"scripts", "references", "assets"} {
entries, err := os.ReadDir(filepath.Join(dir, sub))
if err != nil {
continue
}
for _, e := range entries {
if e.IsDir() {
continue
}
resources = append(resources, sub+"/"+e.Name())
}
}
sort.Strings(resources)
return resources
}
// SkillCatalog contains valid skills and diagnostics for ignored invalid
// entries, so one malformed skill cannot hide the rest.
type SkillCatalog struct {
dir string
skills map[string]Skill
diagnostics []error
}
func DiscoverSkills(dir string) (*SkillCatalog, error) {
dir, err := filepath.Abs(strings.TrimSpace(dir))
if err != nil {
return nil, err
}
catalog := &SkillCatalog{dir: dir, skills: make(map[string]Skill)}
entries, err := os.ReadDir(dir)
if errors.Is(err, fs.ErrNotExist) {
return catalog, nil
}
if err != nil {
return nil, fmt.Errorf("read skills directory: %w", err)
}
for _, entry := range entries {
name := entry.Name()
// Follow symlinks so users can point at shared skill repositories.
// The link name (not the target) is the canonical skill name.
info, err := os.Stat(filepath.Join(dir, name))
if err != nil {
if errors.Is(err, fs.ErrNotExist) {
continue
}
catalog.diagnostics = append(catalog.diagnostics, fmt.Errorf("skill %q: %w", name, err))
continue
}
if !info.IsDir() {
continue
}
if !skillName.MatchString(name) {
catalog.diagnostics = append(catalog.diagnostics, fmt.Errorf("invalid skill directory %q", name))
continue
}
skill, err := parseSkill(filepath.Join(dir, name, skillFilename), name)
if errors.Is(err, fs.ErrNotExist) {
continue
}
if err != nil {
catalog.diagnostics = append(catalog.diagnostics, err)
continue
}
catalog.skills[skill.Name] = skill
}
return catalog, nil
}
// LoadDefaultSkills discovers skills from the spec's scopes, merged with
// deterministic precedence. Roots are scanned lowest-precedence first so later
// roots override earlier ones on name collisions (recording a diagnostic):
//
// 1. ~/.agents/skills/ (user, cross-client)
// 2. user Ollama skills dir (user, Ollama-owned; SkillsDir)
// 3. <project>/.agents/skills/ (project, cross-client)
// 4. <project>/.ollama/skills/ (project, Ollama-owned)
//
// Project-level overrides user-level, and within a scope Ollama-owned
// directories override .agents/skills/. projectDir is the agent's working
// directory at startup (discovery is a session-start snapshot per the spec).
func LoadDefaultSkills(projectDir string) (*SkillCatalog, error) {
roots, err := defaultSkillRoots(projectDir)
if err != nil {
return nil, err
}
catalog := &SkillCatalog{skills: make(map[string]Skill)}
bundled, err := bundledSkillCreator()
if err != nil {
return nil, err
}
catalog.skills[bundled.Name] = bundled
if err := installBundledSkillCreator(); err != nil {
catalog.diagnostics = append(catalog.diagnostics, err)
}
for _, root := range roots {
sub, err := DiscoverSkills(root.path)
if err != nil {
catalog.diagnostics = append(catalog.diagnostics, fmt.Errorf("discover skills in %s: %w", root.path, err))
continue
}
catalog.diagnostics = append(catalog.diagnostics, sub.diagnostics...)
for _, skill := range sub.skills {
// Name collisions across roots are expected precedence resolution,
// not errors: later (higher-precedence) roots legitimately override
// earlier ones. The skill is still loaded; no diagnostic needed.
catalog.skills[skill.Name] = skill
}
}
return catalog, nil
}
func bundledSkillCreator() (Skill, error) {
skill, err := parseSkillContent("", bundledSkillCreatorName, bundledSkillCreatorContent)
if err != nil {
return Skill{}, fmt.Errorf("load bundled %s skill: %w", bundledSkillCreatorName, err)
}
return skill, nil
}
func installBundledSkillCreator() error {
dir, err := SkillsDir()
if err != nil {
return fmt.Errorf("resolve bundled skill directory: %w", err)
}
path := filepath.Join(dir, bundledSkillCreatorName, skillFilename)
if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil {
return fmt.Errorf("create bundled skill directory: %w", err)
}
contents, err := os.ReadFile(path)
if err == nil && string(contents) == bundledSkillCreatorContent {
return nil
}
if err != nil && !errors.Is(err, fs.ErrNotExist) {
return fmt.Errorf("read bundled skill: %w", err)
}
if err := os.WriteFile(path, []byte(bundledSkillCreatorContent), 0o644); err != nil {
return fmt.Errorf("write bundled skill: %w", err)
}
return nil
}
type skillRoot struct {
path string
}
// defaultSkillRoots returns skill directories ordered lowest- to
// highest-precedence. Non-existent directories are scanned harmlessly
// (DiscoverSkills skips them).
func defaultSkillRoots(projectDir string) ([]skillRoot, error) {
var roots []skillRoot
if home, err := os.UserHomeDir(); err == nil && home != "" {
roots = append(roots, skillRoot{path: filepath.Join(home, ".agents", "skills")})
}
userOllama, err := SkillsDir()
if err != nil {
return nil, err
}
roots = append(roots, skillRoot{path: userOllama})
projectDir = strings.TrimSpace(projectDir)
if projectDir != "" {
if abs, err := filepath.Abs(projectDir); err == nil {
roots = append(roots,
skillRoot{path: filepath.Join(abs, ".agents", "skills")},
skillRoot{path: filepath.Join(abs, ".ollama", "skills")},
)
}
}
return roots, nil
}
func (c *SkillCatalog) Dir() string {
if c == nil {
return ""
}
return c.dir
}
func (c *SkillCatalog) List() []Skill {
if c == nil {
return nil
}
list := make([]Skill, 0, len(c.skills))
for _, skill := range c.skills {
list = append(list, skill)
}
sort.Slice(list, func(i, j int) bool { return list[i].Name < list[j].Name })
return list
}
func (c *SkillCatalog) Diagnostics() []error {
if c == nil {
return nil
}
return append([]error(nil), c.diagnostics...)
}
func (c *SkillCatalog) Load(name string) (Skill, error) {
name = strings.TrimSpace(name)
if !skillName.MatchString(name) {
return Skill{}, fmt.Errorf("invalid skill name %q", name)
}
if c == nil {
return Skill{}, errors.New("skills are unavailable")
}
skill, ok := c.skills[name]
if !ok {
return Skill{}, fmt.Errorf("skill %q not found in %s", name, c.dir)
}
return skill, nil
}
// SystemContext advertises the catalog without expanding full instructions in
// every request. The skill call is the explicit loading boundary.
func (c *SkillCatalog) SystemContext() string {
list := c.List()
if len(list) == 0 {
return ""
}
lines := []string{"<available_skills>"}
for _, skill := range list {
description := skill.Description
if description == "" {
description = "No description provided."
}
lines = append(lines, fmt.Sprintf("- %s: %s", skill.Name, description))
}
lines = append(lines, "</available_skills>", "Load a matching skill with the skill tool before following its instructions. Skills only provide instructions; use ordinary tools for filesystem or network access, with their normal approval rules.")
return strings.Join(lines, "\n")
}
func parseSkill(path, directoryName string) (Skill, error) {
// Stat (not Lstat) so a symlinked SKILL.md resolves to its target file.
info, err := os.Stat(path)
if err != nil {
return Skill{}, err
}
if !info.Mode().IsRegular() {
return Skill{}, fmt.Errorf("skill %q: %s is not a regular file", directoryName, skillFilename)
}
if info.Size() > maxSkillBytes {
return Skill{}, fmt.Errorf("skill %q: %s exceeds %d bytes", directoryName, skillFilename, maxSkillBytes)
}
data, err := os.ReadFile(path)
if err != nil {
return Skill{}, fmt.Errorf("read skill %q: %w", directoryName, err)
}
return parseSkillContent(path, directoryName, string(data))
}
func parseSkillContent(path, directoryName, input string) (Skill, error) {
instructions := strings.TrimSpace(input)
if instructions == "" {
return Skill{}, fmt.Errorf("skill %q: %s is empty", directoryName, skillFilename)
}
if !strings.HasPrefix(instructions, "---\n") && !strings.HasPrefix(instructions, "---\r\n") {
return Skill{}, fmt.Errorf("skill %q: missing YAML front matter", directoryName)
}
metadata, body, err := skillFrontMatter(instructions)
if err != nil {
return Skill{}, fmt.Errorf("skill %q: %w", directoryName, err)
}
if metadata.Name == "" {
return Skill{}, fmt.Errorf("skill %q: front matter requires name", directoryName)
}
if metadata.Description == "" {
return Skill{}, fmt.Errorf("skill %q: front matter requires description", directoryName)
}
if !skillName.MatchString(metadata.Name) {
return Skill{}, fmt.Errorf("skill %q: invalid front matter name %q", directoryName, metadata.Name)
}
if metadata.Name != directoryName {
return Skill{}, fmt.Errorf("skill %q: front matter name %q must match directory name", directoryName, metadata.Name)
}
skill := Skill{Name: metadata.Name, Description: metadata.Description, Path: path}
instructions = body
if strings.TrimSpace(instructions) == "" {
return Skill{}, fmt.Errorf("skill %q: instructions are empty", directoryName)
}
skill.Instructions = strings.TrimSpace(instructions)
return skill, nil
}
type skillFrontMatterMetadata struct {
Name string `yaml:"name"`
Description string `yaml:"description"`
Metadata map[string]any `yaml:"metadata"`
}
func skillFrontMatter(input string) (skillFrontMatterMetadata, string, error) {
input = strings.ReplaceAll(input, "\r\n", "\n")
lines := strings.Split(input, "\n")
if len(lines) < 3 || lines[0] != "---" {
return skillFrontMatterMetadata{}, "", errors.New("invalid front matter")
}
for i := 1; i < len(lines); i++ {
if lines[i] == "---" {
var metadata skillFrontMatterMetadata
if err := yaml.Unmarshal([]byte(strings.Join(lines[1:i], "\n")), &metadata); err != nil {
return skillFrontMatterMetadata{}, "", fmt.Errorf("parse YAML front matter: %w", err)
}
metadata.Name = strings.TrimSpace(metadata.Name)
metadata.Description = strings.TrimSpace(metadata.Description)
return metadata, strings.Join(lines[i+1:], "\n"), nil
}
}
return skillFrontMatterMetadata{}, "", errors.New("front matter is not closed")
}
+293
View File
@@ -0,0 +1,293 @@
package agent
import (
"os"
"path/filepath"
"strings"
"testing"
)
func writeCatalogSkill(t *testing.T, dir, name, content string) {
t.Helper()
path := filepath.Join(dir, name)
if err := os.MkdirAll(path, 0o755); err != nil {
t.Fatal(err)
}
if !strings.HasPrefix(content, "---") {
content = "---\nname: " + name + "\ndescription: Test skill.\n---\n" + content
}
if err := os.WriteFile(filepath.Join(path, skillFilename), []byte(content), 0o644); err != nil {
t.Fatal(err)
}
}
func TestDiscoverAndLoadSkills(t *testing.T) {
dir := t.TempDir()
writeCatalogSkill(t, dir, "release-notes", "---\nname: release-notes\ndescription: Draft concise release notes.\nmetadata:\n author: Ollama\n labels:\n - release\n - docs\n---\n# Release notes\n\nUse short bullets.")
catalog, err := DiscoverSkills(dir)
if err != nil {
t.Fatal(err)
}
list := catalog.List()
if len(list) != 1 || list[0].Name != "release-notes" || list[0].Description != "Draft concise release notes." {
t.Fatalf("skills = %#v", list)
}
skill, err := catalog.Load("release-notes")
if err != nil {
t.Fatal(err)
}
if !strings.Contains(skill.Content(), `<skill name="release-notes">`) || !strings.Contains(skill.Content(), "Use short bullets.") {
t.Fatalf("skill content = %q", skill.Content())
}
if context := catalog.SystemContext(); !strings.Contains(context, "release-notes: Draft concise release notes.") || !strings.Contains(context, "normal approval rules") {
t.Fatalf("system context = %q", context)
}
}
func TestDiscoverSkillsSkipsMalformedEntries(t *testing.T) {
dir := t.TempDir()
writeCatalogSkill(t, dir, "valid", "do the useful thing")
writeCatalogSkill(t, dir, "mismatched", "---\nname: whatever\ndescription: wrong name\n---\nbody")
// Genuinely malformed front matter (a line without a key:value pair) is still rejected.
writeCatalogSkill(t, dir, "broken", "---\nname: broken\ndescription\n---\nnope")
writeCatalogSkill(t, dir, "missing-name", "---\ndescription: missing name\n---\nbody")
writeCatalogSkill(t, dir, "missing-description", "---\nname: missing-description\n---\nbody")
writeCatalogSkill(t, dir, "bad-name", "---\nname: bad_name\ndescription: invalid name\n---\nbody")
writeCatalogSkill(t, dir, "under_score", "---\nname: under_score\ndescription: invalid directory\n---\nbody")
if err := os.MkdirAll(filepath.Join(dir, "no-front-matter"), 0o755); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(filepath.Join(dir, "no-front-matter", skillFilename), []byte("body"), 0o644); err != nil {
t.Fatal(err)
}
catalog, err := DiscoverSkills(dir)
if err != nil {
t.Fatal(err)
}
if got, want := len(catalog.List()), 1; got != want {
t.Fatalf("valid skills = %d, want %d", got, want)
}
if got, want := len(catalog.Diagnostics()), 7; got != want {
t.Fatalf("diagnostics = %d, want %d: %#v", got, want, catalog.Diagnostics())
}
if _, err := catalog.Load("broken"); err == nil || !strings.Contains(err.Error(), "not found") {
t.Fatalf("load broken error = %v", err)
}
if _, err := catalog.Load("../valid"); err == nil || !strings.Contains(err.Error(), "invalid skill name") {
t.Fatalf("unsafe name error = %v", err)
}
}
func TestDiscoverSkillsFollowsSymlinks(t *testing.T) {
dir := t.TempDir()
target := t.TempDir()
writeCatalogSkill(t, target, "shared", "---\nname: shared\ndescription: From a linked repo.\n---\nshared instructions")
if err := os.Symlink(filepath.Join(target, "shared"), filepath.Join(dir, "shared")); err != nil {
t.Skipf("symlink not supported: %v", err)
}
catalog, err := DiscoverSkills(dir)
if err != nil {
t.Fatal(err)
}
list := catalog.List()
if len(list) != 1 || list[0].Name != "shared" || list[0].Description != "From a linked repo." {
t.Fatalf("symlinked skills = %#v", list)
}
if !strings.Contains(list[0].Content(), "shared instructions") {
t.Fatalf("symlinked skill content = %q", list[0].Content())
}
}
func TestLoadDefaultSkillsContinuesAfterBadRoot(t *testing.T) {
home := t.TempDir()
t.Setenv("HOME", home)
t.Setenv("USERPROFILE", home)
project := t.TempDir()
writeCatalogSkill(t, filepath.Join(project, ".ollama", "skills"), "release-notes", "project instructions")
badRoot := filepath.Join(t.TempDir(), "not-a-directory")
if err := os.WriteFile(badRoot, []byte("not a directory"), 0o644); err != nil {
t.Fatal(err)
}
t.Setenv(SkillsDirEnv, badRoot)
catalog, err := LoadDefaultSkills(project)
if err != nil {
t.Fatal(err)
}
if _, err := catalog.Load("release-notes"); err != nil {
t.Fatalf("valid skill was hidden by bad root: %v", err)
}
if _, err := catalog.Load(bundledSkillCreatorName); err != nil {
t.Fatalf("bundled skill was hidden by bad root: %v", err)
}
var foundDiagnostic bool
for _, diagnostic := range catalog.Diagnostics() {
if strings.Contains(diagnostic.Error(), badRoot) {
foundDiagnostic = true
break
}
}
if !foundDiagnostic {
t.Fatalf("diagnostics = %#v, want bad root %q", catalog.Diagnostics(), badRoot)
}
}
func TestLoadDefaultSkillsInstallsBundledSkillCreator(t *testing.T) {
dir := t.TempDir()
t.Setenv(SkillsDirEnv, dir)
catalog, err := LoadDefaultSkills("")
if err != nil {
t.Fatal(err)
}
skill, err := catalog.Load(bundledSkillCreatorName)
if err != nil {
t.Fatal(err)
}
path := filepath.Join(dir, bundledSkillCreatorName, skillFilename)
contents, err := os.ReadFile(path)
if err != nil {
t.Fatal(err)
}
if string(contents) != bundledSkillCreatorContent {
t.Fatalf("installed skill = %q, want bundled contents", contents)
}
if skill.Path != path {
t.Fatalf("skill path = %q, want %q", skill.Path, path)
}
if !strings.Contains(skill.Content(), "Skill directory: "+filepath.Dir(path)) {
t.Fatalf("skill content does not identify its directory: %q", skill.Content())
}
}
func TestLoadDefaultSkillsUpdatesExistingSkillCreator(t *testing.T) {
dir := t.TempDir()
t.Setenv(SkillsDirEnv, dir)
writeCatalogSkill(t, dir, bundledSkillCreatorName, "custom instructions")
if _, err := LoadDefaultSkills(""); err != nil {
t.Fatal(err)
}
contents, err := os.ReadFile(filepath.Join(dir, bundledSkillCreatorName, skillFilename))
if err != nil {
t.Fatal(err)
}
if string(contents) != bundledSkillCreatorContent {
t.Fatalf("installed skill = %q, want bundled contents", contents)
}
}
func TestSkillsDirUsesOverrideAndXDG(t *testing.T) {
base := t.TempDir()
override := filepath.Join(base, "skills-override")
t.Setenv(SkillsDirEnv, override)
got, err := SkillsDir()
if err != nil {
t.Fatal(err)
}
want, err := filepath.Abs(override)
if err != nil {
t.Fatal(err)
}
if got != want {
t.Fatalf("SkillsDir override = %q, want %q", got, want)
}
t.Setenv(SkillsDirEnv, "")
xdg := filepath.Join(base, "xdg")
t.Setenv("XDG_CONFIG_HOME", xdg)
if got, err := SkillsDir(); err != nil || got != filepath.Join(xdg, "ollama", "skills") {
t.Fatalf("SkillsDir xdg = %q, want %q, %v", got, filepath.Join(xdg, "ollama", "skills"), err)
}
t.Setenv("XDG_CONFIG_HOME", "")
home := filepath.Join(base, "home")
t.Setenv("HOME", home)
t.Setenv("USERPROFILE", home)
if got, err := SkillsDir(); err != nil || got != filepath.Join(home, ".ollama", "skills") {
t.Fatalf("SkillsDir default = %q, want %q, %v", got, filepath.Join(home, ".ollama", "skills"), err)
}
}
func TestLoadDefaultSkillsPrecedenceAndCollisions(t *testing.T) {
home := t.TempDir()
t.Setenv("HOME", home)
t.Setenv("USERPROFILE", home) // Windows: os.UserHomeDir uses %USERPROFILE%
userOllama := t.TempDir()
t.Setenv(SkillsDirEnv, userOllama)
userAgents := filepath.Join(home, ".agents", "skills")
project := t.TempDir()
projectAgents := filepath.Join(project, ".agents", "skills")
projectOllama := filepath.Join(project, ".ollama", "skills")
// release-notes exists in all four roots; project ollama must win.
writeCatalogSkill(t, userAgents, "release-notes", "from user agents")
writeCatalogSkill(t, userOllama, "release-notes", "from user ollama")
writeCatalogSkill(t, projectOllama, "release-notes", "from project ollama")
// code-review exists in both project roots; project ollama beats project agents.
writeCatalogSkill(t, projectAgents, "code-review", "from project agents")
writeCatalogSkill(t, projectOllama, "code-review", "from project ollama")
// unique appears only in user ollama (via env override).
writeCatalogSkill(t, userOllama, "unique", "only here")
catalog, err := LoadDefaultSkills(project)
if err != nil {
t.Fatal(err)
}
rn, err := catalog.Load("release-notes")
if err != nil || !strings.Contains(rn.Instructions, "from project ollama") || !strings.Contains(rn.Path, ".ollama") {
t.Fatalf("release-notes = %#v, want project ollama to win", rn)
}
cr, err := catalog.Load("code-review")
if err != nil || !strings.Contains(cr.Instructions, "from project ollama") {
t.Fatalf("code-review = %#v, want project ollama to win over project agents", cr)
}
if _, err := catalog.Load("unique"); err != nil {
t.Fatalf("unique should load from user ollama: %v", err)
}
// Collisions are resolved silently by precedence — no diagnostics.
for _, d := range catalog.Diagnostics() {
if strings.Contains(d.Error(), "shadows") {
t.Fatalf("unexpected shadow diagnostic: %v", d)
}
}
}
func TestSkillContentListsDirectoryAndResources(t *testing.T) {
root := t.TempDir()
skillDir := filepath.Join(root, "pdf-processing")
if err := os.MkdirAll(filepath.Join(skillDir, "scripts"), 0o755); err != nil {
t.Fatal(err)
}
if err := os.MkdirAll(filepath.Join(skillDir, "references"), 0o755); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(filepath.Join(skillDir, "SKILL.md"), []byte("---\nname: pdf-processing\ndescription: Handle PDFs.\n---\nHandle PDFs."), 0o644); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(filepath.Join(skillDir, "scripts", "extract.py"), []byte("#!/usr/bin/env python3"), 0o755); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(filepath.Join(skillDir, "references", "ref.md"), []byte("ref"), 0o644); err != nil {
t.Fatal(err)
}
catalog, err := DiscoverSkills(root)
if err != nil {
t.Fatal(err)
}
skill, err := catalog.Load("pdf-processing")
if err != nil {
t.Fatal(err)
}
content := skill.Content()
if !strings.Contains(content, "Skill directory:") || !strings.Contains(content, skillDir) {
t.Fatalf("content missing skill directory: %q", content)
}
if !strings.Contains(content, "<file>scripts/extract.py</file>") || !strings.Contains(content, "<file>references/ref.md</file>") {
t.Fatalf("content missing resource listing: %q", content)
}
}
+450
View File
@@ -0,0 +1,450 @@
package tools
import (
"context"
"errors"
"fmt"
"os"
"os/exec"
"path/filepath"
"runtime"
"strings"
"time"
"unicode/utf8"
"github.com/ollama/ollama/agent"
"github.com/ollama/ollama/api"
)
const (
bashTimeout = 3 * time.Minute
bashWaitDelay = 1 * time.Second
maxBashOutputBytes = 60_000
)
type Bash struct{}
func (b *Bash) Name() string {
return shellToolName()
}
func (b *Bash) Description() string {
return shellToolDescription()
}
func (b *Bash) Schema() api.ToolFunction {
props := api.NewToolPropertiesMap()
props.Set("command", api.ToolProperty{
Type: api.PropertyType{"string"},
Description: shellCommandDescription(),
})
return api.ToolFunction{
Name: b.Name(),
Description: b.Description(),
Parameters: api.ToolFunctionParameters{
Type: "object",
Properties: props,
Required: []string{"command"},
},
}
}
func (b *Bash) RequiresApproval(map[string]any) bool {
return true
}
// ApprovalScope scopes shell approval to the exact, trimmed command string
// using a NUL separator: "<tool>\x00<command>". "Always allow this command"
// matches ONLY that precise string — any whitespace, quoting, or casing
// variant re-prompts. The NUL separator is safe because a shell command
// string cannot contain a literal NUL.
func (b *Bash) ApprovalScope(args map[string]any) string {
name := b.Name()
if command, ok := args["command"].(string); ok {
command = strings.TrimSpace(command)
if command != "" {
return name + "\x00" + command
}
}
return name
}
func (b *Bash) Execute(ctx context.Context, toolCtx agent.ToolContext, args map[string]any) (agent.ToolResult, error) {
// TODO: use shared agent.RequiredStringArg for the "command" parameter (see agent package cleanup plan).
command, ok := args["command"].(string)
if !ok || strings.TrimSpace(command) == "" {
return agent.ToolResult{}, fmt.Errorf("command parameter is required")
}
if err := rejectUnsafeShellCommand(command); err != nil {
return agent.ToolResult{}, err
}
ctx, cancel := context.WithTimeout(ctx, bashTimeout)
defer cancel()
cwdFile, err := os.CreateTemp("", "ollama-agent-cwd-*")
if err != nil {
return agent.ToolResult{}, err
}
cwdPath := cwdFile.Name()
_ = cwdFile.Close()
defer os.Remove(cwdPath)
cmd := newBashCommand(ctx, command, cwdPath)
cmd.WaitDelay = bashWaitDelay
cmd.Cancel = func() error {
return killBashCommand(cmd)
}
if toolCtx.WorkingDir != "" {
cmd.Dir = toolCtx.WorkingDir
}
var stdout, stderr boundedOutput
stdout.Limit = maxBashOutputBytes
stderr.Limit = maxBashOutputBytes
cmd.Stdout = &stdout
cmd.Stderr = &stderr
err = runBashCommand(cmd)
finalWorkingDir := readFinalWorkingDir(cwdPath)
var sb strings.Builder
if stdout.Len() > 0 {
sb.WriteString(stdout.String("stdout"))
}
if stderr.Len() > 0 {
if sb.Len() > 0 {
sb.WriteString("\n")
}
sb.WriteString("stderr:\n")
sb.WriteString(stderr.String("stderr"))
}
if err != nil {
if ctx.Err() == context.DeadlineExceeded {
return agent.ToolResult{Content: bashContentWithError(sb.String(), "Error: command timed out after "+bashTimeout.String()), WorkingDir: finalWorkingDir}, nil
}
if ctx.Err() == context.Canceled {
return agent.ToolResult{Content: bashContentWithError(sb.String(), "Error: command was canceled"), WorkingDir: finalWorkingDir}, nil
}
if errors.Is(err, exec.ErrWaitDelay) {
_ = killBashCommand(cmd)
return agent.ToolResult{Content: bashContentWithError(sb.String(), "Error: command output pipes did not close after "+bashWaitDelay.String()), WorkingDir: finalWorkingDir}, nil
}
if exitErr, ok := err.(*exec.ExitError); ok {
return agent.ToolResult{Content: bashContentWithError(sb.String(), fmt.Sprintf("Exit code: %d", exitErr.ExitCode())), WorkingDir: finalWorkingDir}, nil
}
return agent.ToolResult{Content: sb.String(), WorkingDir: finalWorkingDir}, fmt.Errorf("executing command: %w", err)
}
if sb.Len() == 0 {
return agent.ToolResult{Content: "(no output)", WorkingDir: finalWorkingDir}, nil
}
return agent.ToolResult{Content: sb.String(), WorkingDir: finalWorkingDir}, nil
}
func bashContentWithError(content, msg string) string {
if content == "" {
return msg
}
return content + "\n\n" + msg
}
// rejectUnsafeShellCommand applies a best-effort blocklist for obviously
// destructive or credential-exfiltrating commands. It is defense-in-depth
// ONLY: the interactive approval prompt is the real security control, and
// this check must not be relied upon as a sandbox. Sophisticated or novel
// dangerous commands (e.g. find / -delete, dd, fork bombs, custom binaries)
// are NOT caught here and will simply be routed through approval like any
// other command. Keep the approval prompt as the gate.
func rejectUnsafeShellCommand(command string) error {
switch {
case hasUnsafeRecursiveDelete(command):
return fmt.Errorf("refusing to run unsafe command: recursive delete target is too broad")
case readsCredentialPath(command):
return fmt.Errorf("refusing to run unsafe command: credential file reads are not allowed")
default:
return nil
}
}
func hasUnsafeRecursiveDelete(command string) bool {
// Check each command segment independently. shellSafetyText flattens
// separators (; & | newlines) to spaces, which would otherwise let the
// rm target scan bleed across command boundaries — e.g.
// "rm -rf build && echo ~/.ssh/config" flattened to one token stream
// would treat the unrelated ~/.ssh/config (a ~/-prefixed "unsafe
// target") as an rm argument. Splitting on separators first restores
// command boundaries while still catching multi-target single commands
// like "rm -rf build /etc".
for _, segment := range shellSegments(command) {
fields := shellSafetyFields(segment)
for i, field := range fields {
if isRMCommand(field) && rmCommandDeletesUnsafeTarget(fields[i+1:]) {
return true
}
if isPowerShellDeleteCommand(field) && powerShellDeleteCommandDeletesUnsafeTarget(fields[i+1:]) {
return true
}
}
}
return false
}
// shellSegments splits a command on shell control operators (;, &, |, &&,
// ||) and newlines, returning the individual command segments. It operates on
// the lowercased raw command before quote/separator normalization so that
// command boundaries are preserved for per-segment checks. Subshell parens are
// intentionally NOT treated as separators: splitting on them would fragment
// command substitutions like "rm -rf $(echo /)" into "rm -rf $" and "echo /",
// hiding the destructive "/" target from the per-segment scan. Empty segments
// are dropped.
func shellSegments(command string) []string {
command = strings.ToLower(command)
var segments []string
for _, segment := range strings.FieldsFunc(command, func(r rune) bool {
switch r {
case ';', '&', '|', '\n', '\r':
return true
}
return false
}) {
if segment = strings.TrimSpace(segment); segment != "" {
segments = append(segments, segment)
}
}
return segments
}
func rmCommandDeletesUnsafeTarget(fields []string) bool {
var flags string
for _, field := range fields {
if field == "--" {
continue
}
if strings.HasPrefix(field, "-") {
flags += field
continue
}
if strings.Contains(flags, "r") && strings.Contains(flags, "f") && isUnsafeDeleteTarget(field) {
return true
}
}
return false
}
func powerShellDeleteCommandDeletesUnsafeTarget(fields []string) bool {
var recurse, force bool
var targets []string
for _, field := range fields {
switch field {
case "-r", "-recurse", "-recursive":
recurse = true
case "-f", "-force":
force = true
default:
if !strings.HasPrefix(field, "-") {
targets = append(targets, field)
}
}
}
if !recurse || !force {
return false
}
for _, target := range targets {
if isUnsafeDeleteTarget(target) {
return true
}
}
return false
}
func readsCredentialPath(command string) bool {
fields := shellSafetyFields(command)
if !hasCredentialReadVerb(fields) {
return false
}
normalized := shellSafetyText(command)
for _, fragment := range []string{
"/.ssh/id_rsa",
"/.ssh/id_dsa",
"/.ssh/id_ecdsa",
"/.ssh/id_ed25519",
"/.ssh/config",
"/.ssh/known_hosts",
"/.aws/credentials",
"/.aws/config",
"/.config/gcloud/application_default_credentials.json",
"/.kube/config",
"/.netrc",
"/.npmrc",
"/.docker/config.json",
"/.config/gh/hosts.yml",
"/.gnupg/",
"/etc/shadow",
} {
if strings.Contains(normalized, fragment) {
return true
}
}
return false
}
func hasCredentialReadVerb(fields []string) bool {
for _, field := range fields {
switch field {
case "cat", "less", "more", "head", "tail", "type", "get-content", "gc", "select-string", "grep", "rg", "sed", "awk":
return true
case "env", "printenv":
return true
}
}
return false
}
func isRMCommand(field string) bool {
return field == "rm" || strings.HasSuffix(field, "/rm")
}
func isPowerShellDeleteCommand(field string) bool {
switch field {
case "remove-item", "del", "erase", "rd", "rmdir":
return true
default:
return false
}
}
func isUnsafeDeleteTarget(target string) bool {
if target == "." || target == "./" || target == "*" {
return true
}
if target == "/*" {
return true
}
target = strings.TrimSuffix(target, "/*")
for _, prefix := range []string{"~/", "$home/", "${home}/", "$env:home/", "$env:userprofile/", "%userprofile%/"} {
if strings.HasPrefix(target, prefix) {
return true
}
}
for _, prefix := range []string{"/etc/", "/bin/", "/sbin/", "/usr/", "/var/", "/lib/", "/library/", "/system/", "/applications/", "c:/windows/", "c:/program files/"} {
if strings.HasPrefix(target, prefix) {
return true
}
}
for _, exact := range []string{"/", "~", "$home", "${home}", "$env:home", "$env:userprofile", "%userprofile%", "c:", "c:/", "/etc", "/bin", "/sbin", "/usr", "/var", "/lib", "/library", "/system", "/applications", "c:/windows", "c:/program files"} {
if target == exact {
return true
}
}
return false
}
func shellSafetyFields(command string) []string {
return strings.Fields(shellSafetyText(command))
}
func shellSafetyText(command string) string {
command = strings.ToLower(command)
return strings.NewReplacer(
"\\", "/",
"\n", " ",
"\t", " ",
";", " ",
"&", " ",
"|", " ",
"(", " ",
")", " ",
"\"", "",
"'", "",
"`", "",
).Replace(command)
}
func readFinalWorkingDir(path string) string {
content, err := os.ReadFile(path)
if err != nil {
return ""
}
workingDir := strings.TrimPrefix(string(content), "\ufeff")
workingDir = strings.TrimSpace(workingDir)
if workingDir == "" {
return ""
}
workingDir = normalizeBashWorkingDir(workingDir)
info, err := os.Stat(workingDir)
if err != nil || !info.IsDir() {
return ""
}
return workingDir
}
func normalizeBashWorkingDir(workingDir string) string {
if runtime.GOOS == "windows" && len(workingDir) >= 3 && workingDir[0] == '/' && workingDir[2] == '/' && isASCIIAlpha(workingDir[1]) {
workingDir = strings.ToUpper(string(workingDir[1])) + ":" + workingDir[2:]
}
workingDir = filepath.Clean(filepath.FromSlash(workingDir))
if runtime.GOOS == "windows" && len(workingDir) >= 2 && workingDir[1] == ':' && isASCIIAlpha(workingDir[0]) {
workingDir = strings.ToUpper(string(workingDir[0])) + workingDir[1:]
}
return workingDir
}
func isASCIIAlpha(b byte) bool {
return (b >= 'a' && b <= 'z') || (b >= 'A' && b <= 'Z')
}
type boundedOutput struct {
Limit int
buf []byte
omitted int
}
func (b *boundedOutput) Write(p []byte) (int, error) {
if b.Limit <= 0 {
b.omitted += len(p)
return len(p), nil
}
remaining := b.Limit - len(b.buf)
if remaining <= 0 {
b.omitted += len(p)
return len(p), nil
}
if len(p) <= remaining {
b.buf = append(b.buf, p...)
return len(p), nil
}
writeLen := utf8SafePrefixLen(p[:remaining])
b.buf = append(b.buf, p[:writeLen]...)
b.omitted += len(p) - writeLen
return len(p), nil
}
func (b *boundedOutput) Len() int {
return len(b.buf) + b.omitted
}
func (b *boundedOutput) String(label string) string {
safeLen := utf8SafePrefixLen(b.buf)
content := string(b.buf[:safeLen])
omitted := b.omitted + len(b.buf) - safeLen
if omitted == 0 {
return content
}
return content + agent.TruncMarker(label, safeLen, 0, omitted, false, "")
}
func utf8SafePrefixLen(p []byte) int {
if len(p) == 0 {
return 0
}
for i := 0; i < len(p); {
r, size := utf8.DecodeRune(p[i:])
if r == utf8.RuneError && size == 1 {
return i
}
i += size
}
return len(p)
}
+258
View File
@@ -0,0 +1,258 @@
package tools
import (
"context"
"os"
"path/filepath"
"runtime"
"strings"
"testing"
"unicode/utf8"
"github.com/ollama/ollama/agent"
)
func TestBashReportsFinalWorkingDir(t *testing.T) {
root := t.TempDir()
subdir := filepath.Join(root, "sub")
if err := os.Mkdir(subdir, 0o755); err != nil {
t.Fatal(err)
}
result, err := (&Bash{}).Execute(context.Background(), agent.ToolContext{WorkingDir: root}, map[string]any{
"command": shellTestCommand("cd sub && pwd", "Set-Location sub; Get-Location"),
})
if err != nil {
t.Fatal(err)
}
wantDir, err := filepath.EvalSymlinks(subdir)
if err != nil {
t.Fatal(err)
}
if result.WorkingDir != wantDir {
t.Fatalf("working dir = %q, want %q", result.WorkingDir, wantDir)
}
if !strings.Contains(result.Content, "sub") {
t.Fatalf("content = %q, want pwd output", result.Content)
}
}
func TestBashBoundsOutputWhileRunning(t *testing.T) {
result, err := (&Bash{}).Execute(context.Background(), agent.ToolContext{WorkingDir: t.TempDir()}, map[string]any{
"command": shellTestCommand("yes x | head -c 70000", "[Console]::Out.Write(('x' * 70000))"),
})
if err != nil {
t.Fatal(err)
}
if !strings.Contains(result.Content, "[stdout truncated: showing first ~") || !strings.Contains(result.Content, "omitted ~") || !strings.Contains(result.Content, " tokens.]") {
t.Fatalf("content = %q, want stdout truncation marker", result.Content)
}
if count, want := strings.Count(result.Content, "x"), shellTestCapturedXCount(); count != want {
t.Fatalf("captured x count = %d, want %d", count, want)
}
if len(result.Content) > maxBashOutputBytes+200 {
t.Fatalf("content length = %d, want bounded output", len(result.Content))
}
}
func TestBoundedOutputTruncatesAtUTF8Boundary(t *testing.T) {
var out boundedOutput
out.Limit = len([]byte("abc")) + 1
if _, err := out.Write([]byte("abcédef")); err != nil {
t.Fatal(err)
}
content := out.String("stdout")
if !utf8.ValidString(content) {
t.Fatalf("content is not valid UTF-8: %q", content)
}
if strings.ContainsRune(content, utf8.RuneError) {
t.Fatalf("content contains replacement rune: %q", content)
}
if !strings.HasPrefix(content, "abc\n\n[stdout truncated:") {
t.Fatalf("content = %q, want complete ASCII prefix and truncation marker", content)
}
}
func TestBoundedOutputKeepsCompleteUTF8AtBoundary(t *testing.T) {
var out boundedOutput
out.Limit = len([]byte("abcé"))
if _, err := out.Write([]byte("abcédef")); err != nil {
t.Fatal(err)
}
if content := out.String("stdout"); !strings.HasPrefix(content, "abcé\n\n[stdout truncated:") {
t.Fatalf("content = %q, want complete UTF-8 prefix", content)
}
}
func TestBoundedOutputTrimsTrailingPartialUTF8(t *testing.T) {
var out boundedOutput
out.Limit = 4
if _, err := out.Write([]byte{'a', 'b', 'c', 0xc3}); err != nil {
t.Fatal(err)
}
if _, err := out.Write([]byte{0xa9}); err != nil {
t.Fatal(err)
}
if content := out.String("stdout"); !utf8.ValidString(content) || !strings.HasPrefix(content, "abc\n\n[stdout truncated:") {
t.Fatalf("content = %q, want valid UTF-8 with partial suffix trimmed", content)
}
}
func TestUTF8SafePrefixRejectsMalformedLeadByte(t *testing.T) {
input := []byte{'a', 0xc0, 0x80, 'b'}
if got := utf8SafePrefixLen(input); got != 1 {
t.Fatalf("safe prefix length = %d, want 1", got)
}
}
func TestBoundedOutputDropsMalformedUTF8(t *testing.T) {
var out boundedOutput
out.Limit = 4
if _, err := out.Write([]byte{'a', 0xc0, 0x80, 'b'}); err != nil {
t.Fatal(err)
}
content := out.String("stdout")
if !utf8.ValidString(content) {
t.Fatalf("content is not valid UTF-8: %q", content)
}
if strings.ContainsRune(content, utf8.RuneError) {
t.Fatalf("content contains replacement rune: %q", content)
}
if !strings.HasPrefix(content, "a\n\n[stdout truncated:") {
t.Fatalf("content = %q, want valid prefix and truncation marker", content)
}
}
func TestBashReportsCanceledCommand(t *testing.T) {
ctx, cancel := context.WithCancel(context.Background())
cancel()
result, err := (&Bash{}).Execute(ctx, agent.ToolContext{WorkingDir: t.TempDir()}, map[string]any{
"command": shellTestCommand("sleep 10", "Start-Sleep -Seconds 10"),
})
if err != nil {
t.Fatal(err)
}
if !strings.Contains(result.Content, "Error: command was canceled") {
t.Fatalf("content = %q, want canceled message", result.Content)
}
if strings.Contains(result.Content, "Exit code: -1") {
t.Fatalf("content = %q, should not mask cancellation as exit code", result.Content)
}
}
func TestRejectUnsafeShellCommand(t *testing.T) {
tests := []struct {
name string
command string
wantErr bool
}{
{name: "rm root", command: "rm -rf /", wantErr: true},
{name: "sudo rm root", command: "sudo rm -rf -- /", wantErr: true},
{name: "rm home", command: "rm -fr $HOME", wantErr: true},
{name: "rm root wildcard", command: "rm -rf /*", wantErr: true},
{name: "rm system subdir", command: "rm -rf /etc/ssh", wantErr: true},
{name: "rm cwd", command: "rm -rf .", wantErr: true},
{name: "powershell remove root", command: `Remove-Item -Recurse -Force C:\`, wantErr: true},
{name: "powershell remove system subdir", command: `Remove-Item -Recurse -Force C:\Windows\Temp`, wantErr: true},
{name: "ssh private key", command: "cat ~/.ssh/id_rsa", wantErr: true},
{name: "aws credentials", command: "Get-Content $HOME/.aws/credentials", wantErr: true},
{name: "shadow", command: "head /etc/shadow", wantErr: true},
{name: "netrc", command: "cat ~/.netrc", wantErr: true},
{name: "docker config", command: "cat ~/.docker/config.json", wantErr: true},
{name: "gnupg dir", command: "cat ~/.gnupg/private-keys-v1.d/key", wantErr: true},
{name: "gh hosts", command: "cat ~/.config/gh/hosts.yml", wantErr: true},
{name: "ssh config", command: "cat ~/.ssh/config", wantErr: true},
{name: "printenv dump", command: "printenv", wantErr: false},
{name: "delete build dir", command: "rm -rf build", wantErr: false},
{name: "read project file", command: "cat README.md", wantErr: false},
{name: "mention key text", command: "rg id_rsa docs", wantErr: false},
{name: "env example", command: "cat .env.example", wantErr: false},
{name: "rm build then unrelated tilde path", command: "rm -rf build && echo ~/.ssh/config", wantErr: false},
{name: "rm build then unrelated slash path", command: "rm -rf build; cat /etc/passwd", wantErr: false},
{name: "rm build then unrelated star glob", command: "rm -rf build && ls *.go", wantErr: false},
{name: "rm multiple targets one unsafe", command: "rm -rf build /etc", wantErr: true},
{name: "rm unsafe then safe piped", command: "rm -rf / | tee log", wantErr: true},
{name: "rm unsafe via command substitution", command: "rm -rf $(echo /)", wantErr: true},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
err := rejectUnsafeShellCommand(tt.command)
if tt.wantErr && err == nil {
t.Fatal("expected unsafe command to be rejected")
}
if !tt.wantErr && err != nil {
t.Fatalf("command rejected: %v", err)
}
})
}
}
func TestBashRejectsUnsafeCommandBeforeExecution(t *testing.T) {
_, err := (&Bash{}).Execute(context.Background(), agent.ToolContext{WorkingDir: t.TempDir()}, map[string]any{
"command": "rm -rf /",
})
if err == nil || !strings.Contains(err.Error(), "refusing to run unsafe command") {
t.Fatalf("err = %v, want unsafe command rejection", err)
}
}
func shellTestCommand(unix, windows string) string {
if runtime.GOOS == "windows" {
return windows
}
return unix
}
func shellTestCapturedXCount() int {
if runtime.GOOS == "windows" {
return maxBashOutputBytes
}
return maxBashOutputBytes / 2
}
func TestReadFinalWorkingDirRejectsInvalidPaths(t *testing.T) {
dir := t.TempDir()
cwdFile := filepath.Join(dir, "cwd")
notDir := filepath.Join(dir, "file.txt")
if err := os.WriteFile(notDir, []byte("not a dir"), 0o644); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(cwdFile, []byte(notDir+"\n"), 0o644); err != nil {
t.Fatal(err)
}
if got := readFinalWorkingDir(cwdFile); got != "" {
t.Fatalf("regular file cwd = %q, want empty", got)
}
if err := os.WriteFile(cwdFile, []byte(filepath.Join(dir, "missing")+"\n"), 0o644); err != nil {
t.Fatal(err)
}
if got := readFinalWorkingDir(cwdFile); got != "" {
t.Fatalf("missing cwd = %q, want empty", got)
}
if err := os.WriteFile(cwdFile, []byte(dir+"\n"), 0o644); err != nil {
t.Fatal(err)
}
if got := readFinalWorkingDir(cwdFile); got != dir {
t.Fatalf("directory cwd = %q, want %q", got, dir)
}
}
func TestNormalizeBashWorkingDirWindowsDriveLetter(t *testing.T) {
if runtime.GOOS != "windows" {
t.Skip("windows path normalization")
}
got := normalizeBashWorkingDir("/c/Users/jdoe/project")
want := filepath.Clean(`C:\Users\jdoe\project`)
if got != want {
t.Fatalf("working dir = %q, want %q", got, want)
}
}
+49
View File
@@ -0,0 +1,49 @@
//go:build !windows
package tools
import (
"context"
"os/exec"
"strings"
"syscall"
)
func shellToolName() string {
return "bash"
}
func shellToolDescription() string {
return "Execute a bash command on the system. Use this to inspect files, run tests, and perform development tasks."
}
func shellCommandDescription() string {
return "The bash command to execute."
}
func newBashCommand(ctx context.Context, command, cwdPath string) *exec.Cmd {
script := command + "\n__ollama_status=$?\npwd -P > " + shellQuote(cwdPath) + "\nexit $__ollama_status"
cmd := exec.CommandContext(ctx, "bash", "-c", script)
configureBashCommand(cmd)
return cmd
}
func shellQuote(value string) string {
return "'" + strings.ReplaceAll(value, "'", "'\\''") + "'"
}
func configureBashCommand(cmd *exec.Cmd) {
cmd.SysProcAttr = &syscall.SysProcAttr{Setpgid: true}
}
func runBashCommand(cmd *exec.Cmd) error {
return cmd.Run()
}
func killBashCommand(cmd *exec.Cmd) error {
if cmd == nil || cmd.Process == nil {
return nil
}
_ = syscall.Kill(-cmd.Process.Pid, syscall.SIGKILL)
return nil
}
+40
View File
@@ -0,0 +1,40 @@
//go:build !windows
package tools
import (
"context"
"os/exec"
"strings"
"testing"
"time"
"github.com/ollama/ollama/agent"
)
func TestConfigureBashCommandSetsProcessGroup(t *testing.T) {
cmd := exec.Command("bash", "-c", "true")
configureBashCommand(cmd)
if cmd.SysProcAttr == nil || !cmd.SysProcAttr.Setpgid {
t.Fatalf("configureBashCommand should start bash in a new process group")
}
}
func TestBashWaitDelayBoundsBackgroundOutputPipe(t *testing.T) {
start := time.Now()
result, err := (&Bash{}).Execute(context.Background(), agent.ToolContext{WorkingDir: t.TempDir()}, map[string]any{
"command": "sleep 5 & echo done",
})
if err != nil {
t.Fatal(err)
}
if elapsed := time.Since(start); elapsed > bashWaitDelay+2*time.Second {
t.Fatalf("command elapsed = %s, want bounded near %s", elapsed, bashWaitDelay)
}
if !strings.Contains(result.Content, "done") {
t.Fatalf("content = %q, want command output", result.Content)
}
if !strings.Contains(result.Content, "output pipes did not close") {
t.Fatalf("content = %q, want wait delay message", result.Content)
}
}
+134
View File
@@ -0,0 +1,134 @@
//go:build windows
package tools
import (
"context"
"os/exec"
"strings"
"sync"
"unsafe"
"golang.org/x/sys/windows"
)
var bashJobHandles sync.Map
func shellToolName() string {
return "powershell"
}
func shellToolDescription() string {
return "Execute a PowerShell command on the system. Use this to inspect files, run tests, and perform development tasks."
}
func shellCommandDescription() string {
return "The PowerShell command to execute."
}
func newBashCommand(ctx context.Context, command, cwdPath string) *exec.Cmd {
return exec.CommandContext(
ctx,
"powershell.exe",
"-NoLogo",
"-NoProfile",
"-NonInteractive",
"-ExecutionPolicy",
"Bypass",
"-Command",
powerShellCommandScript(command, cwdPath),
)
}
func powerShellCommandScript(command, cwdPath string) string {
cwdPath = powerShellSingleQuote(cwdPath)
return strings.Join([]string{
"$__ollama_status = 0",
". {",
"try {",
command,
" $__ollama_success = $?",
" $__ollama_last_exit = $global:LASTEXITCODE",
" if ($__ollama_success) {",
" $__ollama_status = 0",
" } elseif ($__ollama_last_exit -is [int] -and $__ollama_last_exit -ne 0) {",
" $__ollama_status = $__ollama_last_exit",
" } else {",
" $__ollama_status = 1",
" }",
"} catch {",
" Write-Error $_",
" $__ollama_status = 1",
"} finally {",
" try { [System.IO.File]::WriteAllText(" + cwdPath + ", (Get-Location).ProviderPath, [System.Text.Encoding]::UTF8) } catch {}",
"}",
"} | Out-String -Stream -Width 4096",
"exit $__ollama_status",
}, "\n")
}
func powerShellSingleQuote(value string) string {
return "'" + strings.ReplaceAll(value, "'", "''") + "'"
}
func runBashCommand(cmd *exec.Cmd) error {
if err := cmd.Start(); err != nil {
return err
}
if job, err := createBashJob(cmd.Process.Pid); err == nil {
bashJobHandles.Store(cmd.Process.Pid, job)
defer releaseBashJob(cmd.Process.Pid)
}
return cmd.Wait()
}
func killBashCommand(cmd *exec.Cmd) error {
if cmd == nil || cmd.Process == nil {
return nil
}
releaseBashJob(cmd.Process.Pid)
_ = cmd.Process.Kill()
return nil
}
func createBashJob(pid int) (windows.Handle, error) {
job, err := windows.CreateJobObject(nil, nil)
if err != nil {
return 0, err
}
info := windows.JOBOBJECT_EXTENDED_LIMIT_INFORMATION{}
info.BasicLimitInformation.LimitFlags = windows.JOB_OBJECT_LIMIT_KILL_ON_JOB_CLOSE
if _, err := windows.SetInformationJobObject(
job,
windows.JobObjectExtendedLimitInformation,
uintptr(unsafe.Pointer(&info)),
uint32(unsafe.Sizeof(info)),
); err != nil {
_ = windows.CloseHandle(job)
return 0, err
}
process, err := windows.OpenProcess(windows.PROCESS_SET_QUOTA|windows.PROCESS_TERMINATE, false, uint32(pid))
if err != nil {
_ = windows.CloseHandle(job)
return 0, err
}
defer windows.CloseHandle(process)
if err := windows.AssignProcessToJobObject(job, process); err != nil {
_ = windows.CloseHandle(job)
return 0, err
}
return job, nil
}
func releaseBashJob(pid int) {
value, ok := bashJobHandles.LoadAndDelete(pid)
if !ok {
return
}
if job, ok := value.(windows.Handle); ok {
_ = windows.CloseHandle(job)
}
}
+15
View File
@@ -0,0 +1,15 @@
//go:build windows
package tools
import (
"strings"
"testing"
)
func TestPowerShellCommandScriptUsesWideOutString(t *testing.T) {
script := powerShellCommandScript("Get-ChildItem", `C:\cwd.txt`)
if !strings.Contains(script, "Out-String -Stream -Width 4096") {
t.Fatalf("script = %q, want explicit Out-String width", script)
}
}
+558
View File
@@ -0,0 +1,558 @@
package tools
import (
"bufio"
"context"
"fmt"
"io"
"os"
"path/filepath"
"strconv"
"strings"
"github.com/ollama/ollama/agent"
"github.com/ollama/ollama/api"
)
const (
maxReadBytes = 200000
)
type Read struct{}
func (r *Read) Name() string {
return "read"
}
func (r *Read) Description() string {
return "Read a text file from the current working directory."
}
func (r *Read) Schema() api.ToolFunction {
props := api.NewToolPropertiesMap()
props.Set("path", api.ToolProperty{
Type: api.PropertyType{"string"},
Description: "Path to the file to read, relative to the working directory.",
})
props.Set("start", api.ToolProperty{
Type: api.PropertyType{"integer"},
Description: "Optional 1-based line to start reading from.",
})
props.Set("end", api.ToolProperty{
Type: api.PropertyType{"integer"},
Description: "Optional 1-based inclusive line to stop reading at.",
})
return api.ToolFunction{
Name: r.Name(),
Description: r.Description(),
Parameters: api.ToolFunctionParameters{
Type: "object",
Properties: props,
Required: []string{"path"},
},
}
}
func (r *Read) RequiresApproval(map[string]any) bool {
return true
}
func (r *Read) Execute(ctx context.Context, toolCtx agent.ToolContext, args map[string]any) (agent.ToolResult, error) {
// TODO: use shared agent.RequiredStringArg / agent.OptionalIntArg for args (see agent package cleanup plan).
path, ok := args["path"].(string)
if !ok || strings.TrimSpace(path) == "" {
return agent.ToolResult{}, fmt.Errorf("path parameter is required")
}
file, info, err := openRegularFile(toolCtx.WorkingDir, path, true)
if err != nil {
return agent.ToolResult{}, err
}
defer file.Close()
selection, err := readSelectionFromArgs(args)
if err != nil {
return agent.ToolResult{}, err
}
if !selection.enabled && info.Size() > maxReadBytes {
return agent.ToolResult{}, fmt.Errorf("%s is too large to read (%d bytes)", path, info.Size())
}
select {
case <-ctx.Done():
return agent.ToolResult{}, ctx.Err()
default:
}
var content string
if selection.enabled {
content, err = readLineSelection(file, selection)
} else {
var contentBytes []byte
contentBytes, err = readAllWithinLimit(file, maxReadBytes)
content = string(contentBytes)
}
if err != nil {
return agent.ToolResult{}, err
}
return agent.ToolResult{Content: content}, nil
}
type Edit struct{}
func (e *Edit) Name() string {
return "edit"
}
func (e *Edit) Description() string {
return "Edit a text file in the current working directory by replacing exact text."
}
func (e *Edit) Schema() api.ToolFunction {
props := api.NewToolPropertiesMap()
props.Set("path", api.ToolProperty{
Type: api.PropertyType{"string"},
Description: "Path to the file to edit, relative to the working directory.",
})
props.Set("old_text", api.ToolProperty{
Type: api.PropertyType{"string"},
Description: "Exact text to replace.",
})
props.Set("new_text", api.ToolProperty{
Type: api.PropertyType{"string"},
Description: "Replacement text.",
})
props.Set("replace_all", api.ToolProperty{
Type: api.PropertyType{"boolean"},
Description: "Replace every occurrence. Defaults to false and requires old_text to match exactly once.",
})
return api.ToolFunction{
Name: e.Name(),
Description: e.Description(),
Parameters: api.ToolFunctionParameters{
Type: "object",
Properties: props,
Required: []string{"path", "old_text", "new_text"},
},
}
}
func (e *Edit) RequiresApproval(map[string]any) bool {
return true
}
func (e *Edit) Execute(ctx context.Context, toolCtx agent.ToolContext, args map[string]any) (agent.ToolResult, error) {
// TODO: use shared agent.RequiredStringArg / agent.OptionalBoolArg for args (see agent package cleanup plan).
path, ok := args["path"].(string)
if !ok || strings.TrimSpace(path) == "" {
return agent.ToolResult{}, fmt.Errorf("path parameter is required")
}
oldText, ok := args["old_text"].(string)
if !ok || oldText == "" {
return agent.ToolResult{}, fmt.Errorf("old_text parameter is required")
}
newText, ok := args["new_text"].(string)
if !ok {
return agent.ToolResult{}, fmt.Errorf("new_text parameter is required")
}
replaceAll, _ := args["replace_all"].(bool)
if err := rejectFinalSymlink(toolCtx.WorkingDir, path); err != nil {
return agent.ToolResult{}, err
}
file, info, err := openRegularFile(toolCtx.WorkingDir, path, false)
if err != nil {
return agent.ToolResult{}, err
}
if info.Size() > maxReadBytes {
file.Close()
return agent.ToolResult{}, fmt.Errorf("%s is too large to edit (%d bytes)", path, info.Size())
}
select {
case <-ctx.Done():
file.Close()
return agent.ToolResult{}, ctx.Err()
default:
}
contentBytes, err := readAllWithinLimit(file, maxReadBytes)
if closeErr := file.Close(); err == nil && closeErr != nil {
err = closeErr
}
if err != nil {
return agent.ToolResult{}, err
}
content := string(contentBytes)
matches := strings.Count(content, oldText)
if matches == 0 {
return agent.ToolResult{}, fmt.Errorf("old_text was not found in %s", path)
}
if matches > 1 && !replaceAll {
return agent.ToolResult{}, fmt.Errorf("old_text matched %d times in %s; set replace_all to true to replace every match", matches, path)
}
var updated string
if replaceAll {
updated = strings.ReplaceAll(content, oldText, newText)
} else {
updated = strings.Replace(content, oldText, newText, 1)
}
if len(updated) > maxReadBytes {
return agent.ToolResult{}, fmt.Errorf("edited content is too large (%d bytes)", len(updated))
}
if err := writeFileAtomic(toolCtx.WorkingDir, path, []byte(updated), info.Mode().Perm()); err != nil {
return agent.ToolResult{}, err
}
return agent.ToolResult{Content: fmt.Sprintf("Updated %s (%d replacement%s).", path, matches, plural(matches))}, nil
}
func cleanRelativePath(path string) (string, error) {
path = strings.TrimSpace(path)
if path == "" {
return "", fmt.Errorf("path parameter is required")
}
if filepath.IsAbs(path) {
return "", fmt.Errorf("absolute paths are not allowed")
}
cleaned := filepath.Clean(path)
if cleaned == "." || cleaned == ".." || strings.HasPrefix(cleaned, ".."+string(os.PathSeparator)) {
return "", fmt.Errorf("path escapes working directory")
}
return cleaned, nil
}
func openRegularFile(workingDir, path string, allowAbsolute bool) (*os.File, os.FileInfo, error) {
path = strings.TrimSpace(path)
if path == "" {
return nil, nil, fmt.Errorf("path parameter is required")
}
if allowAbsolute && filepath.IsAbs(path) {
cleaned := filepath.Clean(path)
info, err := os.Lstat(cleaned)
if err != nil {
return nil, nil, err
}
if info.Mode()&os.ModeSymlink != 0 {
return nil, nil, fmt.Errorf("%s is a symlink; read the target file directly", path)
}
if err := rejectNonRegularFile(path, info); err != nil {
return nil, nil, err
}
file, err := os.Open(cleaned)
if err != nil {
return nil, nil, err
}
info, err = file.Stat()
if err != nil {
file.Close()
return nil, nil, err
}
if err := rejectNonRegularFile(path, info); err != nil {
file.Close()
return nil, nil, err
}
return file, info, nil
}
rel, err := cleanRelativePath(path)
if err != nil {
return nil, nil, err
}
root, err := openWorkingRoot(workingDir)
if err != nil {
return nil, nil, err
}
defer root.Close()
if _, err := regularRootFileInfo(root, rel, path); err != nil {
return nil, nil, err
}
file, err := root.Open(rel)
if err != nil {
return nil, nil, rootPathError(err)
}
info, err := file.Stat()
if err != nil {
file.Close()
return nil, nil, err
}
if err := rejectNonRegularFile(path, info); err != nil {
file.Close()
return nil, nil, err
}
return file, info, nil
}
func regularRootFileInfo(root *os.Root, rel, path string) (os.FileInfo, error) {
info, err := root.Lstat(rel)
if err != nil {
return nil, rootPathError(err)
}
// Reject symlinks outright. os.Root.Open follows symlinks via openat
// without O_NOFOLLOW, so a symlink inside the working root that points
// outside it (e.g. ./notes -> ~/.ssh/id_rsa) would otherwise be read
// transparently, bypassing the working-directory confinement that the
// bash denylist enforces for direct credential reads. The caller must
// operate on the real target file instead.
if info.Mode()&os.ModeSymlink != 0 {
return nil, fmt.Errorf("%s is a symlink; read the target file directly", path)
}
if err := rejectNonRegularFile(path, info); err != nil {
return nil, err
}
return info, nil
}
func rejectNonRegularFile(path string, info os.FileInfo) error {
if info.IsDir() {
return fmt.Errorf("%s is a directory", path)
}
if !info.Mode().IsRegular() {
return fmt.Errorf("%s is not a regular file", path)
}
return nil
}
func writeFileAtomic(workingDir, path string, data []byte, perm os.FileMode) error {
rel, err := cleanRelativePath(path)
if err != nil {
return err
}
root, err := openWorkingRoot(workingDir)
if err != nil {
return err
}
defer root.Close()
if err := rejectRootFinalSymlink(root, rel, path); err != nil {
return err
}
parent, name := filepath.Split(rel)
tmpBase := fmt.Sprintf(".%s.ollama-tmp-%d", name, os.Getpid())
for i := 0; ; i++ {
candidateName := tmpBase
if i > 0 {
candidateName = fmt.Sprintf("%s-%d", tmpBase, i)
}
candidate := filepath.Join(parent, candidateName)
file, err := root.OpenFile(candidate, os.O_WRONLY|os.O_CREATE|os.O_EXCL, perm)
if os.IsExist(err) {
continue
}
if err != nil {
return rootPathError(err)
}
if err := file.Chmod(perm); err != nil {
closeErr := file.Close()
_ = root.Remove(candidate)
if closeErr != nil {
return closeErr
}
return err
}
writeErr := writeAllAndSync(file, data)
closeErr := file.Close()
if writeErr != nil || closeErr != nil {
_ = root.Remove(candidate)
if writeErr != nil {
return writeErr
}
return closeErr
}
if err := root.Rename(candidate, rel); err != nil {
_ = root.Remove(candidate)
return rootPathError(err)
}
return nil
}
}
func rejectFinalSymlink(workingDir, path string) error {
rel, err := cleanRelativePath(path)
if err != nil {
return err
}
root, err := openWorkingRoot(workingDir)
if err != nil {
return err
}
defer root.Close()
return rejectRootFinalSymlink(root, rel, path)
}
func rejectRootFinalSymlink(root *os.Root, rel, path string) error {
info, err := root.Lstat(rel)
if err != nil {
return rootPathError(err)
}
if info.Mode()&os.ModeSymlink != 0 {
return fmt.Errorf("%s is a symlink; edit the target file directly", path)
}
return nil
}
func rootPathError(err error) error {
if err != nil && strings.Contains(err.Error(), "path escapes") {
return fmt.Errorf("path escapes working directory")
}
return err
}
func openWorkingRoot(workingDir string) (*os.Root, error) {
base, err := workingDirAbs(workingDir)
if err != nil {
return nil, err
}
return os.OpenRoot(base)
}
func writeAllAndSync(file *os.File, data []byte) error {
if _, err := file.Write(data); err != nil {
return err
}
return file.Sync()
}
func readAllWithinLimit(reader io.Reader, limit int) ([]byte, error) {
if limit < 0 {
limit = 0
}
content, err := io.ReadAll(io.LimitReader(reader, int64(limit)+1))
if err != nil {
return nil, err
}
if len(content) > limit {
return nil, fmt.Errorf("content is too large (%d byte limit)", limit)
}
return content, nil
}
func workingDirAbs(workingDir string) (string, error) {
base := workingDir
if base == "" {
var err error
base, err = os.Getwd()
if err != nil {
return "", err
}
}
return canonicalPath(base)
}
func canonicalPath(path string) (string, error) {
abs, err := filepath.Abs(path)
if err != nil {
return "", err
}
resolved, err := filepath.EvalSymlinks(abs)
if err == nil {
return resolved, nil
}
return abs, nil
}
type readSelection struct {
enabled bool
start int
end int
}
func readSelectionFromArgs(args map[string]any) (readSelection, error) {
selection := readSelection{start: 1}
if start, ok, err := intReadArg(args, "start"); err != nil {
return readSelection{}, err
} else if ok {
selection.enabled = true
selection.start = start
}
if end, ok, err := intReadArg(args, "end"); err != nil {
return readSelection{}, err
} else if ok {
selection.enabled = true
selection.end = end
}
if !selection.enabled {
return selection, nil
}
if selection.start < 1 {
return readSelection{}, fmt.Errorf("start must be greater than 0")
}
if selection.end > 0 && selection.end < selection.start {
return readSelection{}, fmt.Errorf("end must be greater than or equal to start")
}
return selection, nil
}
func readLineSelection(file *os.File, selection readSelection) (string, error) {
reader := bufio.NewReader(file)
var b strings.Builder
for lineNo := 1; ; {
line, err := reader.ReadSlice('\n')
if lineNo >= selection.start && (selection.end == 0 || lineNo <= selection.end) {
if b.Len()+len(line) > maxReadBytes {
return "", fmt.Errorf("selected content is too large (%d byte limit)", maxReadBytes)
}
b.Write(line)
}
if err != nil {
if err == bufio.ErrBufferFull {
continue
}
if err == io.EOF {
break
}
return "", err
}
if selection.end > 0 && lineNo >= selection.end {
break
}
lineNo++
}
return b.String(), nil
}
func intReadArg(args map[string]any, key string) (int, bool, error) {
value, ok := args[key]
if !ok {
return 0, false, nil
}
switch v := value.(type) {
case int:
return v, true, nil
case int64:
return int(v), true, nil
case float64:
if v != float64(int(v)) {
return 0, true, fmt.Errorf("%s must be a whole number", key)
}
return int(v), true, nil
case string:
v = strings.TrimSpace(v)
if v == "" {
return 0, false, nil
}
n, err := strconv.Atoi(v)
if err != nil {
return 0, true, fmt.Errorf("%s must be a whole number", key)
}
return n, true, nil
default:
return 0, true, fmt.Errorf("%s must be a whole number", key)
}
}
func plural(n int) string {
if n == 1 {
return ""
}
return "s"
}
+338
View File
@@ -0,0 +1,338 @@
package tools
import (
"context"
"io"
"os"
"path/filepath"
"strings"
"testing"
"github.com/ollama/ollama/agent"
)
func TestEditReplacesUniqueText(t *testing.T) {
dir := t.TempDir()
path := filepath.Join(dir, "note.txt")
if err := os.WriteFile(path, []byte("hello world\n"), 0o644); err != nil {
t.Fatal(err)
}
result, err := (&Edit{}).Execute(context.Background(), agent.ToolContext{WorkingDir: dir}, map[string]any{
"path": "note.txt",
"old_text": "hello",
"new_text": "hi",
})
if err != nil {
t.Fatal(err)
}
if !strings.Contains(result.Content, "Updated note.txt") {
t.Fatalf("result = %q", result.Content)
}
content, err := os.ReadFile(path)
if err != nil {
t.Fatal(err)
}
if string(content) != "hi world\n" {
t.Fatalf("content = %q", content)
}
}
func TestEditRequiresUniqueMatchByDefault(t *testing.T) {
dir := t.TempDir()
path := filepath.Join(dir, "note.txt")
if err := os.WriteFile(path, []byte("same same\n"), 0o644); err != nil {
t.Fatal(err)
}
_, err := (&Edit{}).Execute(context.Background(), agent.ToolContext{WorkingDir: dir}, map[string]any{
"path": "note.txt",
"old_text": "same",
"new_text": "other",
})
if err == nil {
t.Fatal("expected ambiguous edit to fail")
}
if !strings.Contains(err.Error(), "matched 2 times") {
t.Fatalf("err = %v", err)
}
}
func TestEditRejectsEscapingPath(t *testing.T) {
dir := t.TempDir()
_, err := (&Edit{}).Execute(context.Background(), agent.ToolContext{WorkingDir: dir}, map[string]any{
"path": "../outside.txt",
"old_text": "old",
"new_text": "new",
})
if err == nil {
t.Fatal("expected escaping path to fail")
}
if !strings.Contains(err.Error(), "path escapes working directory") {
t.Fatalf("err = %v", err)
}
}
func TestEditRejectsSymlinkEscape(t *testing.T) {
dir := t.TempDir()
outside := t.TempDir()
if err := os.WriteFile(filepath.Join(outside, "note.txt"), []byte("old\n"), 0o644); err != nil {
t.Fatal(err)
}
if err := os.Symlink(outside, filepath.Join(dir, "link")); err != nil {
t.Skipf("symlinks unavailable: %v", err)
}
_, err := (&Edit{}).Execute(context.Background(), agent.ToolContext{WorkingDir: dir}, map[string]any{
"path": filepath.Join("link", "note.txt"),
"old_text": "old",
"new_text": "new",
})
if err == nil {
t.Fatal("expected symlink escape to fail")
}
if !strings.Contains(err.Error(), "path escapes working directory") {
t.Fatalf("err = %v", err)
}
content, err := os.ReadFile(filepath.Join(outside, "note.txt"))
if err != nil {
t.Fatal(err)
}
if string(content) != "old\n" {
t.Fatalf("outside content changed to %q", content)
}
}
func TestEditRejectsFinalSymlink(t *testing.T) {
dir := t.TempDir()
target := filepath.Join(dir, "target.txt")
if err := os.WriteFile(target, []byte("old\n"), 0o644); err != nil {
t.Fatal(err)
}
link := filepath.Join(dir, "link.txt")
if err := os.Symlink("target.txt", link); err != nil {
t.Skipf("symlinks unavailable: %v", err)
}
_, err := (&Edit{}).Execute(context.Background(), agent.ToolContext{WorkingDir: dir}, map[string]any{
"path": "link.txt",
"old_text": "old",
"new_text": "new",
})
if err == nil {
t.Fatal("expected final symlink edit to fail")
}
if !strings.Contains(err.Error(), "is a symlink") {
t.Fatalf("err = %v", err)
}
content, err := os.ReadFile(target)
if err != nil {
t.Fatal(err)
}
if string(content) != "old\n" {
t.Fatalf("target content changed to %q", content)
}
info, err := os.Lstat(link)
if err != nil {
t.Fatal(err)
}
if info.Mode()&os.ModeSymlink == 0 {
t.Fatalf("link mode = %v, want symlink", info.Mode())
}
}
func TestReadRejectsParentOutsideCurrentWorkingDir(t *testing.T) {
root := t.TempDir()
subdir := filepath.Join(root, "sub")
if err := os.Mkdir(subdir, 0o755); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(filepath.Join(root, "note.txt"), []byte("hello"), 0o644); err != nil {
t.Fatal(err)
}
_, err := (&Read{}).Execute(context.Background(), agent.ToolContext{WorkingDir: subdir}, map[string]any{
"path": "../note.txt",
})
if err == nil {
t.Fatal("expected parent path to fail")
}
if !strings.Contains(err.Error(), "path escapes working directory") {
t.Fatalf("err = %v", err)
}
}
func TestReadRequiresApproval(t *testing.T) {
if !agent.ToolRequiresApproval((&Read{}), map[string]any{"path": "note.txt"}) {
t.Fatal("read should require approval")
}
}
func TestReadDefaultsToEntireFile(t *testing.T) {
dir := t.TempDir()
content := "one\ntwo\nthree\n"
if err := os.WriteFile(filepath.Join(dir, "note.txt"), []byte(content), 0o644); err != nil {
t.Fatal(err)
}
result, err := (&Read{}).Execute(context.Background(), agent.ToolContext{WorkingDir: dir}, map[string]any{
"path": "note.txt",
})
if err != nil {
t.Fatal(err)
}
if result.Content != content {
t.Fatalf("content = %q", result.Content)
}
}
func TestReadAllowsAbsolutePath(t *testing.T) {
dir := t.TempDir()
path := filepath.Join(dir, "note.txt")
content := "one\ntwo\nthree\n"
if err := os.WriteFile(path, []byte(content), 0o644); err != nil {
t.Fatal(err)
}
result, err := (&Read{}).Execute(context.Background(), agent.ToolContext{WorkingDir: t.TempDir()}, map[string]any{
"path": path,
})
if err != nil {
t.Fatal(err)
}
if result.Content != content {
t.Fatalf("content = %q", result.Content)
}
}
func TestReadRejectsAbsoluteSymlink(t *testing.T) {
dir := t.TempDir()
target := filepath.Join(dir, "target.txt")
if err := os.WriteFile(target, []byte("hello\n"), 0o644); err != nil {
t.Fatal(err)
}
link := filepath.Join(dir, "alias")
if err := os.Symlink(target, link); err != nil {
t.Skipf("symlinks unavailable: %v", err)
}
_, err := (&Read{}).Execute(context.Background(), agent.ToolContext{WorkingDir: t.TempDir()}, map[string]any{
"path": link,
})
if err == nil {
t.Fatal("expected absolute symlink to be rejected")
}
if !strings.Contains(err.Error(), "symlink") {
t.Fatalf("err = %v, want symlink rejection", err)
}
}
func TestReadStartEnd(t *testing.T) {
dir := t.TempDir()
if err := os.WriteFile(filepath.Join(dir, "note.txt"), []byte("one\ntwo\nthree\nfour\n"), 0o644); err != nil {
t.Fatal(err)
}
result, err := (&Read{}).Execute(context.Background(), agent.ToolContext{WorkingDir: dir}, map[string]any{
"path": "note.txt",
"start": 2,
"end": 3,
})
if err != nil {
t.Fatal(err)
}
if result.Content != "two\nthree\n" {
t.Fatalf("content = %q", result.Content)
}
}
func TestReadStartOnly(t *testing.T) {
dir := t.TempDir()
if err := os.WriteFile(filepath.Join(dir, "note.txt"), []byte("one\ntwo\nthree\nfour\n"), 0o644); err != nil {
t.Fatal(err)
}
result, err := (&Read{}).Execute(context.Background(), agent.ToolContext{WorkingDir: dir}, map[string]any{
"path": "note.txt",
"start": 3,
})
if err != nil {
t.Fatal(err)
}
if result.Content != "three\nfour\n" {
t.Fatalf("content = %q", result.Content)
}
}
func TestReadEndOnly(t *testing.T) {
dir := t.TempDir()
if err := os.WriteFile(filepath.Join(dir, "note.txt"), []byte("one\ntwo\nthree\nfour\n"), 0o644); err != nil {
t.Fatal(err)
}
result, err := (&Read{}).Execute(context.Background(), agent.ToolContext{WorkingDir: dir}, map[string]any{
"path": "note.txt",
"end": 2,
})
if err != nil {
t.Fatal(err)
}
if result.Content != "one\ntwo\n" {
t.Fatalf("content = %q", result.Content)
}
}
func TestReadSelectionRejectsHugeSingleLine(t *testing.T) {
dir := t.TempDir()
if err := os.WriteFile(filepath.Join(dir, "note.txt"), []byte(strings.Repeat("x", maxReadBytes+1)), 0o644); err != nil {
t.Fatal(err)
}
_, err := (&Read{}).Execute(context.Background(), agent.ToolContext{WorkingDir: dir}, map[string]any{
"path": "note.txt",
"start": 1,
"end": 1,
})
if err == nil {
t.Fatal("expected huge selected line to fail")
}
if !strings.Contains(err.Error(), "selected content is too large") {
t.Fatalf("err = %v", err)
}
}
func TestReadAllWithinLimitRejectsGrowingRead(t *testing.T) {
reader := io.MultiReader(
strings.NewReader(strings.Repeat("x", maxReadBytes)),
strings.NewReader("x"),
)
_, err := readAllWithinLimit(reader, maxReadBytes)
if err == nil {
t.Fatal("expected over-limit read to fail")
}
if !strings.Contains(err.Error(), "content is too large") {
t.Fatalf("err = %v", err)
}
}
func TestReadRejectsInvalidRange(t *testing.T) {
dir := t.TempDir()
if err := os.WriteFile(filepath.Join(dir, "note.txt"), []byte("one\ntwo\n"), 0o644); err != nil {
t.Fatal(err)
}
_, err := (&Read{}).Execute(context.Background(), agent.ToolContext{WorkingDir: dir}, map[string]any{
"path": "note.txt",
"start": 4,
"end": 2,
})
if err == nil {
t.Fatal("expected invalid range to fail")
}
if !strings.Contains(err.Error(), "end must") {
t.Fatalf("err = %v", err)
}
}
+121
View File
@@ -0,0 +1,121 @@
//go:build !windows
package tools
import (
"context"
"os"
"path/filepath"
"strings"
"syscall"
"testing"
"time"
"github.com/ollama/ollama/agent"
)
func TestOpenRegularFileRejectsFIFO(t *testing.T) {
dir := t.TempDir()
path := filepath.Join(dir, "pipe")
if err := syscall.Mkfifo(path, 0o600); err != nil {
t.Skipf("mkfifo unavailable: %v", err)
}
done := make(chan error, 1)
go func() {
file, _, err := openRegularFile(dir, "pipe", false)
if file != nil {
file.Close()
}
done <- err
}()
select {
case err := <-done:
if err == nil {
t.Fatal("expected FIFO to be rejected")
}
if !strings.Contains(err.Error(), "not a regular file") {
t.Fatalf("err = %v", err)
}
case <-time.After(time.Second):
t.Fatal("openRegularFile blocked on FIFO")
}
}
func TestEditPreservesModeDespiteUmask(t *testing.T) {
oldUmask := syscall.Umask(0o077)
defer syscall.Umask(oldUmask)
dir := t.TempDir()
path := filepath.Join(dir, "note.txt")
if err := os.WriteFile(path, []byte("hello\n"), 0o666); err != nil {
t.Fatal(err)
}
if err := os.Chmod(path, 0o666); err != nil {
t.Fatal(err)
}
_, err := (&Edit{}).Execute(context.Background(), agent.ToolContext{WorkingDir: dir}, map[string]any{
"path": "note.txt",
"old_text": "hello",
"new_text": "hi",
})
if err != nil {
t.Fatal(err)
}
info, err := os.Stat(path)
if err != nil {
t.Fatal(err)
}
if got := info.Mode().Perm(); got != 0o666 {
t.Fatalf("mode = %#o, want 0666", got)
}
}
func TestReadRejectsSymlinkEscapingWorkingDir(t *testing.T) {
root := t.TempDir()
secret := filepath.Join(t.TempDir(), "secret.txt")
if err := os.WriteFile(secret, []byte("top secret\n"), 0o600); err != nil {
t.Fatal(err)
}
link := filepath.Join(root, "notes")
if err := os.Symlink(secret, link); err != nil {
t.Fatal(err)
}
_, err := (&Read{}).Execute(context.Background(), agent.ToolContext{WorkingDir: root}, map[string]any{
"path": "notes",
})
if err == nil {
t.Fatal("expected symlink escaping working dir to be rejected")
}
if !strings.Contains(err.Error(), "symlink") {
t.Fatalf("err = %v, want symlink rejection", err)
}
}
func TestReadRejectsSymlinkInsideWorkingDirToOutside(t *testing.T) {
root := t.TempDir()
target := filepath.Join(root, "real.txt")
if err := os.WriteFile(target, []byte("hello\n"), 0o644); err != nil {
t.Fatal(err)
}
// A symlink to a sibling file still resolves inside the root; Read must
// reject it regardless, consistent with Edit's rejectFinalSymlink.
link := filepath.Join(root, "alias")
if err := os.Symlink(target, link); err != nil {
t.Fatal(err)
}
_, err := (&Read{}).Execute(context.Background(), agent.ToolContext{WorkingDir: root}, map[string]any{
"path": "alias",
})
if err == nil {
t.Fatal("expected symlink to be rejected even when target is inside root")
}
if !strings.Contains(err.Error(), "symlink") {
t.Fatalf("err = %v, want symlink rejection", err)
}
}
+38
View File
@@ -0,0 +1,38 @@
package tools
import (
"context"
"errors"
"github.com/ollama/ollama/agent"
"github.com/ollama/ollama/api"
)
// Skill is the model-facing adapter for the core agent skill catalog.
// It only supplies instructions; regular tools retain their own approval
// requirements for filesystem or network access.
type Skill struct{ Catalog *agent.SkillCatalog }
func (t *Skill) Name() string { return "skill" }
func (t *Skill) Description() string {
return "Load a named Ollama skill and return its instructions."
}
func (t *Skill) Schema() api.ToolFunction {
props := api.NewToolPropertiesMap()
props.Set("name", api.ToolProperty{Type: api.PropertyType{"string"}, Description: "Name of the skill to load."})
return api.ToolFunction{Name: t.Name(), Description: t.Description(), Parameters: api.ToolFunctionParameters{Type: "object", Properties: props, Required: []string{"name"}}}
}
func (t *Skill) Execute(_ context.Context, _ agent.ToolContext, args map[string]any) (agent.ToolResult, error) {
name, ok := args["name"].(string)
if !ok {
return agent.ToolResult{}, errors.New("name parameter is required")
}
skill, err := t.Catalog.Load(name)
if err != nil {
return agent.ToolResult{}, err
}
return agent.ToolResult{Content: skill.Content()}, nil
}
+34
View File
@@ -0,0 +1,34 @@
package tools
import (
"context"
"os"
"path/filepath"
"strings"
"testing"
"github.com/ollama/ollama/agent"
)
func TestSkillLoadsCoreCatalogWithoutApproval(t *testing.T) {
dir := t.TempDir()
path := filepath.Join(dir, "release-notes")
if err := os.Mkdir(path, 0o755); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(filepath.Join(path, "SKILL.md"), []byte("---\nname: release-notes\ndescription: Draft release notes.\n---\nUse concise bullets."), 0o644); err != nil {
t.Fatal(err)
}
catalog, err := agent.DiscoverSkills(dir)
if err != nil {
t.Fatal(err)
}
tool := &Skill{Catalog: catalog}
if agent.ToolRequiresApproval(tool, map[string]any{"name": "release-notes"}) {
t.Fatal("loading a skill must not change ordinary tool approval semantics")
}
result, err := tool.Execute(context.Background(), agent.ToolContext{}, map[string]any{"name": "release-notes"})
if err != nil || !strings.Contains(result.Content, "Use concise bullets.") {
t.Fatalf("tool result = %#v, %v", result, err)
}
}
+186
View File
@@ -0,0 +1,186 @@
package tools
import (
"context"
"errors"
"fmt"
"net/url"
"strings"
"time"
"github.com/ollama/ollama/agent"
"github.com/ollama/ollama/api"
internalcloud "github.com/ollama/ollama/internal/cloud"
)
const (
maxWebFetchContentRunes = 60_000
webSearchTimeout = 15 * time.Second
webFetchTimeout = 30 * time.Second
)
var ErrWebAuthRequired = errors.New("Not authenticated. Run `ollama signin` and try again.")
type WebSearch struct{}
func (w *WebSearch) Name() string {
return "web_search"
}
func (w *WebSearch) Description() string {
return "Search the web for current information that may not be in the model's training data."
}
func (w *WebSearch) Schema() api.ToolFunction {
props := api.NewToolPropertiesMap()
props.Set("query", api.ToolProperty{
Type: api.PropertyType{"string"},
Description: "The search query to look up on the web.",
})
return api.ToolFunction{
Name: w.Name(),
Description: w.Description(),
Parameters: api.ToolFunctionParameters{
Type: "object",
Properties: props,
Required: []string{"query"},
},
}
}
func (w *WebSearch) RequiresApproval(map[string]any) bool {
return true
}
func (w *WebSearch) Execute(ctx context.Context, _ agent.ToolContext, args map[string]any) (agent.ToolResult, error) {
// TODO: use shared agent.RequiredStringArg for the "query" parameter (see agent package cleanup plan).
if internalcloud.Disabled() {
return agent.ToolResult{}, errors.New(internalcloud.DisabledError("web search is unavailable"))
}
query, ok := args["query"].(string)
if !ok || strings.TrimSpace(query) == "" {
return agent.ToolResult{}, fmt.Errorf("query parameter is required")
}
client, err := api.ClientFromEnvironment()
if err != nil {
return agent.ToolResult{}, err
}
ctx, cancel := context.WithTimeout(ctx, webSearchTimeout)
defer cancel()
searchResp, err := client.WebSearchExperimental(ctx, &api.WebSearchRequest{Query: query, MaxResults: 5})
if err != nil {
var authErr api.AuthorizationError
if errors.As(err, &authErr) {
return agent.ToolResult{}, ErrWebAuthRequired
}
return agent.ToolResult{}, err
}
if len(searchResp.Results) == 0 {
return agent.ToolResult{Content: "No results found for query: " + query}, nil
}
var sb strings.Builder
sb.WriteString(fmt.Sprintf("Search results for: %s\n\n", query))
for i, result := range searchResp.Results {
sb.WriteString(fmt.Sprintf("%d. %s\n", i+1, result.Title))
sb.WriteString(fmt.Sprintf(" URL: %s\n", result.URL))
if result.Content != "" {
content := []rune(result.Content)
if len(content) > 300 {
content = append(content[:300], []rune("...")...)
}
sb.WriteString(fmt.Sprintf(" %s\n", string(content)))
}
sb.WriteByte('\n')
}
return agent.ToolResult{Content: sb.String()}, nil
}
type WebFetch struct{}
func (w *WebFetch) Name() string {
return "web_fetch"
}
func (w *WebFetch) Description() string {
return "Fetch and extract text content from a web page."
}
func (w *WebFetch) Schema() api.ToolFunction {
props := api.NewToolPropertiesMap()
props.Set("url", api.ToolProperty{
Type: api.PropertyType{"string"},
Description: "The URL to fetch and extract content from.",
})
return api.ToolFunction{
Name: w.Name(),
Description: w.Description(),
Parameters: api.ToolFunctionParameters{
Type: "object",
Properties: props,
Required: []string{"url"},
},
}
}
func (w *WebFetch) RequiresApproval(map[string]any) bool {
return true
}
func (w *WebFetch) Execute(ctx context.Context, _ agent.ToolContext, args map[string]any) (agent.ToolResult, error) {
// TODO: use shared agent.RequiredStringArg for the "url" parameter (see agent package cleanup plan).
if internalcloud.Disabled() {
return agent.ToolResult{}, errors.New(internalcloud.DisabledError("web fetch is unavailable"))
}
urlStr, ok := args["url"].(string)
if !ok || strings.TrimSpace(urlStr) == "" {
return agent.ToolResult{}, fmt.Errorf("url parameter is required")
}
parsed, err := url.Parse(urlStr)
if err != nil {
return agent.ToolResult{}, fmt.Errorf("invalid URL: %w", err)
}
if scheme := strings.ToLower(parsed.Scheme); scheme != "http" && scheme != "https" {
return agent.ToolResult{}, fmt.Errorf("unsupported URL scheme %q: only http and https are allowed", parsed.Scheme)
}
client, err := api.ClientFromEnvironment()
if err != nil {
return agent.ToolResult{}, err
}
ctx, cancel := context.WithTimeout(ctx, webFetchTimeout)
defer cancel()
fetchResp, err := client.WebFetchExperimental(ctx, &api.WebFetchRequest{URL: urlStr})
if err != nil {
var authErr api.AuthorizationError
if errors.As(err, &authErr) {
return agent.ToolResult{}, ErrWebAuthRequired
}
return agent.ToolResult{}, err
}
var sb strings.Builder
if fetchResp.Title != "" {
sb.WriteString(fmt.Sprintf("Title: %s\n\n", fetchResp.Title))
}
if fetchResp.Content != "" {
sb.WriteString("Content:\n")
sb.WriteString(truncateWebFetchContent(fetchResp.Content))
} else {
sb.WriteString("No content could be extracted from the page.")
}
return agent.ToolResult{Content: sb.String()}, nil
}
func truncateWebFetchContent(content string) string {
return agent.Truncate(content, agent.TruncateConfig{
MaxRunes: maxWebFetchContentRunes,
Label: "tool output",
Hint: "Use a narrower request or search query if more detail is needed.",
})
}
+214
View File
@@ -0,0 +1,214 @@
package tools
import (
"encoding/json"
"errors"
"net/http"
"net/http/httptest"
"os"
"path/filepath"
"strings"
"testing"
coreagent "github.com/ollama/ollama/agent"
"github.com/ollama/ollama/api"
"github.com/ollama/ollama/envconfig"
internalcloud "github.com/ollama/ollama/internal/cloud"
)
func TestWebToolsRequireApproval(t *testing.T) {
if !coreagent.ToolRequiresApproval((&WebSearch{}), map[string]any{"query": "ollama"}) {
t.Fatal("web search should require approval")
}
if !coreagent.ToolRequiresApproval((&WebFetch{}), map[string]any{"url": "https://ollama.com"}) {
t.Fatal("web fetch should require approval")
}
}
var webToolCases = []struct {
name string
tool coreagent.Tool
args map[string]any
path string
operation string
}{
{"search", &WebSearch{}, map[string]any{"query": "ollama"}, "/api/experimental/web_search", "web search is unavailable"},
{"fetch", &WebFetch{}, map[string]any{"url": "https://ollama.com"}, "/api/experimental/web_fetch", "web fetch is unavailable"},
}
// enableWebToolsForTest isolates web tool tests from the runner's cloud
// policy. In particular, Windows can inherit both OLLAMA_NO_CLOUD and a
// server.json from USERPROFILE.
func enableWebToolsForTest(t *testing.T) {
t.Helper()
// Register before t.Setenv so the cache is refreshed after t.Setenv has
// restored the runner's environment during cleanup.
t.Cleanup(envconfig.ReloadServerConfig)
home := t.TempDir()
t.Setenv("HOME", home)
t.Setenv("USERPROFILE", home)
t.Setenv("OLLAMA_NO_CLOUD", "")
envconfig.ReloadServerConfig()
}
// runWebTool executes tool against a stub server that responds to every
// request with status and body, returning the resulting error.
func runWebTool(t *testing.T, tool coreagent.Tool, args map[string]any, path string, status int, body string) error {
t.Helper()
ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.URL.Path != path {
t.Fatalf("path = %q, want %q", r.URL.Path, path)
}
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(status)
_, _ = w.Write([]byte(body))
}))
t.Cleanup(ts.Close)
t.Setenv("OLLAMA_HOST", ts.URL)
_, err := tool.Execute(t.Context(), coreagent.ToolContext{}, args)
return err
}
func TestWebToolsReportAuthenticationError(t *testing.T) {
enableWebToolsForTest(t)
for _, tt := range webToolCases {
t.Run(tt.name, func(t *testing.T) {
err := runWebTool(t, tt.tool, tt.args, tt.path, http.StatusUnauthorized,
`{"error":"unauthorized","signin_url":"https://ollama.com/signin"}`)
if !errors.Is(err, ErrWebAuthRequired) {
t.Fatalf("error = %v, want %v", err, ErrWebAuthRequired)
}
})
}
}
func TestWebToolsPreserveNonAuthenticationErrors(t *testing.T) {
enableWebToolsForTest(t)
for _, tt := range webToolCases {
t.Run(tt.name, func(t *testing.T) {
err := runWebTool(t, tt.tool, tt.args, tt.path, http.StatusTooManyRequests,
`{"error":"web search quota exceeded"}`)
if err == nil {
t.Fatal("expected error")
}
if !strings.Contains(err.Error(), "web search quota exceeded") {
t.Fatalf("error = %q, want original error message", err)
}
})
}
}
func TestWebToolsIgnoreInheritedCloudPolicy(t *testing.T) {
// This cleanup is registered before the test environment, so it restores
// the server config cache after t.Setenv restores the runner's values.
t.Cleanup(envconfig.ReloadServerConfig)
home := t.TempDir()
configPath := filepath.Join(home, ".ollama", "server.json")
if err := os.MkdirAll(filepath.Dir(configPath), 0o755); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(configPath, []byte(`{"disable_ollama_cloud":true}`), 0o644); err != nil {
t.Fatal(err)
}
t.Setenv("HOME", home)
t.Setenv("USERPROFILE", home)
t.Setenv("OLLAMA_NO_CLOUD", "1")
envconfig.ReloadServerConfig()
enableWebToolsForTest(t)
err := runWebTool(t, &WebSearch{}, map[string]any{"query": "ollama"}, "/api/experimental/web_search", http.StatusUnauthorized,
`{"error":"unauthorized","signin_url":"https://ollama.com/signin"}`)
if !errors.Is(err, ErrWebAuthRequired) {
t.Fatalf("error = %v, want %v", err, ErrWebAuthRequired)
}
}
func TestWebFetchRejectsUnsupportedScheme(t *testing.T) {
enableWebToolsForTest(t)
tests := []struct {
name string
url string
wantErr bool
}{
{name: "file scheme", url: "file:///etc/passwd", wantErr: true},
{name: "data scheme", url: "data:text/plain,secret", wantErr: true},
{name: "ftp scheme", url: "ftp://example.com/secret", wantErr: true},
{name: "http allowed", url: "http://example.com", wantErr: false},
{name: "https allowed", url: "https://example.com", wantErr: false},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
_, err := (&WebFetch{}).Execute(t.Context(), coreagent.ToolContext{}, map[string]any{"url": tt.url})
if tt.wantErr && err == nil {
t.Fatal("expected unsupported scheme to be rejected")
}
// For allowed schemes we expect an error only from the missing
// server/auth path, not from scheme validation. The http/https
// cases reach the client and may fail on connection/auth; we only
// assert that the error is NOT a scheme error.
if !tt.wantErr && err != nil && strings.Contains(err.Error(), "unsupported URL scheme") {
t.Fatalf("http/https rejected as unsupported: %v", err)
}
})
}
}
func TestWebFetchBoundsContentBeforeReturning(t *testing.T) {
enableWebToolsForTest(t)
ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.URL.Path != "/api/experimental/web_fetch" {
t.Fatalf("path = %q, want /api/experimental/web_fetch", r.URL.Path)
}
var req api.WebFetchRequest
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
t.Fatal(err)
}
if req.URL != "https://ollama.com" {
t.Fatalf("request URL = %q, want https://ollama.com", req.URL)
}
if err := json.NewEncoder(w).Encode(api.WebFetchResponse{
Title: "Ollama",
Content: strings.Repeat("x", maxWebFetchContentRunes+25),
}); err != nil {
t.Fatal(err)
}
}))
defer ts.Close()
t.Setenv("OLLAMA_HOST", ts.URL)
result, err := (&WebFetch{}).Execute(t.Context(), coreagent.ToolContext{}, map[string]any{
"url": "https://ollama.com",
})
if err != nil {
t.Fatal(err)
}
if !strings.Contains(result.Content, "[tool output truncated: showing first ~") ||
!strings.Contains(result.Content, "omitted ~7 tokens") ||
!strings.Contains(result.Content, "Use a narrower request or search query") {
t.Fatalf("content missing truncation marker: %q", result.Content)
}
if count := strings.Count(result.Content, "x"); count != maxWebFetchContentRunes {
t.Fatalf("captured content count = %d, want %d", count, maxWebFetchContentRunes)
}
}
func TestWebToolsRejectWhenCloudDisabled(t *testing.T) {
t.Setenv("OLLAMA_NO_CLOUD", "1")
for _, tt := range webToolCases {
t.Run(tt.name, func(t *testing.T) {
_, err := tt.tool.Execute(t.Context(), coreagent.ToolContext{}, tt.args)
want := internalcloud.DisabledError(tt.operation)
if err == nil || err.Error() != want {
t.Fatalf("error = %v, want %q", err, want)
}
})
}
}
+12
View File
@@ -777,6 +777,18 @@ func (c *StreamConverter) Process(r api.ChatResponse) []StreamEvent {
}
if r.Message.Thinking != "" && !c.thinkingDone {
if c.textStarted {
events = append(events, StreamEvent{
Event: "content_block_stop",
Data: ContentBlockStopEvent{
Type: "content_block_stop",
Index: c.contentIndex,
},
})
c.contentIndex++
c.textStarted = false
}
if !c.thinkingStarted {
c.thinkingStarted = true
events = append(events, StreamEvent{
+51
View File
@@ -3,6 +3,7 @@ package anthropic
import (
"encoding/base64"
"encoding/json"
"fmt"
"strings"
"testing"
@@ -1140,6 +1141,56 @@ func TestStreamConverter_ThinkingDirectlyFollowedByToolCall(t *testing.T) {
}
}
func TestStreamConverter_TextBeforeThinking(t *testing.T) {
conv := NewStreamConverter("msg_123", "test-model", 0)
responses := []api.ChatResponse{
{Message: api.Message{Role: "assistant", Content: "---\n"}},
{Message: api.Message{Role: "assistant", Thinking: "Let me think."}},
{
Message: api.Message{Role: "assistant", Content: "The answer."},
Done: true,
DoneReason: "stop",
Metrics: api.Metrics{PromptEvalCount: 10, EvalCount: 5},
},
}
var got []string
for _, response := range responses {
for _, event := range conv.Process(response) {
switch data := event.Data.(type) {
case ContentBlockStartEvent:
got = append(got, fmt.Sprintf("%s:%s:%d", event.Event, data.ContentBlock.Type, data.Index))
case ContentBlockDeltaEvent:
got = append(got, fmt.Sprintf("%s:%s:%d", event.Event, data.Delta.Type, data.Index))
case ContentBlockStopEvent:
got = append(got, fmt.Sprintf("%s:%d", event.Event, data.Index))
default:
got = append(got, event.Event)
}
}
}
want := []string{
"message_start",
"content_block_start:text:0",
"content_block_delta:text_delta:0",
"content_block_stop:0",
"content_block_start:thinking:1",
"content_block_delta:thinking_delta:1",
"content_block_stop:1",
"content_block_start:text:2",
"content_block_delta:text_delta:2",
"content_block_stop:2",
"message_delta",
"message_stop",
}
if diff := cmp.Diff(want, got); diff != "" {
t.Fatalf("unexpected stream events (-want +got):\n%s", diff)
}
}
func TestStreamConverter_ToolCallWithUnmarshalableArgs(t *testing.T) {
// Test that unmarshalable arguments (like channels) are handled gracefully
// and don't cause a panic or corrupt stream
+30
View File
@@ -473,6 +473,26 @@ func (c *Client) CloudStatusExperimental(ctx context.Context) (*StatusResponse,
return &status, nil
}
// WebSearchExperimental searches the web through the local server's
// experimental web search endpoint.
func (c *Client) WebSearchExperimental(ctx context.Context, req *WebSearchRequest) (*WebSearchResponse, error) {
var resp WebSearchResponse
if err := c.do(ctx, http.MethodPost, "/api/experimental/web_search", req, &resp); err != nil {
return nil, err
}
return &resp, nil
}
// WebFetchExperimental fetches web page content through the local server's
// experimental web fetch endpoint.
func (c *Client) WebFetchExperimental(ctx context.Context, req *WebFetchRequest) (*WebFetchResponse, error) {
var resp WebFetchResponse
if err := c.do(ctx, http.MethodPost, "/api/experimental/web_fetch", req, &resp); err != nil {
return nil, err
}
return &resp, nil
}
// Signout will signout a client for a local ollama server.
func (c *Client) Signout(ctx context.Context) error {
return c.do(ctx, http.MethodPost, "/api/signout", nil, nil)
@@ -490,3 +510,13 @@ func (c *Client) Whoami(ctx context.Context) (*UserResponse, error) {
}
return &resp, nil
}
// Usage returns the authenticated user's recent activity and included-usage
// limits.
func (c *Client) Usage(ctx context.Context) (*UsageResponse, error) {
var resp UsageResponse
if err := c.do(ctx, http.MethodGet, "/api/usage", nil, &resp); err != nil {
return nil, err
}
return &resp, nil
}
+102
View File
@@ -51,6 +51,32 @@ func TestClientFromEnvironment(t *testing.T) {
}
}
func TestClientUsage(t *testing.T) {
ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodGet || r.URL.Path != "/api/usage" {
t.Fatalf("request = %s %s, want GET /api/usage", r.Method, r.URL.Path)
}
fmt.Fprint(w, `{"activity":{"cost":"0.00709","period":{"type":"last_4_weeks","starting_at":"2026-06-29T00:00:00Z","ending_at":"2026-07-27T00:00:00Z"},"models":[{"name":"qwen3-coder:480b","request_count":1,"cost":"0.00709"}]},"limits":{"session":{"usage":0.006,"models":[]},"weekly":{"usage":0,"models":[]}}}`)
}))
defer ts.Close()
base, err := url.Parse(ts.URL)
if err != nil {
t.Fatal(err)
}
got, err := NewClient(base, ts.Client()).Usage(t.Context())
if err != nil {
t.Fatal(err)
}
if got.Activity.Cost != "0.00709" {
t.Errorf("activity cost = %q, want 0.00709", got.Activity.Cost)
}
if len(got.Activity.Models) != 1 || got.Activity.Models[0].Name != "qwen3-coder:480b" {
t.Errorf("activity models = %#v, want qwen3-coder:480b", got.Activity.Models)
}
}
// testError represents an internal error type with status code and message
// this is used since the error response from the server is not a standard error struct
type testError struct {
@@ -351,6 +377,82 @@ func TestClientDo(t *testing.T) {
}
}
func TestClientWebSearchExperimentalUsesLocalRoute(t *testing.T) {
var gotPath string
var gotMethod string
var gotRequest WebSearchRequest
ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
gotPath = r.URL.Path
gotMethod = r.Method
if err := json.NewDecoder(r.Body).Decode(&gotRequest); err != nil {
t.Fatal(err)
}
if err := json.NewEncoder(w).Encode(WebSearchResponse{
Results: []WebSearchResult{{Title: "Ollama", URL: "https://ollama.com", Content: "models"}},
}); err != nil {
t.Fatal(err)
}
}))
defer ts.Close()
client := NewClient(&url.URL{Scheme: "http", Host: ts.Listener.Addr().String()}, http.DefaultClient)
resp, err := client.WebSearchExperimental(t.Context(), &WebSearchRequest{Query: "ollama", MaxResults: 3})
if err != nil {
t.Fatal(err)
}
if gotMethod != http.MethodPost {
t.Fatalf("method = %q, want POST", gotMethod)
}
if gotPath != "/api/experimental/web_search" {
t.Fatalf("path = %q, want /api/experimental/web_search", gotPath)
}
if gotRequest.Query != "ollama" || gotRequest.MaxResults != 3 {
t.Fatalf("request = %#v", gotRequest)
}
if len(resp.Results) != 1 || resp.Results[0].Title != "Ollama" {
t.Fatalf("response = %#v", resp)
}
}
func TestClientWebFetchExperimentalUsesLocalRoute(t *testing.T) {
var gotPath string
var gotMethod string
var gotRequest WebFetchRequest
ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
gotPath = r.URL.Path
gotMethod = r.Method
if err := json.NewDecoder(r.Body).Decode(&gotRequest); err != nil {
t.Fatal(err)
}
if err := json.NewEncoder(w).Encode(WebFetchResponse{
Title: "Ollama",
Content: "models",
Links: []string{"https://ollama.com/library"},
}); err != nil {
t.Fatal(err)
}
}))
defer ts.Close()
client := NewClient(&url.URL{Scheme: "http", Host: ts.Listener.Addr().String()}, http.DefaultClient)
resp, err := client.WebFetchExperimental(t.Context(), &WebFetchRequest{URL: "https://ollama.com"})
if err != nil {
t.Fatal(err)
}
if gotMethod != http.MethodPost {
t.Fatalf("method = %q, want POST", gotMethod)
}
if gotPath != "/api/experimental/web_fetch" {
t.Fatalf("path = %q, want /api/experimental/web_fetch", gotPath)
}
if gotRequest.URL != "https://ollama.com" {
t.Fatalf("request = %#v", gotRequest)
}
if resp.Title != "Ollama" || resp.Content != "models" {
t.Fatalf("response = %#v", resp)
}
}
type roundTripFunc func(*http.Request) (*http.Response, error)
func (f roundTripFunc) RoundTrip(req *http.Request) (*http.Response, error) {
+69
View File
@@ -868,6 +868,36 @@ type StatusResponse struct {
Cloud CloudStatus `json:"cloud"`
}
// WebSearchRequest is the request for [Client.WebSearchExperimental].
type WebSearchRequest struct {
Query string `json:"query"`
MaxResults int `json:"max_results,omitempty"`
}
// WebSearchResult is a single result from [Client.WebSearchExperimental].
type WebSearchResult struct {
Title string `json:"title"`
URL string `json:"url"`
Content string `json:"content"`
}
// WebSearchResponse is the response from [Client.WebSearchExperimental].
type WebSearchResponse struct {
Results []WebSearchResult `json:"results"`
}
// WebFetchRequest is the request for [Client.WebFetchExperimental].
type WebFetchRequest struct {
URL string `json:"url"`
}
// WebFetchResponse is the response from [Client.WebFetchExperimental].
type WebFetchResponse struct {
Title string `json:"title"`
Content string `json:"content"`
Links []string `json:"links,omitempty"`
}
// GenerateResponse is the response passed into [GenerateResponseFunc].
type GenerateResponse struct {
// Model is the model name that generated the response.
@@ -948,6 +978,45 @@ type UserResponse struct {
Plan string `json:"plan,omitempty"`
}
// UsageResponse reports recent activity and included-usage limits.
type UsageResponse struct {
Activity UsageActivity `json:"activity"`
Limits UsageLimits `json:"limits"`
}
// UsageActivity reports usage activity over a period.
type UsageActivity struct {
Cost string `json:"cost"`
Period UsagePeriod `json:"period"`
Models []UsageModel `json:"models"`
}
// UsagePeriod describes the time window the usage covers.
type UsagePeriod struct {
Type string `json:"type"`
StartingAt time.Time `json:"starting_at"`
EndingAt time.Time `json:"ending_at"`
}
// UsageLimits reports included usage for the current session and week.
type UsageLimits struct {
Session UsageLimit `json:"session"`
Weekly UsageLimit `json:"weekly"`
}
// UsageLimit reports the consumed fraction of an included-usage limit.
type UsageLimit struct {
Usage float64 `json:"usage"`
Models []UsageModel `json:"models"`
}
// UsageModel reports a model's activity.
type UsageModel struct {
Name string `json:"name"`
RequestCount int `json:"request_count"`
Cost string `json:"cost,omitempty"`
}
// Tensor describes the metadata for a given tensor.
type Tensor struct {
Name string `json:"name"`
+4 -4
View File
@@ -22,10 +22,10 @@ const LAUNCH_COMMANDS: LaunchCommand[] = [
iconClassName: "h-7 w-7",
},
{
id: "codex-app",
name: "Codex App",
command: "ollama launch codex-app",
description: "An AI agent you can delegate real work to, by OpenAI",
id: "chatgpt",
name: "ChatGPT",
command: "ollama launch chatgpt",
description: "Complete work with ChatGPT",
icon: "/launch-icons/codex-app.png",
iconClassName: "h-full w-full",
},
+96 -6
View File
@@ -59,7 +59,7 @@ function(ollama_macos_major_version output)
RESULT_VARIABLE _macos_result
ERROR_QUIET)
if(_macos_result EQUAL 0)
string(REGEX MATCH "^[0-9]+" _macos_major "${_macos_version}")
string(REGEX MATCH "^[0-9]+(\\.[0-9]+)?" _macos_major "${_macos_version}")
endif()
set(${output} "${_macos_major}" PARENT_SCOPE)
endfunction()
@@ -72,7 +72,7 @@ function(ollama_macos_sdk_major_version output)
RESULT_VARIABLE _sdk_result
ERROR_QUIET)
if(_sdk_result EQUAL 0)
string(REGEX MATCH "^[0-9]+" _sdk_major "${_sdk_version}")
string(REGEX MATCH "^[0-9]+(\\.[0-9]+)?" _sdk_major "${_sdk_version}")
endif()
set(${output} "${_sdk_major}" PARENT_SCOPE)
endfunction()
@@ -83,7 +83,9 @@ function(ollama_default_mlx_backends output)
ollama_check_metal_toolchain(_metal_version)
ollama_macos_major_version(_macos_major)
ollama_macos_sdk_major_version(_sdk_major)
if(_macos_major AND _sdk_major AND _macos_major GREATER_EQUAL 26 AND _sdk_major GREATER_EQUAL 26)
if(_macos_major AND _sdk_major
AND _macos_major VERSION_GREATER_EQUAL 26.2
AND _sdk_major VERSION_GREATER_EQUAL 26.2)
set(_backends "metal_v4")
else()
set(_backends "metal_v3")
@@ -192,8 +194,16 @@ if(OLLAMA_MLX_BACKENDS)
add_custom_target(ollama-mlx-sources DEPENDS ${_mlx_source_targets})
endif()
set(OLLAMA_BUILD_PARALLEL "" CACHE STRING
"Number of parallel jobs for nested native builds (empty = use generator default)")
set(_native_parallel_args --parallel)
if(NOT OLLAMA_BUILD_PARALLEL STREQUAL "")
list(APPEND _native_parallel_args ${OLLAMA_BUILD_PARALLEL})
endif()
set(OLLAMA_NATIVE_BUILD_TOOL_COMMAND
${CMAKE_COMMAND} --build <BINARY_DIR>)
${CMAKE_COMMAND} --build <BINARY_DIR> ${_native_parallel_args})
set(OLLAMA_NATIVE_BUILD_TARGET_ARG --target)
if(CMAKE_GENERATOR MATCHES "Makefiles")
set(OLLAMA_NATIVE_BUILD_TOOL_COMMAND
@@ -236,6 +246,67 @@ function(ollama_cache_arg_is_set name output)
endif()
endfunction()
function(ollama_backend_cuda_major backend output)
if("${backend}" MATCHES "^cuda_v([0-9]+)$")
set(${output} "${CMAKE_MATCH_1}" PARENT_SCOPE)
else()
set(${output} "" PARENT_SCOPE)
endif()
endfunction()
function(ollama_find_windows_cuda_root major output)
if(NOT WIN32 OR "${major}" STREQUAL "")
set(${output} "" PARENT_SCOPE)
return()
endif()
execute_process(
COMMAND ${CMAKE_COMMAND} -E environment
OUTPUT_VARIABLE _environment)
string(REPLACE "\r\n" "\n" _environment "${_environment}")
string(REPLACE "\r" "\n" _environment "${_environment}")
string(REGEX MATCHALL "CUDA_PATH_V${major}_[0-9]+=[^\n]*" _matches "${_environment}")
set(_best_minor -1)
set(_best_root "")
foreach(_entry IN LISTS _matches)
if(_entry MATCHES "^CUDA_PATH_V${major}_([0-9]+)=(.*)$")
set(_minor "${CMAKE_MATCH_1}")
set(_root "${CMAKE_MATCH_2}")
if(_minor GREATER _best_minor)
set(_best_minor ${_minor})
set(_best_root "${_root}")
endif()
endif()
endforeach()
if(_best_root STREQUAL "" AND DEFINED ENV{CUDA_PATH})
set(_cuda_path "$ENV{CUDA_PATH}")
if(EXISTS "${_cuda_path}/version.json")
file(READ "${_cuda_path}/version.json" _version_json)
if(_version_json MATCHES "\"cuda\"[ \t\r\n]*:[ \t\r\n]*\"${major}\\.")
set(_best_root "${_cuda_path}")
endif()
endif()
endif()
set(${output} "${_best_root}" PARENT_SCOPE)
endfunction()
function(ollama_append_cuda_toolkit_args output backend)
# If CUDAToolkit_ROOT is already explicitly set, just forward it.
ollama_append_cache_arg_if_set(${output} CUDAToolkit_ROOT)
if(NOT DEFINED CUDAToolkit_ROOT OR "${CUDAToolkit_ROOT}" STREQUAL "")
# Auto-discover CUDA toolkit for the requested backend version on Windows.
ollama_backend_cuda_major("${backend}" _cuda_major)
ollama_find_windows_cuda_root("${_cuda_major}" _cuda_root)
if(NOT "${_cuda_root}" STREQUAL "")
ollama_escape_cmake_list("${_cuda_root}" _value)
set(${output} ${${output}} "-DCUDAToolkit_ROOT=${_value}" PARENT_SCOPE)
endif()
endif()
endfunction()
function(ollama_llama_cuda_preset backend output)
ollama_cache_arg_is_set(CMAKE_CUDA_ARCHITECTURES _has_cuda_arch)
if(_has_cuda_arch)
@@ -327,12 +398,28 @@ function(ollama_add_llama_server_build name)
-DCMAKE_OSX_DEPLOYMENT_TARGET=${CMAKE_OSX_DEPLOYMENT_TARGET})
endif()
endif()
# Visual Studio requires -T toolset override to select the correct CUDA toolkit.
# MSBuild's CUDA integration ignores -DCUDAToolkit_ROOT for nvcc selection.
# Prefer user-specified CUDAToolkit_ROOT before falling back to auto-discovery.
set(_generator_args)
if(WIN32 AND CMAKE_GENERATOR MATCHES "Visual Studio")
set(_cuda_root "${CUDAToolkit_ROOT}")
if("${_cuda_root}" STREQUAL "")
ollama_backend_cuda_major("${name}" _cuda_major)
ollama_find_windows_cuda_root("${_cuda_major}" _cuda_root)
endif()
if(NOT "${_cuda_root}" STREQUAL "")
list(APPEND _generator_args -T cuda=${_cuda_root})
endif()
endif()
set(_configure_command ${CMAKE_COMMAND}
${_generator_args}
-S ${CMAKE_SOURCE_DIR}/llama/server
-B <BINARY_DIR>
${_cmake_args})
if(ARG_PRESET)
set(_configure_command ${CMAKE_COMMAND}
${_generator_args}
-S ${CMAKE_SOURCE_DIR}/llama/server
--preset ${ARG_PRESET}
-B <BINARY_DIR>
@@ -544,6 +631,7 @@ if(OLLAMA_HAVE_LLAMA_SERVER)
set(_cuda_args)
ollama_append_cache_arg_if_set(_cuda_args CMAKE_CUDA_ARCHITECTURES)
ollama_append_cache_arg_if_set(_cuda_args CMAKE_CUDA_FLAGS)
ollama_append_cuda_toolkit_args(_cuda_args ${_backend})
ollama_add_llama_server_build(${_backend}
PRESET ${_cuda_preset}
RUNNER_DIR ${_backend}
@@ -555,6 +643,7 @@ if(OLLAMA_HAVE_LLAMA_SERVER)
set(_cuda_args)
ollama_append_cache_arg_if_set(_cuda_args CMAKE_CUDA_ARCHITECTURES)
ollama_append_cache_arg_if_set(_cuda_args CMAKE_CUDA_FLAGS)
ollama_append_cuda_toolkit_args(_cuda_args ${_backend})
ollama_add_llama_server_build(${_backend}
PRESET ${_cuda_preset}
RUNNER_DIR ${_backend}
@@ -664,14 +753,15 @@ foreach(_backend IN LISTS OLLAMA_MLX_BACKENDS)
endif()
ollama_check_metal_toolchain(_metal_version)
ollama_macos_sdk_major_version(_ollama_mlx_sdk_major)
if(_ollama_mlx_sdk_major AND _ollama_mlx_sdk_major GREATER_EQUAL 26)
if(_ollama_mlx_sdk_major
AND _ollama_mlx_sdk_major VERSION_GREATER_EQUAL 26.2)
ollama_add_mlx_build(metal_v4
PRESET mlx_metal_v4
RUNNER_DIR mlx_metal_v4)
list(APPEND _mlx_targets ollama-mlx-metal_v4)
else()
message(FATAL_ERROR
"OLLAMA_MLX_BACKENDS=metal_v4 requires the macOS 26 SDK. "
"OLLAMA_MLX_BACKENDS=metal_v4 requires the macOS 26.2 SDK. "
"Install a newer Xcode or use OLLAMA_MLX_BACKENDS=metal_v3.")
endif()
else()
+114 -31
View File
@@ -102,7 +102,8 @@ install(RUNTIME_DEPENDENCY_SET mlx_runtime_deps
LIBRARY DESTINATION ${OLLAMA_INSTALL_DIR} COMPONENT MLX_VENDOR
)
if(TARGET jaccl)
get_target_property(_MLX_LINK_LIBRARIES mlx LINK_LIBRARIES)
if(TARGET jaccl AND "jaccl" IN_LIST _MLX_LINK_LIBRARIES)
install(TARGETS jaccl
RUNTIME DESTINATION ${OLLAMA_INSTALL_DIR} COMPONENT MLX
LIBRARY DESTINATION ${OLLAMA_INSTALL_DIR} COMPONENT MLX
@@ -123,29 +124,53 @@ endif()
# --component MLX. Headers are installed alongside libmlx in OLLAMA_INSTALL_DIR.
#
# Layout:
# ${OLLAMA_INSTALL_DIR}/include/cccl/{cuda,nv}/ - CCCL headers
# ${OLLAMA_INSTALL_DIR}/include/*.h - CUDA toolkit headers
# ${OLLAMA_INSTALL_DIR}/include/cccl/ - CCCL headers
# ${OLLAMA_INSTALL_DIR}/include/{cute,cutlass}/ - CUTLASS/CUTE headers
# ${OLLAMA_INSTALL_DIR}/include/ - CUDA runtime/core headers
#
# MLX's jit_module.cpp resolves CCCL via
# current_binary_dir()[.parent_path()] / "include" / "cccl"
# On Linux, MLX's jit_module.cpp resolves CCCL via
# current_binary_dir().parent_path() / "include" / "cccl", so we create a
# symlink from lib/ollama/include -> ${OLLAMA_RUNNER_DIR}/include.
# MLX's jit_module.cpp resolves JIT support headers from the backend-local
# include directory. On Linux it also probes current_binary_dir().parent_path()
# / "include", so we create a symlink from lib/ollama/include to the backend
# include directory for archive packaging.
# This will need refinement if we add multiple CUDA versions for MLX in the future.
# CUDA runtime headers are found via CUDA_PATH env var (set by mlxrunner).
if(EXISTS ${CMAKE_BINARY_DIR}/_deps/cccl-src/include/cuda)
install(DIRECTORY ${CMAKE_BINARY_DIR}/_deps/cccl-src/include/cuda
DESTINATION ${OLLAMA_INSTALL_DIR}/include/cccl
set(_mlx_jit_cccl_include_dir "")
if(CUDAToolkit_FOUND)
foreach(_dir ${CUDAToolkit_INCLUDE_DIRS})
if(EXISTS "${_dir}/cccl/cuda/std")
set(_mlx_jit_cccl_include_dir "${_dir}/cccl")
break()
endif()
endforeach()
endif()
if(NOT _mlx_jit_cccl_include_dir AND EXISTS ${CMAKE_BINARY_DIR}/_deps/cccl-src/include/cuda)
set(_mlx_jit_cccl_include_dir "${CMAKE_BINARY_DIR}/_deps/cccl-src/include")
endif()
if(_mlx_jit_cccl_include_dir)
foreach(_cccl_dir cuda nv cub thrust)
if(EXISTS "${_mlx_jit_cccl_include_dir}/${_cccl_dir}")
install(DIRECTORY "${_mlx_jit_cccl_include_dir}/${_cccl_dir}"
DESTINATION ${OLLAMA_INSTALL_DIR}/include/cccl
COMPONENT MLX)
endif()
endforeach()
endif()
if(EXISTS ${CMAKE_BINARY_DIR}/_deps/cutlass-src/include/cute)
install(DIRECTORY ${CMAKE_BINARY_DIR}/_deps/cutlass-src/include/cute
DESTINATION ${OLLAMA_INSTALL_DIR}/include
COMPONENT MLX)
install(DIRECTORY ${CMAKE_BINARY_DIR}/_deps/cccl-src/include/nv
DESTINATION ${OLLAMA_INSTALL_DIR}/include/cccl
install(DIRECTORY ${CMAKE_BINARY_DIR}/_deps/cutlass-src/include/cutlass
DESTINATION ${OLLAMA_INSTALL_DIR}/include
COMPONENT MLX)
endif()
# Install minimal CUDA toolkit headers needed by MLX JIT kernels.
# These are the transitive closure of includes from mlx/backend/cuda/device/*.cuh.
# Install CUDA runtime/core headers needed by MLX JIT kernels.
# NVIDIA's NVRTC bundled-header model is CUDA Runtime + CCCL, not the entire
# toolkit include tree. Keep CCCL coherent above, include CUTLASS/CUTE above,
# and avoid shipping unrelated SDK headers such as NPP, CUPTI, cuRAND, NVML,
# cuBLAS, cuSPARSE, and cuSOLVER.
# The Go mlxrunner sets CUDA_PATH to OLLAMA_INSTALL_DIR so MLX finds them at
# $CUDA_PATH/include/*.h via NVRTC --include-path.
# $CUDA_PATH/include via NVRTC --include-path.
if(CUDAToolkit_FOUND)
# CUDAToolkit_INCLUDE_DIRS may be a semicolon-separated list
# (e.g. ".../include;.../include/cccl"). Find the entry that
@@ -161,39 +186,97 @@ if(CUDAToolkit_FOUND)
message(WARNING "Could not find cuda_runtime_api.h in CUDAToolkit_INCLUDE_DIRS: ${CUDAToolkit_INCLUDE_DIRS}")
else()
set(_dst "${OLLAMA_INSTALL_DIR}/include")
set(_MLX_JIT_CUDA_HEADERS
set(_mlx_jit_cuda_headers
builtin_types.h
channel_descriptor.h
common_functions.h
cooperative_groups.h
cuComplex.h
cuda.h
cudaTypedefs.h
cuda_awbarrier.h
cuda_awbarrier_helpers.h
cuda_awbarrier_primitives.h
cuda_bf16.h
cuda_bf16.hpp
cuda_device_runtime_api.h
cuda_fp16.h
cuda_fp16.hpp
cuda_fp4.h
cuda_fp4.hpp
cuda_fp6.h
cuda_fp6.hpp
cuda_fp8.h
cuda_fp8.hpp
cuda_fp16.h
cuda_fp16.hpp
cuda_occupancy.h
cuda_pipeline.h
cuda_pipeline_helpers.h
cuda_pipeline_primitives.h
cuda_runtime.h
cuda_runtime_api.h
cuda_stdint.h
cudart_platform.h
device_atomic_functions.h
device_atomic_functions.hpp
device_double_functions.h
device_functions.h
device_launch_parameters.h
device_types.h
driver_functions.h
driver_types.h
fatbinary_section.h
host_config.h
host_defines.h
library_types.h
math_constants.h
math_functions.h
mma.h
nvrtc_device_runtime.h
sm_20_atomic_functions.h
sm_20_atomic_functions.hpp
sm_20_intrinsics.h
sm_20_intrinsics.hpp
sm_30_intrinsics.h
sm_30_intrinsics.hpp
sm_32_atomic_functions.h
sm_32_atomic_functions.hpp
sm_32_intrinsics.h
sm_32_intrinsics.hpp
sm_35_atomic_functions.h
sm_35_intrinsics.h
sm_60_atomic_functions.h
sm_60_atomic_functions.hpp
sm_61_intrinsics.h
sm_61_intrinsics.hpp
surface_indirect_functions.h
surface_types.h
target
texture_indirect_functions.h
texture_types.h
vector_functions.h
vector_functions.hpp
vector_types.h
)
foreach(_hdr ${_MLX_JIT_CUDA_HEADERS})
install(FILES "${_cuda_inc}/${_hdr}"
vector_types.h)
set(_mlx_jit_cuda_header_paths "")
foreach(_header IN LISTS _mlx_jit_cuda_headers)
if(EXISTS "${_cuda_inc}/${_header}")
list(APPEND _mlx_jit_cuda_header_paths "${_cuda_inc}/${_header}")
endif()
endforeach()
if(_mlx_jit_cuda_header_paths)
install(FILES ${_mlx_jit_cuda_header_paths}
DESTINATION ${_dst}
COMPONENT MLX)
endif()
foreach(_runtime_dir cooperative_groups crt)
if(EXISTS "${_cuda_inc}/${_runtime_dir}")
install(DIRECTORY "${_cuda_inc}/${_runtime_dir}"
DESTINATION ${_dst}
COMPONENT MLX)
endif()
endforeach()
# Subdirectory headers.
install(DIRECTORY "${_cuda_inc}/cooperative_groups"
DESTINATION ${_dst}
COMPONENT MLX
FILES_MATCHING PATTERN "*.h")
install(FILES "${_cuda_inc}/crt/host_defines.h"
DESTINATION "${_dst}/crt"
COMPONENT MLX)
if(NOT WIN32 AND NOT APPLE)
install(CODE "
set(_link \"${CMAKE_INSTALL_PREFIX}/${OLLAMA_LIB_DIR}/include\")
+1 -1
View File
@@ -55,7 +55,7 @@
"inherits": [ "default" ],
"binaryDir": "${sourceDir}/../../build/metal-v4",
"cacheVariables": {
"CMAKE_OSX_DEPLOYMENT_TARGET": "26.0",
"CMAKE_OSX_DEPLOYMENT_TARGET": "26.2",
"OLLAMA_RUNNER_DIR": "mlx_metal_v4"
}
}
+863
View File
@@ -0,0 +1,863 @@
package cmd
import (
"context"
"errors"
"fmt"
"net/http"
"os"
"runtime"
"slices"
"strconv"
"strings"
"time"
"github.com/spf13/cobra"
coreagent "github.com/ollama/ollama/agent"
agenttools "github.com/ollama/ollama/agent/tools"
"github.com/ollama/ollama/api"
"github.com/ollama/ollama/cmd/config"
"github.com/ollama/ollama/cmd/launch"
agentchat "github.com/ollama/ollama/cmd/tui/chat"
"github.com/ollama/ollama/format"
internalcloud "github.com/ollama/ollama/internal/cloud"
"github.com/ollama/ollama/internal/modelref"
"github.com/ollama/ollama/types/model"
)
type agentTUIOptions struct {
Model string
OpenModelPicker bool
System string
Format string
Options map[string]any
Think *api.ThinkValue
KeepAlive *api.Duration
ContextWindowTokens int
AllowAllTools bool
ToolsDisabled bool
MultiModal bool
}
func registerAgentFlags(cmd *cobra.Command) {
cmd.Flags().String("model", "", "Model to use")
cmd.Flags().String("keepalive", "", "Duration to keep a model loaded (e.g. 5m)")
cmd.Flags().String("format", "", "Response format (e.g. json)")
cmd.Flags().String("think", "", "Enable thinking mode: true/false or high/medium/low for supported models")
cmd.Flags().Lookup("think").NoOptDefVal = "true"
cmd.Flags().Bool("auto-approve-tools", false, "Allow agent tools to run without prompting")
cmd.Flags().Bool("yolo", false, "Alias for --auto-approve-tools")
cmd.Flags().Bool("no-tools", false, "Disable agent tools")
}
func AgentHandler(cmd *cobra.Command, _ []string) error {
opts := agentTUIOptions{
Model: strings.TrimSpace(config.LastModel()),
Options: map[string]any{},
}
thinkExplicit, err := applyAgentFlags(cmd, &opts)
if err != nil {
return err
}
if strings.TrimSpace(opts.Model) == "" {
opts.OpenModelPicker = true
} else if cmd.Flags().Lookup("model") == nil || !cmd.Flags().Lookup("model").Changed {
opts.OpenModelPicker = true
}
client, err := api.ClientFromEnvironment()
if err != nil {
return err
}
if opts.OpenModelPicker {
modelName, err := selectAgentModel(cmd.Context(), client, opts.Model)
if errors.Is(err, launch.ErrCancelled) {
return nil
}
if err != nil {
return err
}
opts.Model = modelName
opts.OpenModelPicker = false
}
if strings.TrimSpace(opts.Model) != "" {
info, err := prepareAgentModel(cmd, client, &opts, thinkExplicit)
if err != nil {
if handleCloudAuthorizationError(err) {
return nil
}
return err
}
opts.System = info.System
if err := saveLastAgentModel(opts.Model); err != nil {
return err
}
}
if err := GenerateAgentTUI(cmd, client, opts); err != nil {
if handleCloudAuthorizationError(err) {
return nil
}
return fmt.Errorf("error running agent: %w", err)
}
return nil
}
func applyAgentFlags(cmd *cobra.Command, opts *agentTUIOptions) (bool, error) {
if flag := cmd.Flags().Lookup("model"); flag != nil && flag.Changed {
modelName, err := cmd.Flags().GetString("model")
if err != nil {
return false, err
}
modelName = strings.TrimSpace(modelName)
if modelName == "" {
return false, errors.New("--model cannot be empty")
}
opts.Model = modelName
opts.OpenModelPicker = false
}
format, err := cmd.Flags().GetString("format")
if err != nil {
return false, err
}
opts.Format = format
thinkExplicit := false
thinkFlag := cmd.Flags().Lookup("think")
if thinkFlag != nil && thinkFlag.Changed {
thinkExplicit = true
thinkStr, err := cmd.Flags().GetString("think")
if err != nil {
return false, err
}
switch thinkStr {
case "", "true":
opts.Think = &api.ThinkValue{Value: true}
case "false":
opts.Think = &api.ThinkValue{Value: false}
case "high", "medium", "low", "max":
opts.Think = &api.ThinkValue{Value: thinkStr}
default:
return false, fmt.Errorf("invalid value for --think: %q (must be true, false, high, medium, low, or max)", thinkStr)
}
}
keepAlive, err := cmd.Flags().GetString("keepalive")
if err != nil {
return false, err
}
if keepAlive != "" {
d, err := time.ParseDuration(keepAlive)
if err != nil {
return false, err
}
opts.KeepAlive = &api.Duration{Duration: d}
}
autoApprove, err := cmd.Flags().GetBool("auto-approve-tools")
if err != nil {
return false, err
}
yolo, err := cmd.Flags().GetBool("yolo")
if err != nil {
return false, err
}
opts.AllowAllTools = autoApprove || yolo
toolsDisabled, err := cmd.Flags().GetBool("no-tools")
if err != nil {
return false, err
}
opts.ToolsDisabled = toolsDisabled
return thinkExplicit, nil
}
func saveLastAgentModel(model string) error {
model = strings.TrimSpace(model)
if model == "" {
return nil
}
return config.SetLastModel(model)
}
func prepareAgentModel(cmd *cobra.Command, client *api.Client, opts *agentTUIOptions, thinkExplicit bool) (*api.ShowResponse, error) {
requestedCloud := modelref.HasExplicitCloudSource(opts.Model)
info, err := func() (*api.ShowResponse, error) {
info, err := client.Show(cmd.Context(), &api.ShowRequest{Model: opts.Model})
var se api.StatusError
if errors.As(err, &se) && se.StatusCode == http.StatusNotFound {
if requestedCloud {
return nil, err
}
if err := PullHandler(cmd, []string{opts.Model}); err != nil {
return nil, err
}
return client.Show(cmd.Context(), &api.ShowRequest{Model: opts.Model})
}
return info, err
}()
if err != nil {
return nil, err
}
ensureCloudStub(cmd.Context(), client, opts.Model)
opts.Think, err = inferThinkingOption(&info.Capabilities, &runOptions{Model: opts.Model, Think: opts.Think}, thinkExplicit)
if err != nil {
return nil, err
}
opts.MultiModal = showResponseSupportsMultimodal(info)
opts.ContextWindowTokens = showResponseContextWindow(info)
return info, nil
}
func GenerateAgentTUI(cmd *cobra.Command, client *api.Client, opts agentTUIOptions) error {
cwd := agentWorkingDir()
contextWindowForModel := func(ctx context.Context, model string, fallback int) int {
return agentContextWindowForModel(ctx, client, model, fallback)
}
skillCatalog, err := coreagent.LoadDefaultSkills(cwd)
if err != nil {
return fmt.Errorf("load agent skills: %w", err)
}
for _, diagnostic := range skillCatalog.Diagnostics() {
fmt.Fprintf(os.Stderr, "\033[1mwarning:\033[0m ignored invalid agent skill: %v\n", diagnostic)
}
var registry *coreagent.Registry
registryForModel := func(ctx context.Context, model string) *coreagent.Registry {
return agentToolsRegistry(ctx, client, model, skillCatalog)
}
if opts.Model != "" {
registry = agentToolsRegistry(cmd.Context(), client, opts.Model, skillCatalog)
}
systemPrompt := agentSystemPromptWithWorkingDir(opts.Model, opts.System, agentSkillSystemContext(skillCatalog, registry, opts.ToolsDisabled), cwd)
_, err = agentchat.Run(cmd.Context(), agentchat.Options{
Model: opts.Model,
Client: client,
Tools: registry,
ToolRegistryForModel: registryForModel,
ToolsDisabled: opts.ToolsDisabled,
MultiModalForModel: func(ctx context.Context, model string) bool {
return agentModelSupportsMultimodal(ctx, client, model)
},
ModelOptions: func(ctx context.Context) ([]agentchat.ModelOption, error) {
return agentModelOptions(ctx, client)
},
OnModelSelected: func(_ context.Context, model string) error {
return config.SetLastModel(model)
},
SystemPromptForModel: func(ctx context.Context, model string, registry *coreagent.Registry, toolsDisabled bool) string {
return agentSystemPromptWithWorkingDir(model, agentSystemFromShow(ctx, client, model), agentSkillSystemContext(skillCatalog, registry, toolsDisabled), cwd)
},
Skills: skillCatalog,
SystemPrompt: systemPrompt,
WorkingDir: cwd,
Format: opts.Format,
Options: opts.Options,
Think: opts.Think,
KeepAlive: opts.KeepAlive,
MultiModal: opts.MultiModal,
AllowAllTools: opts.AllowAllTools,
ContextWindowTokens: opts.ContextWindowTokens,
Compactor: &coreagent.SimpleCompactor{
Client: client,
Options: coreagent.CompactionOptions{ContextWindowTokens: opts.ContextWindowTokens},
},
ContextWindowTokensForModel: func(ctx context.Context, model string, fallback int) int {
return contextWindowForModel(ctx, model, fallback)
},
PreloadModel: func(ctx context.Context, model string, think *api.ThinkValue) (int, error) {
return preloadAgentModelIfLocal(ctx, client, opts, model, think)
},
CheckCloudModel: func(ctx context.Context, model, requiredPlan string) error {
return ensureCloudModelAccess(ctx, client, model, requiredPlan)
},
OpenBrowser: launch.OpenBrowser,
PollCloudAuth: func(ctx context.Context) (string, bool, error) {
user, err := client.Whoami(ctx)
if err != nil {
return "", false, err
}
if user == nil || user.Name == "" {
return "", false, nil
}
return user.Name, true, nil
},
})
return err
}
func agentSkillSystemContext(catalog *coreagent.SkillCatalog, registry *coreagent.Registry, toolsDisabled bool) string {
if toolsDisabled || registry == nil {
return ""
}
if _, ok := registry.Get("skill"); !ok {
return ""
}
return catalog.SystemContext()
}
func selectAgentModel(ctx context.Context, client *api.Client, current string) (string, error) {
models, err := agentModelOptions(ctx, client)
if err != nil {
return "", err
}
if len(models) == 0 {
return "", errors.New("no models available, run 'ollama pull <model>' first")
}
items := agentSelectionItems(models)
switch {
case launch.DefaultSingleSelectorWithUpdates != nil:
return launch.DefaultSingleSelectorWithUpdates("Select model to run:", items, current, nil)
case launch.DefaultSingleSelector != nil:
return launch.DefaultSingleSelector("Select model to run:", items, current)
default:
return "", errors.New("no selector configured")
}
}
func agentSelectionItems(models []agentchat.ModelOption) []launch.SelectionItem {
items := make([]launch.SelectionItem, 0, len(models))
for _, model := range models {
items = append(items, launch.SelectionItem{
Name: model.Name,
Description: agentSelectionDescription(model),
Recommended: model.Recommended,
AvailabilityBadge: model.AvailabilityBadge,
})
}
return items
}
func agentSelectionDescription(model agentchat.ModelOption) string {
return strings.TrimSpace(model.Description)
}
var agentGetwd = os.Getwd
func agentWorkingDir() string {
cwd, err := agentGetwd()
if err != nil {
return ""
}
return cwd
}
func agentSystemPromptWithWorkingDir(modelName string, modelSystem string, extra string, workingDir string) string {
return agentSystemPromptAtWithWorkingDir(time.Now(), modelName, modelSystem, extra, workingDir)
}
func agentSystemPromptAtWithWorkingDir(now time.Time, modelName string, modelSystem string, extra string, workingDir string) string {
var parts []string
parts = append(parts, agentDefaultSystemPromptWithWorkingDir(now, modelName, workingDir))
if strings.TrimSpace(modelSystem) != "" {
parts = append(parts, strings.TrimSpace(modelSystem))
}
if strings.TrimSpace(extra) != "" {
parts = append(parts, strings.TrimSpace(extra))
}
return strings.Join(parts, "\n\n")
}
func agentDefaultSystemPromptWithWorkingDir(now time.Time, modelName string, workingDir string) string {
date := now.Format("Monday, January 2, 2006")
shellName := "bash"
if runtime.GOOS == "windows" {
shellName = "PowerShell"
}
parts := []string{
"You are running in Ollama, in a harness to help the user accomplish tasks, and the model is " + modelName + ".",
"",
"Current date: " + date + ".",
"",
}
parts = append(parts,
"Be concise, practical, and action-oriented. Use tools when they materially help. Verify current or fast-changing facts with web tools when available; otherwise state uncertainty.",
"",
"Use "+shellName+" carefully. Prefer read-only inspection first. Stay within the current working directory unless explicitly asked. Surface intent before risky actions such as writes, deletes, moves, installs, git state changes, service changes, sudo, secrets access, network scripts, or commands outside the working directory. Request approval when required and do not work around denied approvals.",
"",
"Tell the user about meaningful changes, verification, failures, blockers, assumptions, and risks. Summarize routine tool output instead of dumping it.",
)
if workingDir != "" {
parts = append(parts, "Current working directory: "+strconv.Quote(workingDir)+".")
}
return strings.Join(parts, "\n")
}
func agentSystemFromShow(ctx context.Context, client *api.Client, modelName string) string {
if client == nil || strings.TrimSpace(modelName) == "" {
return ""
}
resp, err := client.Show(ctx, &api.ShowRequest{Model: modelName})
if err != nil {
fmt.Fprintf(os.Stderr, "\033[1mwarning:\033[0m could not load model system prompt: %v\n", err)
return ""
}
return resp.System
}
func agentToolsRegistry(ctx context.Context, client *api.Client, modelName string, skillCatalog *coreagent.SkillCatalog) *coreagent.Registry {
supportsTools, err := agentModelSupportsTools(ctx, client, modelName)
if err != nil {
fmt.Fprintf(os.Stderr, "\033[1mwarning:\033[0m could not check model capabilities: %v\n", err)
}
if !supportsTools {
return nil
}
registry := &coreagent.Registry{}
if os.Getenv("OLLAMA_AGENT_DISABLE_SHELL") == "" {
registry.Register(&agenttools.Bash{})
}
registry.Register(&agenttools.Read{})
registry.Register(&agenttools.Edit{})
if len(skillCatalog.List()) > 0 {
registry.Register(&agenttools.Skill{Catalog: skillCatalog})
}
if os.Getenv("OLLAMA_AGENT_DISABLE_WEBSEARCH") == "" {
if disabled, known := agentCloudStatusDisabled(ctx, client); !known || !disabled {
registry.Register(&agenttools.WebSearch{})
registry.Register(&agenttools.WebFetch{})
} else {
fmt.Fprintf(os.Stderr, "%s\n", internalcloud.DisabledError("web search is unavailable"))
}
}
return registry
}
func agentModelSupportsTools(ctx context.Context, client *api.Client, modelName string) (bool, error) {
if client == nil || strings.TrimSpace(modelName) == "" {
return false, nil
}
resp, err := client.Show(ctx, &api.ShowRequest{Model: modelName})
if err != nil {
return false, err
}
return slices.Contains(resp.Capabilities, model.CapabilityTools), nil
}
func agentModelSupportsMultimodal(ctx context.Context, client *api.Client, modelName string) bool {
if client == nil || strings.TrimSpace(modelName) == "" {
return false
}
resp, err := client.Show(ctx, &api.ShowRequest{Model: modelName})
if err != nil {
fmt.Fprintf(os.Stderr, "\033[1mwarning:\033[0m could not check model capabilities: %v\n", err)
return false
}
return showResponseSupportsMultimodal(resp)
}
func showResponseSupportsMultimodal(resp *api.ShowResponse) bool {
if resp == nil {
return false
}
if slices.Contains(resp.Capabilities, model.CapabilityVision) || slices.Contains(resp.Capabilities, model.CapabilityAudio) {
return true
}
if len(resp.ProjectorInfo) != 0 {
return true
}
for key := range resp.ModelInfo {
if strings.Contains(key, ".vision.") {
return true
}
}
return false
}
func agentContextWindowForModel(ctx context.Context, client *api.Client, modelName string, fallback int) int {
if client == nil || strings.TrimSpace(modelName) == "" {
return fallback
}
if tokens := loadedContextWindowForModel(ctx, client, modelName); tokens > 0 {
return tokens
}
if modelref.HasExplicitCloudSource(modelName) {
if tokens := agentRecommendationContextWindowForModel(ctx, client, modelName); tokens > 0 {
return tokens
}
}
resp, err := client.Show(ctx, &api.ShowRequest{Model: modelName})
if err != nil {
return fallback
}
if tokens := showResponseContextWindow(resp); tokens > 0 {
return tokens
}
return fallback
}
func agentRecommendationContextWindowForModel(ctx context.Context, client *api.Client, modelName string) int {
if client == nil {
return 0
}
recs, err := client.ModelRecommendationsExperimental(ctx)
if err != nil || recs == nil {
return 0
}
return contextWindowFromRecommendations(modelName, recs.Recommendations)
}
func contextWindowFromRecommendations(modelName string, recommendations []api.ModelRecommendation) int {
for _, rec := range recommendations {
if rec.ContextLength <= 0 {
continue
}
if sameModelRef(modelName, rec.Model) {
return rec.ContextLength
}
}
return 0
}
func sameModelRef(a, b string) bool {
a = comparableModelRef(a)
b = comparableModelRef(b)
if strings.EqualFold(a, b) {
return true
}
pa, errA := modelref.ParseRef(a)
pb, errB := modelref.ParseRef(b)
if errA != nil || errB != nil {
return false
}
if !strings.EqualFold(pa.Base, pb.Base) {
return false
}
return pa.Source == pb.Source ||
pa.Source == modelref.ModelSourceUnspecified ||
pb.Source == modelref.ModelSourceUnspecified
}
func comparableModelRef(value string) string {
value = strings.TrimSpace(value)
if strings.HasSuffix(strings.ToLower(value), ":latest") {
return strings.TrimSpace(value[:len(value)-len(":latest")])
}
return value
}
func showResponseContextWindow(resp *api.ShowResponse) int {
if resp == nil {
return 0
}
if resp.Details.ContextLength > 0 {
return resp.Details.ContextLength
}
if n, ok := numericModelInfo(resp.ModelInfo["general.context_length"]); ok {
return n
}
best := 0
for key, value := range resp.ModelInfo {
if key != "context_length" && !strings.HasSuffix(key, ".context_length") {
continue
}
if n, ok := numericModelInfo(value); ok && n > best {
best = n
}
}
return best
}
func numericModelInfo(value any) (int, bool) {
switch v := value.(type) {
case int:
return v, v > 0
case int32:
return int(v), v > 0
case int64:
return int(v), v > 0
case uint:
return int(v), v > 0
case uint32:
return int(v), v > 0
case uint64:
return int(v), v > 0
case float64:
return int(v), v > 0
case string:
n, err := strconv.Atoi(strings.TrimSpace(v))
return n, err == nil && n > 0
default:
return 0, false
}
}
func preloadAgentModelIfLocal(ctx context.Context, client *api.Client, opts agentTUIOptions, modelName string, think *api.ThinkValue) (int, error) {
modelName = strings.TrimSpace(modelName)
if client == nil || modelName == "" {
return 0, nil
}
if modelref.HasExplicitCloudSource(modelName) {
return 0, nil
}
info, err := client.Show(ctx, &api.ShowRequest{Model: modelName})
if err != nil {
return 0, err
}
if info.RemoteHost != "" {
return 0, nil
}
if err := client.Generate(ctx, &api.GenerateRequest{
Model: modelName,
KeepAlive: opts.KeepAlive,
Options: opts.Options,
Think: think,
}, func(api.GenerateResponse) error {
return nil
}); err != nil {
return 0, err
}
return loadedContextWindowForModel(ctx, client, modelName), nil
}
func loadedContextWindowForModel(ctx context.Context, client *api.Client, modelName string) int {
if client == nil || strings.TrimSpace(modelName) == "" {
return 0
}
resp, err := client.ListRunning(ctx)
if err != nil {
return 0
}
return processContextWindowForModel(modelName, resp)
}
func processContextWindowForModel(modelName string, resp *api.ProcessResponse) int {
if resp == nil {
return 0
}
for _, running := range resp.Models {
if running.ContextLength <= 0 {
continue
}
if sameModelRef(modelName, running.Name) || sameModelRef(modelName, running.Model) {
return running.ContextLength
}
}
return 0
}
func agentModelOptions(ctx context.Context, client *api.Client) ([]agentchat.ModelOption, error) {
if client == nil {
return nil, errors.New("model picker requires an API client")
}
list, err := client.List(ctx)
if err != nil {
return nil, err
}
seen := make(map[string]struct{})
var options []agentchat.ModelOption
add := func(name, description string, recommended bool, requiredPlan string, cloud bool) {
name = strings.TrimSpace(name)
if name == "" {
return
}
key := strings.ToLower(name)
if _, ok := seen[key]; ok {
return
}
seen[key] = struct{}{}
options = append(options, agentchat.ModelOption{
Name: name,
Description: strings.TrimSpace(description),
Recommended: recommended,
RequiredPlan: requiredPlan,
Cloud: cloud,
})
}
if disabled, known := agentCloudStatusDisabled(ctx, client); !known || !disabled {
if recs, err := client.ModelRecommendationsExperimental(ctx); err == nil {
for _, rec := range recs.Recommendations {
name := strings.TrimSpace(rec.Model)
if !modelref.HasExplicitCloudSource(name) {
continue
}
add(name, agentRecommendationDescription(rec), true, strings.TrimSpace(rec.RequiredPlan), true)
}
}
}
local := slices.Clone(list.Models)
slices.SortStableFunc(local, func(a, b api.ListModelResponse) int {
return strings.Compare(strings.ToLower(a.Name), strings.ToLower(b.Name))
})
for _, model := range local {
name := strings.TrimSpace(model.Name)
if name == "" {
name = strings.TrimSpace(model.Model)
}
name = strings.TrimSuffix(name, ":latest")
if modelref.HasExplicitCloudSource(name) {
add(name, agentCloudModelDescription(model), false, "", true)
continue
}
add(name, agentLocalModelDescription(model), false, "", false)
}
badges, signInURLs := cloudAvailabilityBadges(ctx, client, options)
for i := range options {
options[i].AvailabilityBadge = badges[options[i].Name]
options[i].SignInURL = signInURLs[options[i].Name]
}
return options, nil
}
func cloudAvailabilityBadges(ctx context.Context, client *api.Client, options []agentchat.ModelOption) (map[string]string, map[string]string) {
badges := make(map[string]string)
signInURLs := make(map[string]string)
hasCloud := false
for _, opt := range options {
if opt.Cloud {
hasCloud = true
break
}
}
if !hasCloud {
return badges, signInURLs
}
if disabled, known := agentCloudStatusDisabled(ctx, client); known && disabled {
return badges, signInURLs
}
whoamiCtx, cancel := context.WithTimeout(ctx, 3*time.Second)
defer cancel()
user, err := client.Whoami(whoamiCtx)
if err != nil {
var authErr api.AuthorizationError
signInURL := ""
if errors.As(err, &authErr) && (authErr.StatusCode == http.StatusUnauthorized || authErr.SigninURL != "") {
if authErr.SigninURL != "" {
signInURL = authErr.SigninURL
}
} else {
return badges, signInURLs
}
for _, opt := range options {
if opt.Cloud {
badges[opt.Name] = "Sign in required"
if signInURL != "" {
signInURLs[opt.Name] = signInURL
}
}
}
return badges, signInURLs
}
signedIn := user != nil && user.Name != ""
for _, opt := range options {
if !opt.Cloud {
continue
}
if !signedIn {
badges[opt.Name] = "Sign in required"
} else if opt.RequiredPlan != "" && !launch.PlanSatisfies(user.Plan, opt.RequiredPlan) {
badges[opt.Name] = "Upgrade required"
}
}
return badges, signInURLs
}
func agentRecommendationDescription(rec api.ModelRecommendation) string {
var parts []string
if description := strings.TrimSpace(rec.Description); description != "" {
parts = append(parts, description)
} else {
parts = append(parts, "cloud")
}
if rec.ContextLength > 0 {
parts = append(parts, format.HumanNumber(uint64(rec.ContextLength))+" ctx")
}
return strings.Join(parts, " - ")
}
func agentLocalModelDescription(model api.ListModelResponse) string {
desc := agentModelArchDescription(model)
if desc == "" {
return "local"
}
return "local - " + desc
}
func agentCloudModelDescription(model api.ListModelResponse) string {
return agentModelArchDescription(model)
}
func agentModelArchDescription(model api.ListModelResponse) string {
var details []string
if model.Details.Family != "" {
details = append(details, model.Details.Family)
}
if ps := humanizedParameterSize(model.Details.ParameterSize); ps != "" {
details = append(details, ps)
}
if model.Details.QuantizationLevel != "" {
details = append(details, model.Details.QuantizationLevel)
}
var parts []string
if len(details) > 0 {
parts = append(parts, strings.Join(details, " "))
}
if model.Details.ContextLength > 0 {
parts = append(parts, format.HumanNumber(uint64(model.Details.ContextLength))+" ctx")
}
return strings.Join(parts, " - ")
}
func humanizedParameterSize(s string) string {
s = strings.TrimSpace(s)
if s == "" {
return ""
}
if f, err := strconv.ParseFloat(s, 64); err == nil {
return format.HumanNumber(uint64(f))
}
return s
}
func agentCloudStatusDisabled(ctx context.Context, client *api.Client) (disabled bool, known bool) {
if internalcloud.Disabled() {
return true, true
}
status, err := client.CloudStatusExperimental(ctx)
if err != nil {
var statusErr api.StatusError
if errors.As(err, &statusErr) && statusErr.StatusCode == http.StatusNotFound {
return false, false
}
return false, false
}
return status.Cloud.Disabled, true
}
func ensureCloudModelAccess(ctx context.Context, client *api.Client, modelName, requiredPlan string) error {
if client == nil {
return errors.New("no API client available")
}
if disabled, known := agentCloudStatusDisabled(ctx, client); known && disabled {
return errors.New("remote inference is unavailable")
}
user, err := client.Whoami(ctx)
if err != nil {
return err
}
if user != nil && user.Name != "" {
if requiredPlan != "" && !launch.PlanSatisfies(user.Plan, requiredPlan) {
return fmt.Errorf("plan upgrade required: %s needs plan %s, you have %s", modelName, requiredPlan, user.Plan)
}
return nil
}
return fmt.Errorf("%s requires sign in", modelName)
}
+186
View File
@@ -0,0 +1,186 @@
package cmd
import (
"errors"
"os"
"path/filepath"
"strconv"
"strings"
"testing"
"time"
"github.com/spf13/cobra"
coreagent "github.com/ollama/ollama/agent"
agenttools "github.com/ollama/ollama/agent/tools"
"github.com/ollama/ollama/api"
"github.com/ollama/ollama/cmd/config"
agentchat "github.com/ollama/ollama/cmd/tui/chat"
)
func TestAgentSystemPromptIncludesSessionWorkingDirOnce(t *testing.T) {
workingDir := t.TempDir()
prompt := agentSystemPromptAtWithWorkingDir(
time.Date(2026, time.July, 14, 0, 0, 0, 0, time.UTC),
"test-model",
"model instruction",
"caller instruction",
workingDir,
)
workingDirInstruction := "Current working directory: " + strconv.Quote(workingDir) + "."
if got := strings.Count(prompt, workingDirInstruction); got != 1 {
t.Fatalf("working directory instruction count = %d, want 1:\n%s", got, prompt)
}
for _, want := range []string{"model instruction", "caller instruction"} {
if !strings.Contains(prompt, want) {
t.Fatalf("prompt missing %q:\n%s", want, prompt)
}
}
}
func TestAgentWorkingDirIgnoresGetwdFailure(t *testing.T) {
original := agentGetwd
agentGetwd = func() (string, error) {
return "", errors.New("getwd failed")
}
t.Cleanup(func() {
agentGetwd = original
})
if got := agentWorkingDir(); got != "" {
t.Fatalf("working directory = %q, want empty on getwd failure", got)
}
}
func TestAgentSystemPromptIncludesSkillCatalog(t *testing.T) {
dir := t.TempDir()
if err := os.Mkdir(filepath.Join(dir, "release-notes"), 0o755); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(filepath.Join(dir, "release-notes", "SKILL.md"), []byte("---\nname: release-notes\ndescription: Draft releases.\n---\nUse bullets."), 0o644); err != nil {
t.Fatal(err)
}
catalog, err := coreagent.DiscoverSkills(dir)
if err != nil {
t.Fatal(err)
}
got := agentSystemPromptAtWithWorkingDir(time.Date(2026, 7, 14, 0, 0, 0, 0, time.UTC), "model", "", catalog.SystemContext(), "")
if !strings.Contains(got, "release-notes: Draft releases.") || !strings.Contains(got, "normal approval rules") {
t.Fatalf("system prompt missing skill context: %q", got)
}
}
func TestAgentSkillSystemContextRequiresAvailableEnabledSkillTool(t *testing.T) {
dir := t.TempDir()
if err := os.Mkdir(filepath.Join(dir, "release-notes"), 0o755); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(filepath.Join(dir, "release-notes", "SKILL.md"), []byte("---\nname: release-notes\ndescription: Draft releases.\n---\nUse bullets."), 0o644); err != nil {
t.Fatal(err)
}
catalog, err := coreagent.DiscoverSkills(dir)
if err != nil {
t.Fatal(err)
}
registry := &coreagent.Registry{}
registry.Register(&agenttools.Skill{Catalog: catalog})
if got := agentSkillSystemContext(catalog, registry, false); !strings.Contains(got, "release-notes: Draft releases.") {
t.Fatalf("enabled skill context = %q", got)
}
if got := agentSkillSystemContext(catalog, registry, true); got != "" {
t.Fatalf("disabled tools should omit skill context, got %q", got)
}
if got := agentSkillSystemContext(catalog, &coreagent.Registry{}, false); got != "" {
t.Fatalf("unavailable skill tool should omit skill context, got %q", got)
}
}
func TestAgentSelectionItemsUseLaunchSections(t *testing.T) {
items := agentSelectionItems([]agentchat.ModelOption{
{Name: "glm-5.2:cloud", Description: "cloud", Recommended: true, Cloud: true},
{Name: "llama3.2", Description: "local"},
})
if len(items) != 2 {
t.Fatalf("items = %d, want 2", len(items))
}
if !items[0].Recommended {
t.Fatalf("cloud recommendation should be pinned: %#v", items[0])
}
if items[1].Recommended {
t.Fatalf("local selected model should stay in launch More section: %#v", items[1])
}
if items[1].Description != "local" {
t.Fatalf("selected model description = %q, want plain description", items[1].Description)
}
}
func TestContextWindowFromRecommendationsMatchesCloudModel(t *testing.T) {
got := contextWindowFromRecommendations("glm-5.2:cloud", []api.ModelRecommendation{
{Model: "gemma4:cloud", ContextLength: 32768},
{Model: "glm-5.2:cloud", ContextLength: 1048576},
})
if got != 1048576 {
t.Fatalf("context window = %d, want 1048576", got)
}
}
func TestShowResponseContextWindowReadsArchitectureContextLength(t *testing.T) {
got := showResponseContextWindow(&api.ShowResponse{
ModelInfo: map[string]any{
"qwen3.context_length": uint32(262144),
"qwen3.rope.scaling.original_context_length": uint32(32768),
},
})
if got != 262144 {
t.Fatalf("context window = %d, want 262144", got)
}
}
func TestProcessContextWindowForModelMatchesLatestAlias(t *testing.T) {
got := processContextWindowForModel("ornith", &api.ProcessResponse{
Models: []api.ProcessModelResponse{
{Name: "other:latest", Model: "other:latest", ContextLength: 32768},
{Name: "ornith:latest", Model: "ornith:latest", ContextLength: 262144},
},
})
if got != 262144 {
t.Fatalf("context window = %d, want 262144", got)
}
}
func TestSaveLastAgentModel(t *testing.T) {
setCmdTestHome(t, t.TempDir())
if err := saveLastAgentModel(" qwen3:8b "); err != nil {
t.Fatalf("saveLastAgentModel returned error: %v", err)
}
if got := config.LastModel(); got != "qwen3:8b" {
t.Fatalf("last model = %q, want qwen3:8b", got)
}
if err := saveLastAgentModel(" "); err != nil {
t.Fatalf("saveLastAgentModel blank returned error: %v", err)
}
if got := config.LastModel(); got != "qwen3:8b" {
t.Fatalf("blank save changed last model to %q", got)
}
}
func TestApplyAgentFlagsNoTools(t *testing.T) {
cmd := &cobra.Command{}
registerAgentFlags(cmd)
if err := cmd.Flags().Set("no-tools", "true"); err != nil {
t.Fatal(err)
}
var opts agentTUIOptions
if _, err := applyAgentFlags(cmd, &opts); err != nil {
t.Fatalf("applyAgentFlags returned error: %v", err)
}
if !opts.ToolsDisabled {
t.Fatal("--no-tools should disable tools")
}
}
+130 -66
View File
@@ -27,6 +27,7 @@ import (
"strings"
"sync/atomic"
"syscall"
"text/tabwriter"
"time"
"github.com/containerd/console"
@@ -55,7 +56,6 @@ import (
"github.com/ollama/ollama/types/model"
"github.com/ollama/ollama/types/syncmap"
"github.com/ollama/ollama/version"
xcmd "github.com/ollama/ollama/x/cmd"
xcreate "github.com/ollama/ollama/x/create"
xcreateclient "github.com/ollama/ollama/x/create/client"
"github.com/ollama/ollama/x/imagegen"
@@ -96,6 +96,8 @@ func init() {
}
launch.DefaultConfirmPrompt = tui.RunConfirmWithOptions
launch.DefaultSpinner = tui.RunSpinner
}
func runTUISingleSelector(title string, items []launch.SelectionItem, current string, updates <-chan []launch.SelectionItem) (string, error) {
@@ -885,11 +887,6 @@ func RunHandler(cmd *cobra.Command, args []string) error {
return imagegen.RunCLI(cmd, name, opts.Prompt, interactive, opts.KeepAlive)
}
// Check for experimental flag
isExperimental, _ := cmd.Flags().GetBool("experimental")
yoloMode, _ := cmd.Flags().GetBool("experimental-yolo")
enableWebsearch, _ := cmd.Flags().GetBool("experimental-websearch")
if interactive {
if err := loadOrUnloadModel(cmd, &opts); err != nil {
var sErr api.AuthorizationError
@@ -916,11 +913,6 @@ func RunHandler(cmd *cobra.Command, args []string) error {
}
}
// Use experimental agent loop with tools
if isExperimental {
return xcmd.GenerateInteractive(cmd, opts.Model, opts.WordWrap, opts.Options, opts.Think, opts.HideThinking, opts.KeepAlive, yoloMode, enableWebsearch)
}
return generateInteractive(cmd, opts)
}
if err := generate(cmd, opts); err != nil {
@@ -986,6 +978,100 @@ func SignoutHandler(cmd *cobra.Command, args []string) error {
return nil
}
func UsageHandler(cmd *cobra.Command, args []string) error {
out := cmd.OutOrStdout()
client, err := api.ClientFromEnvironment()
if err != nil {
return err
}
usage, err := client.Usage(cmd.Context())
if err != nil {
var aErr api.AuthorizationError
if errors.As(err, &aErr) && aErr.StatusCode == http.StatusUnauthorized {
fmt.Fprintln(out, "You need to be signed in to Ollama to view usage.")
fmt.Fprintln(out)
if aErr.SigninURL != "" {
_ = browser.OpenURL(aErr.SigninURL)
fmt.Fprintf(out, ConnectInstructions, aErr.SigninURL)
}
return nil
}
return err
}
fmt.Fprintln(out, "Usage")
details := tabwriter.NewWriter(out, 0, 4, 2, ' ', 0)
fmt.Fprintf(details, " Period\t%s to %s\n", usage.Activity.Period.StartingAt.Format("2006-01-02"), usage.Activity.Period.EndingAt.Format("2006-01-02"))
fmt.Fprintf(details, " Spend\t$%s\n", usage.Activity.Cost)
if err := details.Flush(); err != nil {
return err
}
if len(usage.Activity.Models) == 0 && usageLimitEmpty(usage.Limits.Session) && usageLimitEmpty(usage.Limits.Weekly) {
fmt.Fprintln(out)
fmt.Fprintln(out, "No usage recorded for this period.")
return nil
}
if len(usage.Activity.Models) > 0 {
fmt.Fprintln(out)
fmt.Fprintln(out, "Activity")
table := tabwriter.NewWriter(out, 0, 4, 2, ' ', 0)
fmt.Fprintln(table, " Model\tRequests\tSpend")
for _, m := range usage.Activity.Models {
fmt.Fprintf(table, " %s\t%d\t$%s\n", usageModelName(m.Name), m.RequestCount, m.Cost)
}
if err := table.Flush(); err != nil {
return err
}
}
if err := writeUsageLimit(out, "Session", usage.Limits.Session); err != nil {
return err
}
if err := writeUsageLimit(out, "Weekly", usage.Limits.Weekly); err != nil {
return err
}
return nil
}
func usageLimitEmpty(limit api.UsageLimit) bool {
return limit.Usage == 0 && len(limit.Models) == 0
}
func usageModelName(name string) string {
switch name {
case "web search":
return "Web Search"
case "web fetch":
return "Web Fetch"
default:
return name
}
}
func writeUsageLimit(out io.Writer, name string, limit api.UsageLimit) error {
if usageLimitEmpty(limit) {
return nil
}
fmt.Fprintln(out)
fmt.Fprintln(out, name)
table := tabwriter.NewWriter(out, 0, 4, 2, ' ', 0)
fmt.Fprintf(table, " Used\t%.1f%%\n", limit.Usage*100)
if len(limit.Models) > 0 {
fmt.Fprintln(table, " Model\tRequests")
}
for _, m := range limit.Models {
fmt.Fprintf(table, " %s\t%d\n", usageModelName(m.Name), m.RequestCount)
}
return table.Flush()
}
func PushHandler(cmd *cobra.Command, args []string) error {
client, err := api.ClientFromEnvironment()
if err != nil {
@@ -2143,72 +2229,32 @@ func ensureServerRunning(ctx context.Context) error {
}
func launchInteractiveModel(cmd *cobra.Command, modelName string) error {
opts := runOptions{
Model: modelName,
WordWrap: os.Getenv("TERM") == "xterm-256color",
Options: map[string]any{},
ShowConnect: true,
}
client, err := api.ClientFromEnvironment()
if err != nil {
return err
}
requestedCloud := modelref.HasExplicitCloudSource(modelName)
info, err := func() (*api.ShowResponse, error) {
showReq := &api.ShowRequest{Name: modelName}
info, err := client.Show(cmd.Context(), showReq)
var se api.StatusError
if errors.As(err, &se) && se.StatusCode == http.StatusNotFound {
if requestedCloud {
return nil, err
}
if err := PullHandler(cmd, []string{modelName}); err != nil {
return nil, err
}
return client.Show(cmd.Context(), &api.ShowRequest{Name: modelName})
}
return info, err
}()
opts := agentTUIOptions{
Model: modelName,
Options: map[string]any{},
}
info, err := prepareAgentModel(cmd, client, &opts, false)
if err != nil {
if handleCloudAuthorizationError(err) {
return nil
}
return err
}
opts.System = info.System
ensureCloudStub(cmd.Context(), client, modelName)
opts.Think, err = inferThinkingOption(&info.Capabilities, &opts, false)
if err != nil {
if err := saveLastAgentModel(opts.Model); err != nil {
return err
}
audioCapable := slices.Contains(info.Capabilities, model.CapabilityAudio)
opts.MultiModal = slices.Contains(info.Capabilities, model.CapabilityVision) || audioCapable
// TODO: remove the projector info and vision info checks below,
// these are left in for backwards compatibility with older servers
// that don't have the capabilities field in the model info
if len(info.ProjectorInfo) != 0 {
opts.MultiModal = true
}
for k := range info.ModelInfo {
if strings.Contains(k, ".vision.") {
opts.MultiModal = true
break
if err := GenerateAgentTUI(cmd, client, opts); err != nil {
if handleCloudAuthorizationError(err) {
return nil
}
}
applyShowResponseToRunOptions(&opts, info)
if err := loadOrUnloadModel(cmd, &opts); err != nil {
return fmt.Errorf("error loading model: %w", err)
}
if err := generateInteractive(cmd, opts); err != nil {
return fmt.Errorf("error running model: %w", err)
return fmt.Errorf("error running agent: %w", err)
}
return nil
}
@@ -2324,7 +2370,7 @@ func runLauncherAction(cmd *cobra.Command, action tui.TUIAction, deps launcherDe
func launcherActionExitsLoop(integration string) bool {
switch integration {
case "codex-app", "vscode":
case "chatgpt", "codex-app", "vscode":
return true
default:
return false
@@ -2413,9 +2459,6 @@ func NewCLI() *cobra.Command {
runCmd.Flags().Bool("hidethinking", false, "Hide thinking output (if provided)")
runCmd.Flags().Bool("truncate", false, "For embedding models: truncate inputs exceeding context length (default: true). Set --truncate=false to error instead")
runCmd.Flags().Int("dimensions", 0, "Truncate output embeddings to specified dimension (embedding models only)")
runCmd.Flags().Bool("experimental", false, "Enable experimental agent loop with tools")
runCmd.Flags().Bool("experimental-yolo", false, "Skip all tool approval prompts (use with caution)")
runCmd.Flags().Bool("experimental-websearch", false, "Enable web search tool in experimental mode")
// Image generation flags (width, height, steps, seed, etc.)
imagegen.RegisterFlags(runCmd)
@@ -2423,6 +2466,15 @@ func NewCLI() *cobra.Command {
runCmd.Flags().Bool("imagegen", false, "Use the imagegen runner for LLM inference")
runCmd.Flags().MarkHidden("imagegen")
agentCmd := &cobra.Command{
Use: "agent",
Short: "Run an agent",
Args: cobra.ExactArgs(0),
PreRunE: checkServerHeartbeat,
RunE: AgentHandler,
}
registerAgentFlags(agentCmd)
stopCmd := &cobra.Command{
Use: "stop MODEL",
Short: "Stop a running model",
@@ -2493,6 +2545,14 @@ func NewCLI() *cobra.Command {
RunE: SignoutHandler,
}
usageCmd := &cobra.Command{
Use: "usage",
Short: "Show your ollama.com usage",
Args: cobra.ExactArgs(0),
PreRunE: checkServerHeartbeat,
RunE: UsageHandler,
}
listCmd := &cobra.Command{
Use: "list",
Aliases: []string{"ls"},
@@ -2553,9 +2613,11 @@ func NewCLI() *cobra.Command {
createCmd,
showCmd,
runCmd,
agentCmd,
stopCmd,
pullCmd,
pushCmd,
usageCmd,
listCmd,
psCmd,
copyCmd,
@@ -2600,6 +2662,7 @@ func NewCLI() *cobra.Command {
createCmd,
showCmd,
runCmd,
agentCmd,
stopCmd,
pullCmd,
pushCmd,
@@ -2607,6 +2670,7 @@ func NewCLI() *cobra.Command {
loginCmd,
signoutCmd,
logoutCmd,
usageCmd,
listCmd,
psCmd,
copyCmd,
+1 -1
View File
@@ -249,7 +249,7 @@ func TestRunLauncherAction_GUIAppsExitTUILoop(t *testing.T) {
cmd := &cobra.Command{}
cmd.SetContext(context.Background())
for _, integration := range []string{"codex-app", "vscode"} {
for _, integration := range []string{"chatgpt", "vscode"} {
continueLoop, err := runLauncherAction(cmd, tui.TUIAction{Kind: tui.TUIActionLaunchIntegration, Integration: integration}, launcherDeps{
resolveRunModel: unexpectedRunModelResolution(t),
launchIntegration: func(ctx context.Context, req launch.IntegrationLaunchRequest) error {
+96
View File
@@ -1398,6 +1398,102 @@ func TestListHandler(t *testing.T) {
}
}
func TestUsageHandler(t *testing.T) {
startsAt := time.Date(2026, time.June, 29, 0, 0, 0, 0, time.UTC)
endsAt := time.Date(2026, time.July, 27, 0, 0, 0, 0, time.UTC)
tests := []struct {
name string
statusCode int
response any
want string
}{
{
name: "activity and limits",
statusCode: http.StatusOK,
response: api.UsageResponse{
Activity: api.UsageActivity{
Cost: "12.34000",
Period: api.UsagePeriod{
Type: "last_4_weeks",
StartingAt: startsAt,
EndingAt: endsAt,
},
Models: []api.UsageModel{{Name: "gpt-oss:120b", RequestCount: 42, Cost: "12.34000"}},
},
Limits: api.UsageLimits{
Session: api.UsageLimit{Usage: 0.006, Models: []api.UsageModel{{Name: "web search", RequestCount: 1}}},
},
},
want: "Usage\n" +
" Period 2026-06-29 to 2026-07-27\n" +
" Spend $12.34000\n\n" +
"Activity\n" +
" Model Requests Spend\n" +
" gpt-oss:120b 42 $12.34000\n\n" +
"Session\n" +
" Used 0.6%\n" +
" Model Requests\n" +
" Web Search 1\n",
},
{
name: "no usage",
statusCode: http.StatusOK,
response: api.UsageResponse{
Activity: api.UsageActivity{
Cost: "0.00000",
Period: api.UsagePeriod{Type: "last_4_weeks", StartingAt: startsAt, EndingAt: endsAt},
Models: []api.UsageModel{},
},
Limits: api.UsageLimits{
Session: api.UsageLimit{Models: []api.UsageModel{}},
Weekly: api.UsageLimit{Models: []api.UsageModel{}},
},
},
want: "Usage\n" +
" Period 2026-06-29 to 2026-07-27\n" +
" Spend $0.00000\n\n" +
"No usage recorded for this period.\n",
},
{
name: "not signed in",
statusCode: http.StatusUnauthorized,
response: map[string]string{"error": "unauthorized"},
want: "You need to be signed in to Ollama to view usage.\n\n",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodGet || r.URL.Path != "/api/usage" {
t.Fatalf("request = %s %s, want GET /api/usage", r.Method, r.URL.Path)
}
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(tt.statusCode)
if err := json.NewEncoder(w).Encode(tt.response); err != nil {
t.Fatal(err)
}
}))
defer server.Close()
t.Setenv("OLLAMA_HOST", server.URL)
cmd := &cobra.Command{}
cmd.SetContext(t.Context())
var out bytes.Buffer
cmd.SetOut(&out)
if err := UsageHandler(cmd, nil); err != nil {
t.Fatal(err)
}
if got := out.String(); got != tt.want {
t.Errorf("unexpected output (-want +got):\n%s", cmp.Diff(tt.want, got))
}
})
}
}
func TestCreateHandler(t *testing.T) {
tests := []struct {
name string
+189
View File
@@ -0,0 +1,189 @@
package filedata
import (
"errors"
"fmt"
"io"
"net/http"
"net/url"
"os"
"path/filepath"
"regexp"
"slices"
"strings"
"github.com/ollama/ollama/api"
)
type File struct {
Path string
Data api.ImageData
}
func NormalizePath(fp string) string {
fp = strings.Trim(fp, "\"")
fp = strings.NewReplacer(
"\\ ", " ",
"\\(", "(",
"\\)", ")",
"\\[", "[",
"\\]", "]",
"\\{", "{",
"\\}", "}",
"\\$", "$",
"\\&", "&",
"\\;", ";",
"\\'", "'",
"\\\\", "\\",
"\\*", "*",
"\\?", "?",
"\\~", "~",
).Replace(fp)
if u, err := url.Parse(fp); err == nil && strings.EqualFold(u.Scheme, "file") {
return normalizeFileURL(u)
} else if normalized, ok := normalizeMalformedFileURL(fp); ok {
return normalized
}
return fp
}
// fileExtractRe matches file:// URLs and filesystem paths ending in image/audio
// extensions. Hoisted to package scope so the per-keystroke slash-completion
// path (chat.slashInputIsMultimodalFile -> ExtractNames) doesn't recompile it
// on every call.
var fileExtractRe = regexp.MustCompile(`(?:file://\S+?\.(?i:jpg|jpeg|png|webp|wav)\b)|(?:(?:[a-zA-Z]:)?(?:\./|\.\\|/|\\)[\S\\ ]+?\.(?i:jpg|jpeg|png|webp|wav)\b)`)
func ExtractNames(input string) []string {
return fileExtractRe.FindAllString(input, -1)
}
func Extract(input string) (string, []api.ImageData, error) {
cleaned, files, err := ExtractWithFiles(input)
if err != nil {
return "", nil, err
}
data := make([]api.ImageData, 0, len(files))
for _, file := range files {
data = append(data, file.Data)
}
return cleaned, data, nil
}
func ExtractWithFiles(input string) (string, []File, error) {
filePaths := ExtractNames(input)
var files []File
for _, fp := range filePaths {
nfp := NormalizePath(fp)
data, err := GetData(nfp)
if errors.Is(err, os.ErrNotExist) {
continue
} else if err != nil {
return "", nil, fmt.Errorf("couldn't process file %q: %w", nfp, err)
}
input = strings.ReplaceAll(input, "'"+nfp+"'", "")
input = strings.ReplaceAll(input, "'"+fp+"'", "")
input = strings.ReplaceAll(input, `"`+nfp+`"`, "")
input = strings.ReplaceAll(input, `"`+fp+`"`, "")
input = strings.ReplaceAll(input, fp, "")
files = append(files, File{Path: nfp, Data: data})
}
return strings.TrimSpace(input), files, nil
}
func GetData(filePath string) ([]byte, error) {
file, err := os.Open(filePath)
if err != nil {
return nil, err
}
defer file.Close()
buf := make([]byte, 512)
_, err = file.Read(buf)
if err != nil {
return nil, err
}
contentType := http.DetectContentType(buf)
allowedTypes := []string{"image/jpeg", "image/jpg", "image/png", "image/webp", "audio/wave"}
if !slices.Contains(allowedTypes, contentType) {
return nil, fmt.Errorf("invalid file type: %s", contentType)
}
info, err := file.Stat()
if err != nil {
return nil, err
}
var maxSize int64 = 100 * 1024 * 1024
if info.Size() > maxSize {
return nil, errors.New("file size exceeds maximum limit (100MB)")
}
buf = make([]byte, info.Size())
_, err = file.Seek(0, 0)
if err != nil {
return nil, err
}
_, err = io.ReadFull(file, buf)
if err != nil {
return nil, err
}
return buf, nil
}
func Kind(path string) string {
if strings.EqualFold(filepath.Ext(path), ".wav") {
return "audio"
}
return "image"
}
func normalizeFileURL(u *url.URL) string {
path := u.Path
if unescaped, err := url.PathUnescape(path); err == nil {
path = unescaped
}
host := u.Host
if unescaped, err := url.PathUnescape(host); err == nil {
host = unescaped
}
if len(host) >= 2 && host[1] == ':' && isASCIIAlpha(host[0]) {
return filepath.Clean(filepath.FromSlash(host + path))
}
if len(path) >= 4 && path[0] == '/' && path[2] == ':' && isASCIIAlpha(path[1]) {
path = path[1:]
}
if u.Host != "" && !strings.EqualFold(u.Host, "localhost") {
return `\\` + u.Host + filepath.FromSlash(path)
}
return filepath.FromSlash(path)
}
func normalizeMalformedFileURL(raw string) (string, bool) {
const prefix = "file://"
if !strings.HasPrefix(strings.ToLower(raw), prefix) {
return "", false
}
path := raw[len(prefix):]
if unescaped, err := url.PathUnescape(path); err == nil {
path = unescaped
}
path = strings.TrimPrefix(path, "localhost")
if len(path) >= 3 && path[0] == '/' && path[2] == ':' && isASCIIAlpha(path[1]) {
path = path[1:]
}
if len(path) >= 2 && path[1] == ':' && isASCIIAlpha(path[0]) {
return filepath.Clean(filepath.FromSlash(path)), true
}
return "", false
}
func isASCIIAlpha(b byte) bool {
return (b >= 'a' && b <= 'z') || (b >= 'A' && b <= 'Z')
}
+223
View File
@@ -0,0 +1,223 @@
package filedata
import (
"net/url"
"os"
"path/filepath"
"strings"
"testing"
)
func TestNormalizePathMalformedWindowsFileURL(t *testing.T) {
got := NormalizePath(`file://C:%5CUsers%5Cjdoe%5CPictures%5Cimg.png`)
want := filepath.Clean(`C:\Users\jdoe\Pictures\img.png`)
if got != want {
t.Fatalf("path = %q, want %q", got, want)
}
}
func TestNormalizePathTwoSlashWindowsFileURL(t *testing.T) {
got := NormalizePath(`file://C:/Users/jdoe/Pictures/img.png`)
want := filepath.Clean(`C:/Users/jdoe/Pictures/img.png`)
if got != want {
t.Fatalf("path = %q, want %q", got, want)
}
}
func TestNormalizePathLocalhostWindowsFileURL(t *testing.T) {
got := NormalizePath(`file://localhost/C:/Users/jdoe/Pictures/img.png`)
want := filepath.Clean(`C:/Users/jdoe/Pictures/img.png`)
if got != want {
t.Fatalf("path = %q, want %q", got, want)
}
}
func TestExtractNames(t *testing.T) {
// Unix style paths
input := ` some preamble
./relative\ path/one.png inbetween1 ./not a valid two.jpg inbetween2 ./1.svg
/unescaped space /three.jpeg inbetween3 /valid\ path/dir/four.png "./quoted with spaces/five.JPG
/unescaped space /six.webp inbetween6 /valid\ path/dir/seven.WEBP`
res := ExtractNames(input)
if len(res) != 7 {
t.Fatalf("len = %d, want 7", len(res))
}
assertContains(t, res[0], "one.png")
assertContains(t, res[1], "two.jpg")
assertContains(t, res[2], "three.jpeg")
assertContains(t, res[3], "four.png")
assertContains(t, res[4], "five.JPG")
assertContains(t, res[5], "six.webp")
assertContains(t, res[6], "seven.WEBP")
assertNotContains(t, res[4], "\"")
for _, r := range res {
assertNotContains(t, r, "inbetween1")
}
assertNotContainsSlice(t, res, "./1.svg")
}
func TestExtractNamesWindowsPaths(t *testing.T) {
input := ` some preamble
c:/users/jdoe/one.png inbetween1 c:/program files/someplace/two.jpg inbetween2
/absolute/nospace/three.jpeg inbetween3 /absolute/with space/four.png inbetween4
./relative\ path/five.JPG inbetween5 "./relative with/spaces/six.png inbetween6
d:\path with\spaces\seven.JPEG inbetween7 c:\users\jdoe\eight.png inbetween8
d:\program files\someplace\nine.png inbetween9 "E:\program files\someplace\ten.PNG
c:/users/jdoe/eleven.webp inbetween11 c:/program files/someplace/twelve.WebP inbetween12
d:\path with\spaces\thirteen.WEBP some ending
`
res := ExtractNames(input)
if len(res) != 13 {
t.Fatalf("len = %d, want 13", len(res))
}
assertNotContainsSlice(t, res, "inbetween2")
assertContains(t, res[0], "one.png")
assertContains(t, res[0], "c:")
assertContains(t, res[1], "two.jpg")
assertContains(t, res[1], "c:")
assertContains(t, res[2], "three.jpeg")
assertContains(t, res[3], "four.png")
assertContains(t, res[4], "five.JPG")
assertContains(t, res[5], "six.png")
assertContains(t, res[6], "seven.JPEG")
assertContains(t, res[6], "d:")
assertContains(t, res[7], "eight.png")
assertContains(t, res[7], "c:")
assertContains(t, res[8], "nine.png")
assertContains(t, res[8], "d:")
assertContains(t, res[9], "ten.PNG")
assertContains(t, res[9], "E:")
assertContains(t, res[10], "eleven.webp")
assertContains(t, res[10], "c:")
assertContains(t, res[11], "twelve.WebP")
assertContains(t, res[11], "c:")
assertContains(t, res[12], "thirteen.WEBP")
assertContains(t, res[12], "d:")
}
func TestExtractNamesDragDropPaths(t *testing.T) {
input := `file:///Users/jdoe/Pictures/one.png file://localhost/C:/Users/jdoe/Pictures/two.webp file:///C:/Users/jdoe/Pictures/three.jpg .\relative\four.png`
res := ExtractNames(input)
if len(res) != 4 {
t.Fatalf("len = %d, want 4", len(res))
}
assertContains(t, res[0], "file:///Users/jdoe/Pictures/one.png")
assertContains(t, res[1], "file://localhost/C:/Users/jdoe/Pictures/two.webp")
assertContains(t, res[2], "file:///C:/Users/jdoe/Pictures/three.jpg")
assertContains(t, res[3], `.\relative\four.png`)
}
func TestNormalizePathFileURL(t *testing.T) {
got := NormalizePath("file:///C:/Users/jdoe/Pictures/img.png")
want := filepath.FromSlash("C:/Users/jdoe/Pictures/img.png")
if got != want {
t.Fatalf("path = %q, want %q", got, want)
}
}
func TestExtractRemovesQuotedFilepath(t *testing.T) {
dir := t.TempDir()
fp := filepath.Join(dir, "img.jpg")
data := make([]byte, 600)
copy(data, []byte{
0xff, 0xd8, 0xff, 0xe0, 0x00, 0x10, 'J', 'F', 'I', 'F',
0x00, 0x01, 0x01, 0x01, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00,
0xff, 0xd9,
})
if err := os.WriteFile(fp, data, 0o600); err != nil {
t.Fatalf("failed to write test image: %v", err)
}
input := "before '" + fp + "' after"
cleaned, imgs, err := Extract(input)
if err != nil {
t.Fatalf("err: %v", err)
}
if len(imgs) != 1 {
t.Fatalf("imgs = %d, want 1", len(imgs))
}
if cleaned != "before after" {
t.Fatalf("cleaned = %q, want %q", cleaned, "before after")
}
}
func TestExtractFileURL(t *testing.T) {
dir := t.TempDir()
fp := filepath.Join(dir, "img.png")
data := make([]byte, 600)
copy(data, []byte{0x89, 'P', 'N', 'G', '\r', '\n', 0x1a, '\n'})
if err := os.WriteFile(fp, data, 0o600); err != nil {
t.Fatalf("failed to write test image: %v", err)
}
fileURL := (&url.URL{Scheme: "file", Path: fp}).String()
cleaned, imgs, err := Extract("before " + fileURL + " after")
if err != nil {
t.Fatalf("err: %v", err)
}
if len(imgs) != 1 {
t.Fatalf("imgs = %d, want 1", len(imgs))
}
if cleaned != "before after" {
t.Fatalf("cleaned = %q, want %q", cleaned, "before after")
}
}
func TestExtractWAV(t *testing.T) {
dir := t.TempDir()
fp := filepath.Join(dir, "sample.wav")
data := make([]byte, 600)
copy(data[:44], []byte{
'R', 'I', 'F', 'F',
0x58, 0x02, 0x00, 0x00,
'W', 'A', 'V', 'E',
'f', 'm', 't', ' ',
0x10, 0x00, 0x00, 0x00,
0x01, 0x00,
0x01, 0x00,
0x80, 0x3e, 0x00, 0x00,
0x00, 0x7d, 0x00, 0x00,
0x02, 0x00,
0x10, 0x00,
'd', 'a', 't', 'a',
0x34, 0x02, 0x00, 0x00,
})
if err := os.WriteFile(fp, data, 0o600); err != nil {
t.Fatalf("failed to write test audio: %v", err)
}
input := "before " + fp + " after"
cleaned, imgs, err := Extract(input)
if err != nil {
t.Fatalf("err: %v", err)
}
if len(imgs) != 1 {
t.Fatalf("imgs = %d, want 1", len(imgs))
}
if cleaned != "before after" {
t.Fatalf("cleaned = %q, want %q", cleaned, "before after")
}
}
func assertContains(t *testing.T, s, want string) {
t.Helper()
if !strings.Contains(s, want) {
t.Fatalf("%q does not contain %q", s, want)
}
}
func assertNotContains(t *testing.T, s, want string) {
t.Helper()
if strings.Contains(s, want) {
t.Fatalf("%q unexpectedly contains %q", s, want)
}
}
func assertNotContainsSlice(t *testing.T, ss []string, want string) {
t.Helper()
for _, s := range ss {
if strings.Contains(s, want) {
t.Fatalf("slice unexpectedly contains %q in %q", want, s)
}
}
}
+3 -3
View File
@@ -20,7 +20,7 @@ const (
)
var (
ErrPlanVerificationUnavailable = errors.New("Could not verify your plan. Try again in a moment.")
ErrPlanVerificationUnavailable = errors.New("Could not verify Ollama plan. Try again in a moment or use a local model.")
errUpgradeCancelled = errors.New("upgrade cancelled")
)
@@ -247,7 +247,7 @@ func (c *launcherClient) ensureCloudModelAccess(ctx context.Context, model strin
c.accountState = &state
}
if state.Status == accountStateUnknown {
return ErrPlanVerificationUnavailable
return nil
}
if state.Status == accountStateSignedOut {
@@ -259,7 +259,7 @@ func (c *launcherClient) ensureCloudModelAccess(ctx context.Context, model strin
c.accountState = &state
}
if state.Status == accountStateUnknown {
return ErrPlanVerificationUnavailable
return nil
}
}
+103 -11
View File
@@ -7,6 +7,7 @@ import (
"path/filepath"
"runtime"
"strconv"
"strings"
"github.com/ollama/ollama/envconfig"
)
@@ -37,17 +38,21 @@ func (c *Claude) findPath() (string, error) {
if runtime.GOOS == "windows" {
name = "claude.exe"
}
fallback := filepath.Join(home, ".claude", "local", name)
if _, err := os.Stat(fallback); err != nil {
return "", err
for _, fallback := range []string{
filepath.Join(home, ".local", "bin", name),
filepath.Join(home, ".claude", "local", name),
} {
if _, err := os.Stat(fallback); err == nil {
return fallback, nil
}
}
return fallback, nil
return "", fmt.Errorf("claude binary not found")
}
func (c *Claude) Run(model string, _ []LaunchModel, args []string) error {
claudePath, err := c.findPath()
claudePath, err := ensureClaudeInstalled()
if err != nil {
return fmt.Errorf("claude is not installed, install from https://code.claude.com/docs/en/quickstart")
return err
}
cmd := exec.Command(claudePath, c.args(model, args)...)
@@ -55,17 +60,104 @@ func (c *Claude) Run(model string, _ []LaunchModel, args []string) error {
cmd.Stdout = os.Stdout
cmd.Stderr = os.Stderr
env := append(os.Environ(),
"ANTHROPIC_BASE_URL="+envconfig.Host().String(),
cmd.Env = append(os.Environ(), c.envVars(model)...)
return cmd.Run()
}
func (c *Claude) envVars(model string) []string {
env := []string{
"ANTHROPIC_BASE_URL=" + envconfig.Host().String(),
"ANTHROPIC_API_KEY=",
"ANTHROPIC_AUTH_TOKEN=ollama",
"CLAUDE_CODE_ATTRIBUTION_HEADER=0",
)
"DISABLE_ERROR_REPORTING=1",
"DISABLE_FEEDBACK_COMMAND=1",
"CLAUDE_CODE_DISABLE_FEEDBACK_SURVEY=1",
}
env = append(env, c.modelEnvVars(model)...)
return env
}
cmd.Env = env
return cmd.Run()
func ensureClaudeInstalled() (string, error) {
if path, err := (&Claude{}).findPath(); err == nil {
return path, nil
}
if err := checkClaudeInstallerDependencies(); err != nil {
return "", err
}
ok, err := ConfirmPrompt("Claude Code is not installed. Install now?")
if err != nil {
return "", err
}
if !ok {
return "", fmt.Errorf("claude installation cancelled")
}
bin, args, err := claudeInstallerCommand(runtime.GOOS)
if err != nil {
return "", err
}
fmt.Fprintf(os.Stderr, "\nInstalling Claude Code...\n")
cmd := exec.Command(bin, args...)
cmd.Stdin = os.Stdin
cmd.Stdout = os.Stdout
cmd.Stderr = os.Stderr
if err := cmd.Run(); err != nil {
return "", fmt.Errorf("failed to install claude: %w", err)
}
path, err := (&Claude{}).findPath()
if err != nil {
return "", fmt.Errorf("claude was installed but the binary was not found on PATH\n\nYou may need to restart your shell")
}
fmt.Fprintf(os.Stderr, "%sClaude Code installed successfully%s\n\n", ansiGreen, ansiReset)
return path, nil
}
func checkClaudeInstallerDependencies() error {
switch runtime.GOOS {
case "windows":
if _, err := exec.LookPath("powershell"); err != nil {
return fmt.Errorf("claude is not installed and required dependencies are missing\n\nInstall the following first:\n PowerShell: https://learn.microsoft.com/powershell/\n\nThen re-run:\n ollama launch claude")
}
default:
var missing []string
if _, err := exec.LookPath("curl"); err != nil {
missing = append(missing, "curl: https://curl.se/")
}
if _, err := exec.LookPath("bash"); err != nil {
missing = append(missing, "bash: https://www.gnu.org/software/bash/")
}
if len(missing) > 0 {
return fmt.Errorf("claude is not installed and required dependencies are missing\n\nInstall the following first:\n %s\n\nThen re-run:\n ollama launch claude", strings.Join(missing, "\n "))
}
}
return nil
}
func claudeInstallerCommand(goos string) (string, []string, error) {
switch goos {
case "windows":
return "powershell", []string{
"-NoProfile",
"-ExecutionPolicy",
"Bypass",
"-Command",
"irm https://claude.ai/install.ps1 | iex",
}, nil
case "darwin", "linux":
return "bash", []string{
"-c",
"curl -fsSL https://claude.ai/install.sh | bash",
}, nil
default:
return "", nil, fmt.Errorf("unsupported platform for claude install: %s", goos)
}
}
// modelEnvVars returns Claude Code env vars that route all model tiers through Ollama.
+269
View File
@@ -1,12 +1,15 @@
package launch
import (
"fmt"
"os"
"path/filepath"
"runtime"
"slices"
"strings"
"testing"
"github.com/ollama/ollama/envconfig"
)
func TestClaudeIntegration(t *testing.T) {
@@ -67,6 +70,28 @@ func TestClaudeFindPath(t *testing.T) {
}
})
t.Run("falls back to ~/.local/bin/claude", func(t *testing.T) {
tmpDir := t.TempDir()
setTestHome(t, tmpDir)
t.Setenv("PATH", t.TempDir()) // empty dir, no claude binary
name := "claude"
if runtime.GOOS == "windows" {
name = "claude.exe"
}
fallback := filepath.Join(tmpDir, ".local", "bin", name)
os.MkdirAll(filepath.Dir(fallback), 0o755)
os.WriteFile(fallback, []byte("#!/bin/sh\n"), 0o755)
got, err := c.findPath()
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if got != fallback {
t.Errorf("findPath() = %q, want %q", got, fallback)
}
})
t.Run("returns error when neither PATH nor fallback exists", func(t *testing.T) {
tmpDir := t.TempDir()
setTestHome(t, tmpDir)
@@ -79,6 +104,210 @@ func TestClaudeFindPath(t *testing.T) {
})
}
func TestEnsureClaudeInstalled(t *testing.T) {
withConfirm := func(t *testing.T, fn func(prompt string) (bool, error)) {
t.Helper()
oldConfirm := DefaultConfirmPrompt
DefaultConfirmPrompt = func(prompt string, options ConfirmOptions) (bool, error) {
return fn(prompt)
}
t.Cleanup(func() { DefaultConfirmPrompt = oldConfirm })
}
t.Run("already installed", func(t *testing.T) {
setTestHome(t, t.TempDir())
tmpDir := t.TempDir()
t.Setenv("PATH", tmpDir)
writeFakeBinary(t, tmpDir, "claude")
withConfirm(t, func(prompt string) (bool, error) {
t.Fatalf("did not expect prompt, got %q", prompt)
return false, nil
})
bin, err := ensureClaudeInstalled()
if err != nil {
t.Fatalf("ensureClaudeInstalled() error = %v", err)
}
if filepath.Base(bin) != "claude" && filepath.Base(bin) != "claude.cmd" {
t.Fatalf("bin = %q, want claude binary", bin)
}
})
t.Run("missing dependencies", func(t *testing.T) {
setTestHome(t, t.TempDir())
t.Setenv("PATH", t.TempDir())
withConfirm(t, func(prompt string) (bool, error) {
t.Fatalf("did not expect prompt, got %q", prompt)
return false, nil
})
_, err := ensureClaudeInstalled()
if err == nil || !strings.Contains(err.Error(), "required dependencies are missing") {
t.Fatalf("expected missing dependency error, got %v", err)
}
})
t.Run("missing and user declines install", func(t *testing.T) {
setTestHome(t, t.TempDir())
tmpDir := t.TempDir()
t.Setenv("PATH", tmpDir)
writeClaudeInstallerDeps(t, tmpDir)
withConfirm(t, func(prompt string) (bool, error) {
if prompt != "Claude Code is not installed. Install now?" {
t.Fatalf("unexpected prompt: %q", prompt)
}
return false, nil
})
_, err := ensureClaudeInstalled()
if err == nil || !strings.Contains(err.Error(), "installation cancelled") {
t.Fatalf("expected cancellation error, got %v", err)
}
})
t.Run("missing and user confirms install succeeds", func(t *testing.T) {
if runtime.GOOS == "windows" {
t.Skip("uses POSIX shell fake binaries")
}
homeDir := t.TempDir()
setTestHome(t, homeDir)
tmpDir := t.TempDir()
t.Setenv("PATH", tmpDir)
writeFakeBinary(t, tmpDir, "curl")
installLog := filepath.Join(tmpDir, "bash.log")
installedClaude := filepath.Join(homeDir, ".local", "bin", "claude")
bashScript := fmt.Sprintf(`#!/bin/sh
echo "$@" >> %q
if [ "$1" = "-c" ]; then
/bin/mkdir -p %q
/bin/cat > %q <<'EOS'
#!/bin/sh
exit 0
EOS
/bin/chmod +x %q
fi
exit 0
`, installLog, filepath.Dir(installedClaude), installedClaude, installedClaude)
if err := os.WriteFile(filepath.Join(tmpDir, "bash"), []byte(bashScript), 0o755); err != nil {
t.Fatalf("failed to write fake bash: %v", err)
}
withConfirm(t, func(prompt string) (bool, error) {
return true, nil
})
bin, err := ensureClaudeInstalled()
if err != nil {
t.Fatalf("ensureClaudeInstalled() error = %v", err)
}
if bin != installedClaude {
t.Fatalf("bin = %q, want %q", bin, installedClaude)
}
logData, err := os.ReadFile(installLog)
if err != nil {
t.Fatalf("failed to read install log: %v", err)
}
if !strings.Contains(string(logData), "https://claude.ai/install.sh") {
t.Fatalf("expected install.sh command in log, got:\n%s", string(logData))
}
})
t.Run("install command fails", func(t *testing.T) {
if runtime.GOOS == "windows" {
t.Skip("uses POSIX shell fake binaries")
}
setTestHome(t, t.TempDir())
tmpDir := t.TempDir()
t.Setenv("PATH", tmpDir)
writeFakeBinary(t, tmpDir, "curl")
if err := os.WriteFile(filepath.Join(tmpDir, "bash"), []byte("#!/bin/sh\nexit 1\n"), 0o755); err != nil {
t.Fatalf("failed to write fake bash: %v", err)
}
withConfirm(t, func(prompt string) (bool, error) {
return true, nil
})
_, err := ensureClaudeInstalled()
if err == nil || !strings.Contains(err.Error(), "failed to install claude") {
t.Fatalf("expected install failure error, got %v", err)
}
})
}
func writeClaudeInstallerDeps(t *testing.T, dir string) {
t.Helper()
if runtime.GOOS == "windows" {
writeFakeBinary(t, dir, "powershell")
return
}
writeFakeBinary(t, dir, "curl")
writeFakeBinary(t, dir, "bash")
}
func TestClaudeInstallerCommand(t *testing.T) {
tests := []struct {
name string
goos string
wantBin string
want string
wantErr string
}{
{
name: "unix",
goos: "linux",
wantBin: "bash",
want: "curl -fsSL https://claude.ai/install.sh | bash",
},
{
name: "macos",
goos: "darwin",
wantBin: "bash",
want: "curl -fsSL https://claude.ai/install.sh | bash",
},
{
name: "windows",
goos: "windows",
wantBin: "powershell",
want: "irm https://claude.ai/install.ps1 | iex",
},
{
name: "unsupported",
goos: "plan9",
wantErr: "unsupported platform",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
bin, args, err := claudeInstallerCommand(tt.goos)
if tt.wantErr != "" {
if err == nil || !strings.Contains(err.Error(), tt.wantErr) {
t.Fatalf("expected error containing %q, got %v", tt.wantErr, err)
}
return
}
if err != nil {
t.Fatalf("claudeInstallerCommand() error = %v", err)
}
if bin != tt.wantBin {
t.Fatalf("bin = %q, want %q", bin, tt.wantBin)
}
if !slices.Contains(args, tt.want) {
t.Fatalf("args = %v, want command containing %q", args, tt.want)
}
})
}
}
func TestClaudeArgs(t *testing.T) {
c := &Claude{}
@@ -93,6 +322,7 @@ func TestClaudeArgs(t *testing.T) {
{"with model and verbose", "llama3.2", []string{"--verbose"}, []string{"--model", "llama3.2", "--verbose"}},
{"empty model with help", "", []string{"--help"}, []string{"--help"}},
{"with allowed tools", "llama3.2", []string{"--allowedTools", "Read,Write,Bash"}, []string{"--model", "llama3.2", "--allowedTools", "Read,Write,Bash"}},
{"with channels", "llama3.2", []string{"--channels", "plugin:telegram@claude-plugins-official"}, []string{"--model", "llama3.2", "--channels", "plugin:telegram@claude-plugins-official"}},
}
for _, tt := range tests {
@@ -105,6 +335,45 @@ func TestClaudeArgs(t *testing.T) {
}
}
func TestClaudeEnvVars(t *testing.T) {
c := &Claude{}
envMap := func(envs []string) map[string]string {
m := make(map[string]string)
for _, e := range envs {
k, v, _ := strings.Cut(e, "=")
m[k] = v
}
return m
}
got := envMap(c.envVars("llama3.2"))
for key, want := range map[string]string{
"ANTHROPIC_BASE_URL": envconfig.Host().String(),
"ANTHROPIC_API_KEY": "",
"ANTHROPIC_AUTH_TOKEN": "ollama",
"CLAUDE_CODE_ATTRIBUTION_HEADER": "0",
"DISABLE_ERROR_REPORTING": "1",
"DISABLE_FEEDBACK_COMMAND": "1",
"CLAUDE_CODE_DISABLE_FEEDBACK_SURVEY": "1",
"ANTHROPIC_DEFAULT_OPUS_MODEL": "llama3.2",
"ANTHROPIC_DEFAULT_SONNET_MODEL": "llama3.2",
"ANTHROPIC_DEFAULT_HAIKU_MODEL": "llama3.2",
"CLAUDE_CODE_SUBAGENT_MODEL": "llama3.2",
} {
if got[key] != want {
t.Errorf("%s = %q, want %q", key, got[key], want)
}
}
// Both variables disable Claude Code feature-flag evaluation, which keeps Channels unavailable.
for _, key := range []string{"CLAUDE_CODE_DISABLE_NONESSENTIAL_TRAFFIC", "DISABLE_TELEMETRY"} {
if _, ok := got[key]; ok {
t.Errorf("%s must not be set by Ollama", key)
}
}
}
func TestClaudeModelEnvVars(t *testing.T) {
c := &Claude{}
+167 -42
View File
@@ -2,6 +2,7 @@ package launch
import (
"encoding/json"
"errors"
"fmt"
"os"
"os/exec"
@@ -16,13 +17,14 @@ import (
)
const (
chatGPTIntegrationName = "chatgpt"
codexAppIntegrationName = "codex-app"
codexAppProfileName = "ollama-launch-codex-app"
codexAppBundleID = "com.openai.codex"
codexAppModelCatalogFilename = "ollama-launch-models.json"
codexAppRestoreHint = "To restore your usual Codex profile, run: ollama launch codex-app --restore"
codexAppConfigurationSuccess = "Codex App profile changed to Ollama."
codexAppRestoreSuccess = "Codex App restored to your usual profile."
codexAppRestoreHint = "To restore your usual ChatGPT profile, run: ollama launch chatgpt --restore"
codexAppConfigurationSuccess = "ChatGPT profile changed to Ollama."
codexAppRestoreSuccess = "ChatGPT restored to your usual profile."
)
var (
@@ -49,7 +51,7 @@ var (
// model while leaving model discovery and switching to Codex's Ollama provider.
type CodexApp struct{}
func (c *CodexApp) String() string { return "Codex App" }
func (c *CodexApp) String() string { return "ChatGPT" }
func (c *CodexApp) Supported() error { return codexAppSupported() }
@@ -68,7 +70,7 @@ func (c *CodexApp) Configure(model string) error {
func (c *CodexApp) ConfigureWithModels(primary string, models []LaunchModel) error {
primary = strings.TrimSpace(primary)
if primary == "" {
return fmt.Errorf("codex-app requires a model")
return fmt.Errorf("chatgpt requires a model")
}
configPath, err := codexConfigPath()
@@ -106,7 +108,10 @@ func (c *CodexApp) CurrentModel() string {
if parsed.RootString(codexRootModelProviderKey) == profileName {
baseURL := parsed.ProviderString(profileName, "base_url")
if codexNormalizeURL(baseURL) == codexNormalizeURL(codexBaseURL()) && codexAppCatalogHealthy(parsed, profileName) {
return strings.TrimSpace(parsed.RootString(codexRootModelKey))
model := strings.TrimSpace(parsed.RootString(codexRootModelKey))
if codexAppCatalogContainsModel(model) {
return model
}
}
}
}
@@ -125,7 +130,11 @@ func (c *CodexApp) CurrentModel() string {
if !codexAppCatalogHealthy(parsed, profileName) {
return ""
}
return strings.TrimSpace(parsed.ProfileString(profileName, codexRootModelKey))
model := strings.TrimSpace(parsed.ProfileString(profileName, codexRootModelKey))
if !codexAppCatalogContainsModel(model) {
return ""
}
return model
}
func codexAppManagedProfileNames() []string {
@@ -169,6 +178,40 @@ func codexAppCatalogHealthy(config codexParsedConfig, profileName string) bool {
return len(catalog.Models) > 0
}
// codexAppCatalogContainsModel reports whether model appears as a slug in the
// Ollama-managed model catalog. When the configured model is not in the catalog
// the user has drifted away from the launch-managed model (e.g. by selecting a
// built-in OpenAI model in the Codex App UI), and the launch config should be
// treated as inactive.
func codexAppCatalogContainsModel(model string) bool {
if strings.TrimSpace(model) == "" {
return false
}
catalogPath, err := codexAppModelCatalogPath()
if err != nil {
return false
}
data, err := os.ReadFile(catalogPath)
if err != nil {
return false
}
var catalog struct {
Models []struct {
Slug string `json:"slug"`
} `json:"models"`
}
if err := json.Unmarshal(data, &catalog); err != nil {
return false
}
target := codexAppCatalogModelKey(model)
for _, m := range catalog.Models {
if codexAppCatalogModelKey(m.Slug) == target {
return true
}
}
return false
}
func writeCodexAppConfig(configPath, model, modelCatalogPath string) error {
baseURL := codexBaseURL()
@@ -209,10 +252,10 @@ func writeCodexAppConfig(configPath, model, modelCatalogPath string) error {
func codexValidateAppConfigText(config codexParsedConfig, model, modelCatalogPath, baseURL string) error {
if got, ok := config.RootStringOK(codexRootProfileKey); ok {
return fmt.Errorf("generated Codex App config still contains legacy profile = %q", got)
return fmt.Errorf("generated ChatGPT config still contains legacy profile = %q", got)
}
if config.Exists("profiles", codexAppProfileName) {
return fmt.Errorf("generated Codex App config still contains legacy profiles.%s table", codexAppProfileName)
return fmt.Errorf("generated ChatGPT config still contains legacy profiles.%s table", codexAppProfileName)
}
for _, check := range []struct {
path []string
@@ -226,14 +269,14 @@ func codexValidateAppConfigText(config codexParsedConfig, model, modelCatalogPat
{[]string{"model_providers", codexAppProfileName, "wire_api"}, "responses"},
} {
if got, ok := config.String(check.path...); !ok || got != check.want {
return fmt.Errorf("generated Codex App config missing %s = %q", strings.Join(check.path, "."), check.want)
return fmt.Errorf("generated ChatGPT config missing %s = %q", strings.Join(check.path, "."), check.want)
}
}
return nil
}
func (c *CodexApp) Onboard() error {
return config.MarkIntegrationOnboarded(codexAppIntegrationName)
return config.MarkIntegrationOnboarded(chatGPTIntegrationName)
}
func (c *CodexApp) RequiresInteractiveOnboarding() bool {
@@ -257,9 +300,9 @@ func (c *CodexApp) Run(_ string, _ []LaunchModel, args []string) error {
return err
}
if len(args) > 0 {
return fmt.Errorf("codex-app does not accept extra arguments")
return fmt.Errorf("chatgpt does not accept extra arguments")
}
return codexAppLaunchOrRestart("Restart Codex to use Ollama?", nil)
return codexAppLaunchOrRestart("Restart ChatGPT to use Ollama?", nil)
}
func (c *CodexApp) Restore() error {
@@ -283,7 +326,7 @@ func (c *CodexApp) Restore() error {
if err := codexAppRemoveOwnedCatalog(); err != nil {
return codexAppRestoreFailure(configPath, err)
}
return codexAppLaunchOrRestart("Restart Codex to use your usual profile?", nil)
return codexAppLaunchOrRestart("Restart ChatGPT to use your usual profile?", nil)
}
return codexAppRestoreFailure(configPath, err)
}
@@ -319,11 +362,11 @@ func (c *CodexApp) Restore() error {
if err := removeCodexAppRestoreState(); err != nil {
return codexAppRestoreFailure(configPath, err)
}
return codexAppLaunchOrRestart("Restart Codex to use your usual profile?", nil)
return codexAppLaunchOrRestart("Restart ChatGPT to use your usual profile?", nil)
}
func codexAppRestoreFailure(configPath string, err error) error {
return fmt.Errorf("restore Codex App config: %w\n\nRestore did not complete. Check these files before retrying:\n Codex config: %s\n Restore state: %s\n Model catalog: %s\n Backups: %s",
return fmt.Errorf("restore ChatGPT config: %w\n\nRestore did not complete. Check these files before retrying:\n Codex config: %s\n Restore state: %s\n Model catalog: %s\n Backups: %s",
err,
configPath,
codexAppRestoreStatePath(),
@@ -337,7 +380,7 @@ func codexAppSupported() error {
case "darwin", "windows":
return nil
default:
return fmt.Errorf("Codex App launch is only supported on macOS and Windows")
return fmt.Errorf("ChatGPT launch is only supported on macOS and Windows")
}
}
@@ -381,7 +424,7 @@ func codexAppModelCatalogPathForConfig(configPath string) string {
func writeCodexAppModelCatalog(path, primary string, models []LaunchModel) error {
if len(models) == 0 {
return fmt.Errorf("codex-app model catalog cannot be empty")
return fmt.Errorf("chatgpt model catalog cannot be empty")
}
baseInstructions := codexAppBaseInstructions()
@@ -544,9 +587,12 @@ func codexAppAppPath() string {
}
func codexAppDarwinAppCandidates() []string {
candidates := []string{"/Applications/Codex.app"}
candidates := []string{"/Applications/ChatGPT.app", "/Applications/Codex.app"}
if home, err := os.UserHomeDir(); err == nil {
candidates = append(candidates, filepath.Join(home, "Applications", "Codex.app"))
candidates = append(candidates,
filepath.Join(home, "Applications", "ChatGPT.app"),
filepath.Join(home, "Applications", "Codex.app"),
)
}
return candidates
}
@@ -558,6 +604,11 @@ func codexAppWindowsAppCandidates() []string {
}
candidates := []string{
filepath.Join(local, "Programs", "ChatGPT", "ChatGPT.exe"),
filepath.Join(local, "Programs", "OpenAI ChatGPT", "ChatGPT.exe"),
filepath.Join(local, "ChatGPT", "ChatGPT.exe"),
filepath.Join(local, "OpenAI ChatGPT", "ChatGPT.exe"),
filepath.Join(local, "OpenAI", "ChatGPT", "ChatGPT.exe"),
filepath.Join(local, "Programs", "Codex", "Codex.exe"),
filepath.Join(local, "Programs", "OpenAI Codex", "Codex.exe"),
filepath.Join(local, "Codex", "Codex.exe"),
@@ -566,6 +617,11 @@ func codexAppWindowsAppCandidates() []string {
filepath.Join(local, "openai-codex-electron", "Codex.exe"),
}
for _, pattern := range []string{
filepath.Join(local, "Programs", "ChatGPT", "app-*", "ChatGPT.exe"),
filepath.Join(local, "Programs", "OpenAI ChatGPT", "app-*", "ChatGPT.exe"),
filepath.Join(local, "ChatGPT", "app-*", "ChatGPT.exe"),
filepath.Join(local, "OpenAI ChatGPT", "app-*", "ChatGPT.exe"),
filepath.Join(local, "OpenAI", "ChatGPT", "app-*", "ChatGPT.exe"),
filepath.Join(local, "Programs", "Codex", "app-*", "Codex.exe"),
filepath.Join(local, "Programs", "OpenAI Codex", "app-*", "Codex.exe"),
filepath.Join(local, "Codex", "app-*", "Codex.exe"),
@@ -628,22 +684,54 @@ func codexAppLaunchOrRestart(prompt string, launchArgs []string) error {
return err
}
if !restart {
fmt.Fprintln(os.Stderr, "\nQuit and reopen Codex when you're ready for the profile change to take effect.")
fmt.Fprintln(os.Stderr, "\nQuit and reopen ChatGPT when you're ready for the profile change to take effect.")
return nil
}
if err := codexAppQuitApp(); err != nil {
return fmt.Errorf("quit Codex: %w", err)
// A single spinner and cancellation channel span the entire restart flow
// (quit, wait, force-quit, wait, reopen) so that one Ctrl+C aborts the
// whole sequence rather than just the currently-active wait. The bubbletea
// spinner closes Cancelled() from its raw-mode Ctrl+C handler; the ANSI
// fallback relies on SIGINT terminating the process directly.
sp := StartSpinner(codexAppRestartMessage)
defer sp.Stop()
cancelled := sp.Cancelled()
isCancelled := func() bool {
if cancelled == nil {
return false
}
select {
case <-cancelled:
return true
default:
return false
}
}
if err := codexAppQuitApp(); err != nil {
return fmt.Errorf("quit ChatGPT: %w", err)
}
if isCancelled() {
return ErrCancelled
}
gracefulErr := waitForCodexAppGracefulExit(codexAppExitTimeout, cancelled)
if isCancelled() {
return ErrCancelled
}
if errors.Is(gracefulErr, ErrCancelled) {
return gracefulErr
}
gracefulErr := waitForCodexAppGracefulExit(codexAppExitTimeout)
if gracefulErr != nil && !codexAppForceQuitSupported() {
return gracefulErr
}
if codexAppForceQuitSupported() && codexAppIsRunning() {
if forceErr := codexAppForceQuit(); forceErr != nil {
return fmt.Errorf("force stop Codex: %w", forceErr)
if isCancelled() {
return ErrCancelled
}
if err := waitForCodexAppExit(codexAppForceExitTimeout); err != nil {
if forceErr := codexAppForceQuit(); forceErr != nil {
return fmt.Errorf("force stop ChatGPT: %w", forceErr)
}
if err := waitForCodexAppExit(codexAppForceExitTimeout, cancelled); err != nil {
return err
}
} else if gracefulErr != nil {
@@ -651,6 +739,10 @@ func codexAppLaunchOrRestart(prompt string, launchArgs []string) error {
return gracefulErr
}
}
if isCancelled() {
return ErrCancelled
}
sp.Stop()
if restartAppID != "" {
return codexAppOpenStart(restartAppID)
}
@@ -664,8 +756,8 @@ func codexAppForceQuitSupported() bool {
return codexAppGOOS == "darwin" || codexAppGOOS == "windows"
}
func waitForCodexAppGracefulExit(timeout time.Duration) error {
return waitForCodexAppCondition(timeout, func() bool {
func waitForCodexAppGracefulExit(timeout time.Duration, cancel <-chan struct{}) error {
return waitForCodexAppCondition(timeout, cancel, func() bool {
if codexAppGOOS == "windows" {
return !codexAppHasWindow()
}
@@ -673,21 +765,41 @@ func waitForCodexAppGracefulExit(timeout time.Duration) error {
})
}
func waitForCodexAppExit(timeout time.Duration) error {
return waitForCodexAppCondition(timeout, func() bool {
func waitForCodexAppExit(timeout time.Duration, cancel <-chan struct{}) error {
return waitForCodexAppCondition(timeout, cancel, func() bool {
return !codexAppIsRunning()
})
}
func waitForCodexAppCondition(timeout time.Duration, done func() bool) error {
// codexAppRestartMessage is the label shown next to the animated spinner while
// the ChatGPT desktop app is quitting before being reopened.
const codexAppRestartMessage = "Restarting ChatGPT..."
// waitForCodexAppCondition polls done at a 200ms cadence until it reports the
// app has exited or timeout elapses. It watches cancel (closed by the spinner
// when the user hits Ctrl+C) and returns ErrCancelled if the flow is aborted.
// The spinner itself is owned by the caller so a single spinner spans the
// whole restart sequence. When timeout is zero the loop never runs, so
// force-quit paths that short-circuit the graceful wait return immediately.
func waitForCodexAppCondition(timeout time.Duration, cancel <-chan struct{}, done func() bool) error {
deadline := time.Now().Add(timeout)
for time.Now().Before(deadline) {
if cancel != nil {
select {
case <-cancel:
return ErrCancelled
default:
}
}
if done() {
return nil
}
codexAppSleep(200 * time.Millisecond)
}
return fmt.Errorf("Codex did not quit; quit it manually and re-run the command")
if done() {
return nil
}
return fmt.Errorf("ChatGPT did not quit; quit it manually and re-run the command")
}
func defaultCodexAppOpenApp(args []string) error {
@@ -710,7 +822,7 @@ func defaultCodexAppOpenApp(args []string) error {
if appID := codexAppStartID(); appID != "" {
return codexAppOpenStart(appID)
}
return fmt.Errorf("Codex executable was not found; open Codex manually once and re-run 'ollama launch codex-app'")
return fmt.Errorf("ChatGPT was not found; install it from https://chatgpt.com/download, then re-run 'ollama launch chatgpt'")
case "darwin":
if path := codexAppAppPath(); path != "" {
cmd := exec.Command("open", path)
@@ -747,14 +859,17 @@ func defaultCodexAppOpenStartAppID(appID string) error {
func defaultCodexAppQuitApp() error {
if codexAppGOOS == "windows" {
script := `Get-Process Codex -ErrorAction SilentlyContinue | Where-Object { $_.MainWindowHandle -ne 0 } | ForEach-Object { [void]$_.CloseMainWindow() }`
script := `Get-Process ChatGPT,Codex -ErrorAction SilentlyContinue | Where-Object { $_.MainWindowHandle -ne 0 } | ForEach-Object { [void]$_.CloseMainWindow() }`
return exec.Command("powershell.exe", "-NoProfile", "-Command", script).Run()
}
scriptErr := exec.Command("osascript", "-e", `tell application "Codex" to quit`).Run()
scriptErr := exec.Command("osascript", "-e", `tell application "ChatGPT" to quit`).Run()
if scriptErr != nil {
scriptErr = exec.Command("osascript", "-e", `tell application id "`+codexAppBundleID+`" to quit`).Run()
}
if scriptErr != nil {
scriptErr = exec.Command("osascript", "-e", `tell application "Codex" to quit`).Run()
}
return scriptErr
}
@@ -793,7 +908,7 @@ func defaultCodexAppHasOpenWindow() bool {
if codexAppGOOS != "windows" {
return codexAppIsRunning()
}
script := `(Get-Process Codex -ErrorAction SilentlyContinue | Where-Object { $_.MainWindowHandle -ne 0 } | Select-Object -First 1).Id`
script := `(Get-Process ChatGPT,Codex -ErrorAction SilentlyContinue | Where-Object { $_.MainWindowHandle -ne 0 } | Select-Object -First 1).Id`
out, err := exec.Command("powershell.exe", "-NoProfile", "-Command", script).Output()
return err == nil && strings.TrimSpace(string(out)) != ""
}
@@ -803,7 +918,11 @@ func defaultCodexAppIsRunning() bool {
case "windows":
return len(codexAppMatchingProcessIDs()) > 0
case "darwin":
out, err := exec.Command("osascript", "-e", `tell application "System Events" to exists process "Codex"`).Output()
out, err := exec.Command("osascript", "-e", `tell application "System Events" to exists process "ChatGPT"`).Output()
if err == nil && strings.TrimSpace(string(out)) == "true" {
return true
}
out, err = exec.Command("osascript", "-e", `tell application "System Events" to exists process "Codex"`).Output()
if err == nil && strings.TrimSpace(string(out)) == "true" {
return true
}
@@ -845,7 +964,7 @@ func codexAppMatchingProcessIDs() []int {
}
func codexAppWindowsMatchingProcessIDs() []int {
script := fmt.Sprintf(`$current = %d; Get-CimInstance Win32_Process -Filter "Name = 'Codex.exe' OR Name = 'codex.exe'" | Where-Object { $_.ProcessId -ne $current -and ((($_.Name -ieq 'Codex.exe') -and (($null -eq $_.CommandLine) -or ($_.CommandLine -notlike '* --type=*'))) -or (($_.Name -ieq 'codex.exe') -and ($_.CommandLine -like '*app-server*'))) } | Select-Object -ExpandProperty ProcessId`, os.Getpid())
script := fmt.Sprintf(`$current = %d; Get-CimInstance Win32_Process -Filter "Name = 'Codex.exe' OR Name = 'codex.exe' OR Name = 'ChatGPT.exe' OR Name = 'chatgpt.exe'" | Where-Object { $_.ProcessId -ne $current -and ((($_.Name -ieq 'Codex.exe' -or $_.Name -ieq 'ChatGPT.exe') -and (($null -eq $_.CommandLine) -or ($_.CommandLine -notlike '* --type=*'))) -or ((($_.Name -ieq 'codex.exe') -or ($_.Name -ieq 'chatgpt.exe')) -and ($_.CommandLine -like '*app-server*'))) } | Select-Object -ExpandProperty ProcessId`, os.Getpid())
out, err := exec.Command("powershell.exe", "-NoProfile", "-Command", script).Output()
if err != nil {
return nil
@@ -865,7 +984,7 @@ func defaultCodexAppRunningAppPath() string {
if codexAppGOOS != "windows" {
return ""
}
script := `(Get-Process Codex -ErrorAction SilentlyContinue | Where-Object { $_.MainWindowHandle -ne 0 -and $_.Path } | Select-Object -First 1 -ExpandProperty Path)`
script := `(Get-Process ChatGPT,Codex -ErrorAction SilentlyContinue | Where-Object { $_.MainWindowHandle -ne 0 -and $_.Path } | Select-Object -First 1 -ExpandProperty Path)`
out, err := exec.Command("powershell.exe", "-NoProfile", "-Command", script).Output()
if err != nil {
return ""
@@ -877,7 +996,7 @@ func defaultCodexAppStartAppID() string {
if codexAppGOOS != "windows" {
return ""
}
script := `(Get-StartApps Codex | Where-Object { $_.Name -eq 'Codex' -or $_.Name -like 'Codex*' } | Select-Object -First 1 -ExpandProperty AppID)`
script := `(Get-StartApps | Where-Object { $_.Name -eq 'ChatGPT' -or $_.Name -like 'ChatGPT*' -or $_.Name -eq 'Codex' -or $_.Name -like 'Codex*' } | Select-Object -First 1 -ExpandProperty AppID)`
out, err := exec.Command("powershell.exe", "-NoProfile", "-Command", script).Output()
if err != nil {
return ""
@@ -895,7 +1014,7 @@ func defaultCodexAppCanOpenBundleID() bool {
}
func codexAppProcessMatches(command string) bool {
if strings.Contains(command, `\Codex.exe`) && strings.Contains(command, " --type=") {
if (strings.Contains(command, `\Codex.exe`) || strings.Contains(command, `\ChatGPT.exe`)) && strings.Contains(command, " --type=") {
return false
}
for _, pattern := range codexAppProcessPatterns() {
@@ -908,8 +1027,14 @@ func codexAppProcessMatches(command string) bool {
func codexAppProcessPatterns() []string {
return []string{
"ChatGPT.app/Contents/MacOS/ChatGPT",
"ChatGPT.app/Contents/Resources/codex app-server",
"Codex.app/Contents/MacOS/Codex",
"Codex.app/Contents/Resources/codex app-server",
`\ChatGPT.exe`,
`resources\chatgpt.exe app-server`,
`resources\chatgpt.exe" app-server`,
`resources\chatgpt.exe" "app-server`,
`\Codex.exe`,
`resources\codex.exe app-server`,
`resources\codex.exe" app-server`,
+207 -2
View File
@@ -2,9 +2,11 @@ package launch
import (
"encoding/json"
"errors"
"fmt"
"os"
"path/filepath"
"slices"
"strings"
"testing"
"time"
@@ -157,6 +159,27 @@ func TestCodexAppInstalledUsesMacBundleIDFallback(t *testing.T) {
}
}
func TestChatGPTMissingAppGivesDownloadRecovery(t *testing.T) {
withCodexAppPlatform(t, "darwin")
oldCanOpenID := codexAppCanOpenID
oldStat := codexAppStat
codexAppCanOpenID = func() bool { return false }
codexAppStat = func(string) (os.FileInfo, error) { return nil, os.ErrNotExist }
t.Cleanup(func() {
codexAppCanOpenID = oldCanOpenID
codexAppStat = oldStat
})
err := EnsureIntegrationInstalled(chatGPTIntegrationName, &CodexApp{})
if err == nil {
t.Fatal("expected missing ChatGPT install error")
}
if !strings.Contains(err.Error(), "chatgpt is not installed") || !strings.Contains(err.Error(), "https://chatgpt.com/download") {
t.Fatalf("missing-app error = %q, want ChatGPT download recovery", err)
}
}
func TestCodexAppConfigureActivatesOllamaProviderWithoutLegacyProfile(t *testing.T) {
tmpDir := t.TempDir()
setTestHome(t, tmpDir)
@@ -291,6 +314,58 @@ func TestCodexAppConfigureUsesAppSpecificProfileWithoutTouchingCLIProfile(t *tes
assertBackupContains(t, filepath.Join(fileutil.BackupDir(), codexAppIntegrationName, "config.toml.*"), `profile = "default"`)
}
func TestCodexAppConfigureIsIdempotentAndPreservesUnrelatedProvider(t *testing.T) {
tmpDir := t.TempDir()
setTestHome(t, tmpDir)
t.Setenv("OLLAMA_HOST", "http://127.0.0.1:9999")
configPath := filepath.Join(tmpDir, ".codex", "config.toml")
if err := os.MkdirAll(filepath.Dir(configPath), 0o755); err != nil {
t.Fatal(err)
}
existing := "[model_providers.custom]\n" +
`name = "Custom"` + "\n" +
`base_url = "https://example.invalid/v1"` + "\n" +
`env_key = "CUSTOM_API_KEY"` + "\n"
if err := os.WriteFile(configPath, []byte(existing), 0o644); err != nil {
t.Fatal(err)
}
app := &CodexApp{}
models := testLaunchModels("llama3.2", "qwen3:8b")
if err := app.ConfigureWithModels("llama3.2", models); err != nil {
t.Fatalf("first ConfigureWithModels returned error: %v", err)
}
first, err := os.ReadFile(configPath)
if err != nil {
t.Fatal(err)
}
if err := app.ConfigureWithModels("llama3.2", models); err != nil {
t.Fatalf("second ConfigureWithModels returned error: %v", err)
}
second, err := os.ReadFile(configPath)
if err != nil {
t.Fatal(err)
}
if string(second) != string(first) {
t.Fatalf("rerun changed generated config:\nfirst:\n%s\nsecond:\n%s", first, second)
}
parsed, err := codexParseConfig(string(second))
if err != nil {
t.Fatal(err)
}
if got := parsed.ProviderString("custom", "env_key"); got != "CUSTOM_API_KEY" {
t.Fatalf("custom provider env_key = %q, want preserved value", got)
}
if parsed.Exists("model_providers", codexAppProfileName, "env_key") {
t.Fatalf("managed local Ollama provider should not require an API key:\n%s", second)
}
if got := app.CurrentModel(); got != "llama3.2" {
t.Fatalf("CurrentModel = %q, want llama3.2", got)
}
}
func TestCodexCLIConfigRefreshLeavesCodexAppConfigActive(t *testing.T) {
tmpDir := t.TempDir()
setTestHome(t, tmpDir)
@@ -626,6 +701,60 @@ func TestCodexAppCurrentModelRequiresHealthyCatalog(t *testing.T) {
}
}
func TestCodexAppCurrentModelDetectsDriftedModel(t *testing.T) {
tmpDir := t.TempDir()
setTestHome(t, tmpDir)
t.Setenv("OLLAMA_HOST", "http://127.0.0.1:11434")
catalogPath := mustWriteCodexAppTestCatalog(t, "llama3.2")
configPath := filepath.Join(tmpDir, ".codex", "config.toml")
if err := os.MkdirAll(filepath.Dir(configPath), 0o755); err != nil {
t.Fatal(err)
}
content := "" +
`model = "gpt-5.5"` + "\n" +
fmt.Sprintf(`model_provider = %q`, codexAppProfileName) + "\n\n" +
fmt.Sprintf(`model_catalog_json = %q`, catalogPath) + "\n\n" +
codexProviderHeaderFor(codexAppProfileName) + "\n" +
`name = "Ollama"` + "\n" +
`base_url = "http://127.0.0.1:11434/v1/"` + "\n" +
`wire_api = "responses"` + "\n"
if err := os.WriteFile(configPath, []byte(content), 0o644); err != nil {
t.Fatal(err)
}
if got := (&CodexApp{}).CurrentModel(); got != "" {
t.Fatalf("CurrentModel = %q, want empty when model has drifted from the Ollama catalog", got)
}
}
func TestCodexAppCurrentModelAcceptsLatestSuffixDrift(t *testing.T) {
tmpDir := t.TempDir()
setTestHome(t, tmpDir)
t.Setenv("OLLAMA_HOST", "http://127.0.0.1:11434")
catalogPath := mustWriteCodexAppTestCatalog(t, "llama3.2")
configPath := filepath.Join(tmpDir, ".codex", "config.toml")
if err := os.MkdirAll(filepath.Dir(configPath), 0o755); err != nil {
t.Fatal(err)
}
content := "" +
`model = "llama3.2:latest"` + "\n" +
fmt.Sprintf(`model_provider = %q`, codexAppProfileName) + "\n\n" +
fmt.Sprintf(`model_catalog_json = %q`, catalogPath) + "\n\n" +
codexProviderHeaderFor(codexAppProfileName) + "\n" +
`name = "Ollama"` + "\n" +
`base_url = "http://127.0.0.1:11434/v1/"` + "\n" +
`wire_api = "responses"` + "\n"
if err := os.WriteFile(configPath, []byte(content), 0o644); err != nil {
t.Fatal(err)
}
if got := (&CodexApp{}).CurrentModel(); got != "llama3.2:latest" {
t.Fatalf("CurrentModel = %q, want llama3.2:latest (:latest suffix should not be treated as drift)", got)
}
}
func TestCodexAppConfigurePopulatesCatalogFromEnrichedModels(t *testing.T) {
tmpDir := t.TempDir()
setTestHome(t, tmpDir)
@@ -1392,6 +1521,61 @@ func TestCodexAppRunWaitsForGracefulExitBeforeReopening(t *testing.T) {
}
}
func TestCodexAppRunCtrlCAbortsEntireRestartFlow(t *testing.T) {
withCodexAppPlatform(t, "darwin")
restoreConfirm := withLaunchConfirmPolicy(launchConfirmPolicy{yes: true})
defer restoreConfirm()
oldSleep := codexAppSleep
oldDefaultSpinner := DefaultSpinner
t.Cleanup(func() {
codexAppSleep = oldSleep
DefaultSpinner = oldDefaultSpinner
})
// Simulate the user pressing Ctrl+C during the graceful-exit wait: the
// shared spinner's cancellation channel is closed on the first poll,
// which only happens inside the wait loop.
cancel := make(chan struct{})
var spinnerStopped bool
codexAppSleep = func(time.Duration) {
select {
case <-cancel:
default:
close(cancel)
}
}
DefaultSpinner = func(string) *Spinner {
return NewSpinner(func() { spinnerStopped = true }, cancel)
}
var calls []string
withCodexAppProcessHooks(t,
func() bool { return true }, // app stays "running" so the wait polls
func() error { calls = append(calls, "quit"); return nil },
func() error { calls = append(calls, "open"); return nil },
)
codexAppExitTimeout = 5 * time.Second
codexAppForceQuit = func() error {
calls = append(calls, "force")
return nil
}
err := (&CodexApp{}).Run("qwen3.5", nil, nil)
if !errors.Is(err, ErrCancelled) {
t.Fatalf("Run error = %v, want ErrCancelled", err)
}
if !spinnerStopped {
t.Fatal("expected the shared spinner to be stopped on cancel")
}
// The flow must abort after quit: no force-quit, no reopen, despite the app
// still being "running" (which would otherwise trigger the force-quit path).
want := []string{"quit"}
if !slices.Equal(calls, want) {
t.Fatalf("calls = %v, want the whole flow to abort after quit: %v", calls, want)
}
}
func TestCodexAppRunForceStopsMacAfterGracefulTimeout(t *testing.T) {
withCodexAppPlatform(t, "darwin")
restoreConfirm := withLaunchConfirmPolicy(launchConfirmPolicy{yes: true})
@@ -1445,7 +1629,7 @@ func TestCodexAppRunReturnsMacForceStopError(t *testing.T) {
}
err := (&CodexApp{}).Run("qwen3.5", nil, nil)
if err == nil || !strings.Contains(err.Error(), "force stop Codex") || !strings.Contains(err.Error(), "operation not permitted") {
if err == nil || !strings.Contains(err.Error(), "force stop ChatGPT") || !strings.Contains(err.Error(), "operation not permitted") {
t.Fatalf("Run error = %v, want force stop failure", err)
}
}
@@ -1587,7 +1771,7 @@ func TestCodexAppRunReturnsWindowsForceStopError(t *testing.T) {
}
err := (&CodexApp{}).Run("qwen3.5", nil, nil)
if err == nil || !strings.Contains(err.Error(), "force stop Codex") || !strings.Contains(err.Error(), "access denied") {
if err == nil || !strings.Contains(err.Error(), "force stop ChatGPT") || !strings.Contains(err.Error(), "access denied") {
t.Fatalf("Run error = %v, want force stop failure", err)
}
}
@@ -1602,6 +1786,8 @@ func TestCodexAppRunRejectsExtraArgs(t *testing.T) {
func TestCodexAppProcessMatchesMainAndAppServer(t *testing.T) {
for _, command := range []string{
"/Applications/ChatGPT.app/Contents/MacOS/ChatGPT",
"/Applications/ChatGPT.app/Contents/Resources/codex app-server --analytics-default-enabled",
"/Applications/Codex.app/Contents/MacOS/Codex",
"/Applications/Codex.app/Contents/Resources/codex app-server --analytics-default-enabled",
`C:\Users\parth\AppData\Local\Programs\Codex\Codex.exe`,
@@ -1614,6 +1800,7 @@ func TestCodexAppProcessMatchesMainAndAppServer(t *testing.T) {
}
for _, command := range []string{
"/Applications/ChatGPT.app/Contents/Frameworks/ChatGPT Helper.app/Contents/MacOS/ChatGPT Helper",
"/Applications/Codex.app/Contents/Frameworks/Codex Helper.app/Contents/MacOS/Codex Helper",
"/Applications/Codex.app/Contents/Frameworks/Electron Framework.framework/Helpers/chrome_crashpad_handler",
`"C:\Program Files\WindowsApps\OpenAI.Codex_26.429.8261.0_x64__2p2nqsd0c76g0\app\Codex.exe" --type=renderer --user-data-dir="C:\Users\parth\AppData\Roaming\Codex"`,
@@ -1625,6 +1812,24 @@ func TestCodexAppProcessMatchesMainAndAppServer(t *testing.T) {
}
}
func TestCodexAppCandidatesIncludeChatGPT(t *testing.T) {
withCodexAppPlatform(t, "darwin")
candidates := codexAppDarwinAppCandidates()
if len(candidates) == 0 || candidates[0] != "/Applications/ChatGPT.app" {
t.Fatalf("darwin candidates = %v, want ChatGPT first", candidates)
}
if !slices.Contains(candidates, "/Applications/Codex.app") {
t.Fatalf("darwin candidates = %v, want legacy Codex app", candidates)
}
withCodexAppPlatform(t, "windows")
local := filepath.Join(t.TempDir(), "LocalAppData")
t.Setenv("LOCALAPPDATA", local)
if candidates := codexAppWindowsAppCandidates(); !slices.Contains(candidates, filepath.Join(local, "Programs", "ChatGPT", "ChatGPT.exe")) {
t.Fatalf("windows candidates = %v, want ChatGPT app", candidates)
}
}
func catalogSlugs(models []map[string]any) []string {
slugs := make([]string, 0, len(models))
for _, model := range models {
+26 -26
View File
@@ -281,7 +281,7 @@ func TestLaunchCmdModelFlagFiltersDisabledCloudFromSavedConfig(t *testing.T) {
case "/api/status":
fmt.Fprintf(w, `{"cloud":{"disabled":true,"source":"config"}}`)
case "/api/show":
fmt.Fprintf(w, `{"model":"llama3.2"}`)
fmt.Fprintf(w, `{"model":"sample-model"}`)
default:
w.WriteHeader(http.StatusNotFound)
}
@@ -294,7 +294,7 @@ func TestLaunchCmdModelFlagFiltersDisabledCloudFromSavedConfig(t *testing.T) {
defer restore()
cmd := LaunchCmd(func(cmd *cobra.Command, args []string) error { return nil }, func(cmd *cobra.Command) {})
cmd.SetArgs([]string{"stubeditor", "--model", "llama3.2"})
cmd.SetArgs([]string{"stubeditor", "--model", "sample-model"})
if err := cmd.Execute(); err != nil {
t.Fatalf("launch command failed: %v", err)
}
@@ -303,14 +303,14 @@ func TestLaunchCmdModelFlagFiltersDisabledCloudFromSavedConfig(t *testing.T) {
if err != nil {
t.Fatalf("failed to reload integration config: %v", err)
}
if diff := cmp.Diff([]string{"llama3.2"}, saved.Models); diff != "" {
if diff := cmp.Diff([]string{"sample-model"}, saved.Models); diff != "" {
t.Fatalf("saved models mismatch (-want +got):\n%s", diff)
}
if diff := cmp.Diff([][]string{{"llama3.2"}}, stub.edited); diff != "" {
if diff := cmp.Diff([][]string{{"sample-model"}}, stub.edited); diff != "" {
t.Fatalf("editor models mismatch (-want +got):\n%s", diff)
}
if stub.ranModel != "llama3.2" {
t.Fatalf("expected launch to run with llama3.2, got %q", stub.ranModel)
if stub.ranModel != "sample-model" {
t.Fatalf("expected launch to run with sample-model, got %q", stub.ranModel)
}
}
@@ -325,9 +325,9 @@ func TestLaunchCmdModelFlagClearsDisabledCloudOverride(t *testing.T) {
case "/api/experimental/model-recommendations":
fmt.Fprint(w, `{"recommendations":[]}`)
case "/api/tags":
fmt.Fprint(w, `{"models":[{"name":"llama3.2"}]}`)
fmt.Fprint(w, `{"models":[{"name":"sample-model"}]}`)
case "/api/show":
fmt.Fprint(w, `{"model":"llama3.2"}`)
fmt.Fprint(w, `{"model":"sample-model"}`)
default:
w.WriteHeader(http.StatusNotFound)
}
@@ -347,7 +347,7 @@ func TestLaunchCmdModelFlagClearsDisabledCloudOverride(t *testing.T) {
DefaultSingleSelector = func(title string, items []SelectionItem, current string) (string, error) {
selectorCalls++
gotCurrent = current
return "llama3.2", nil
return "sample-model", nil
}
cmd := LaunchCmd(func(cmd *cobra.Command, args []string) error { return nil }, func(cmd *cobra.Command) {})
@@ -364,7 +364,7 @@ func TestLaunchCmdModelFlagClearsDisabledCloudOverride(t *testing.T) {
if gotCurrent != "" {
t.Fatalf("expected disabled override to be cleared before selection, got current %q", gotCurrent)
}
if stub.ranModel != "llama3.2" {
if stub.ranModel != "sample-model" {
t.Fatalf("expected launch to run with replacement local model, got %q", stub.ranModel)
}
if !strings.Contains(stderr, "Warning: ignoring --model glm-5:cloud because cloud is disabled") {
@@ -375,7 +375,7 @@ func TestLaunchCmdModelFlagClearsDisabledCloudOverride(t *testing.T) {
if err != nil {
t.Fatalf("failed to reload integration config: %v", err)
}
if diff := cmp.Diff([]string{"llama3.2"}, saved.Models); diff != "" {
if diff := cmp.Diff([]string{"sample-model"}, saved.Models); diff != "" {
t.Fatalf("saved models mismatch (-want +got):\n%s", diff)
}
}
@@ -424,7 +424,7 @@ func TestLaunchCmdYes_AutoConfirmsLaunchPromptPath(t *testing.T) {
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
switch r.URL.Path {
case "/api/show":
fmt.Fprint(w, `{"model":"llama3.2"}`)
fmt.Fprint(w, `{"model":"sample-model"}`)
case "/api/status":
w.WriteHeader(http.StatusNotFound)
fmt.Fprint(w, `{"error":"not found"}`)
@@ -445,16 +445,16 @@ func TestLaunchCmdYes_AutoConfirmsLaunchPromptPath(t *testing.T) {
}
cmd := LaunchCmd(func(cmd *cobra.Command, args []string) error { return nil }, func(cmd *cobra.Command) {})
cmd.SetArgs([]string{"stubeditor", "--model", "llama3.2", "--yes"})
cmd.SetArgs([]string{"stubeditor", "--model", "sample-model", "--yes"})
if err := cmd.Execute(); err != nil {
t.Fatalf("launch command with --yes failed: %v", err)
}
if diff := cmp.Diff([][]string{{"llama3.2"}}, stub.edited); diff != "" {
if diff := cmp.Diff([][]string{{"sample-model"}}, stub.edited); diff != "" {
t.Fatalf("editor models mismatch (-want +got):\n%s", diff)
}
if stub.ranModel != "llama3.2" {
t.Fatalf("expected launch to run with llama3.2, got %q", stub.ranModel)
if stub.ranModel != "sample-model" {
t.Fatalf("expected launch to run with sample-model, got %q", stub.ranModel)
}
}
@@ -513,7 +513,7 @@ func TestLaunchCmdHeadlessWithoutYes_AllowsConfiguredLaunch(t *testing.T) {
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
switch r.URL.Path {
case "/api/show":
fmt.Fprint(w, `{"model":"llama3.2"}`)
fmt.Fprint(w, `{"model":"sample-model"}`)
case "/api/status":
w.WriteHeader(http.StatusNotFound)
fmt.Fprint(w, `{"error":"not found"}`)
@@ -534,15 +534,15 @@ func TestLaunchCmdHeadlessWithoutYes_AllowsConfiguredLaunch(t *testing.T) {
}
cmd := LaunchCmd(func(cmd *cobra.Command, args []string) error { return nil }, func(cmd *cobra.Command) {})
cmd.SetArgs([]string{"stubeditor", "--model", "llama3.2"})
cmd.SetArgs([]string{"stubeditor", "--model", "sample-model"})
err := cmd.Execute()
if err != nil {
t.Fatalf("expected launch command to succeed without --yes when an explicit model is provided, got %v", err)
}
if diff := compareStringSlices(stub.edited, [][]string{{"llama3.2"}}); diff != "" {
if diff := compareStringSlices(stub.edited, [][]string{{"sample-model"}}); diff != "" {
t.Fatalf("unexpected editor writes (-want +got):\n%s", diff)
}
if stub.ranModel != "llama3.2" {
if stub.ranModel != "sample-model" {
t.Fatalf("expected launch to run configured model, got %q", stub.ranModel)
}
}
@@ -551,7 +551,7 @@ func TestLaunchCmdIntegrationArgPromptsForModelWithSavedSelection(t *testing.T)
tmpDir := t.TempDir()
setLaunchTestHome(t, tmpDir)
if err := config.SaveIntegration("stubapp", []string{"llama3.2"}); err != nil {
if err := config.SaveIntegration("stubapp", []string{"sample-model"}); err != nil {
t.Fatalf("failed to seed saved config: %v", err)
}
@@ -560,7 +560,7 @@ func TestLaunchCmdIntegrationArgPromptsForModelWithSavedSelection(t *testing.T)
case "/api/experimental/model-recommendations":
fmt.Fprint(w, `{"recommendations":[]}`)
case "/api/tags":
fmt.Fprint(w, `{"models":[{"name":"llama3.2"},{"name":"qwen3:8b"}]}`)
fmt.Fprint(w, `{"models":[{"name":"sample-model"},{"name":"qwen3:8b"}]}`)
case "/api/show":
fmt.Fprint(w, `{"model":"qwen3:8b"}`)
default:
@@ -589,8 +589,8 @@ func TestLaunchCmdIntegrationArgPromptsForModelWithSavedSelection(t *testing.T)
t.Fatalf("launch command failed: %v", err)
}
if gotCurrent != "llama3.2" {
t.Fatalf("expected selector current model to be saved model llama3.2, got %q", gotCurrent)
if gotCurrent != "sample-model" {
t.Fatalf("expected selector current model to be saved model sample-model, got %q", gotCurrent)
}
if stub.ranModel != "qwen3:8b" {
t.Fatalf("expected launch to run selected model qwen3:8b, got %q", stub.ranModel)
@@ -611,14 +611,14 @@ func TestLaunchCmdHeadlessYes_IntegrationRequiresModelEvenWhenSaved(t *testing.T
withLauncherHooks(t)
withInteractiveSession(t, false)
if err := config.SaveIntegration("stubapp", []string{"llama3.2"}); err != nil {
if err := config.SaveIntegration("stubapp", []string{"sample-model"}); err != nil {
t.Fatalf("failed to seed saved config: %v", err)
}
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
switch r.URL.Path {
case "/api/show":
fmt.Fprint(w, `{"model":"llama3.2"}`)
fmt.Fprint(w, `{"model":"sample-model"}`)
default:
w.WriteHeader(http.StatusNotFound)
}
+133
View File
@@ -0,0 +1,133 @@
package launch
import (
"fmt"
"strings"
"github.com/ollama/ollama/internal/modelref"
)
var deprecatedLaunchModels = map[string]struct{}{
"codellama": {},
"qwen2.5": {},
"qwen2.5-coder": {},
"llama3": {},
"llama3.1": {},
"llama3.2": {},
"llama3.3": {},
"mistral": {},
"starcoder": {},
}
var deprecatedLaunchModelTags = map[string]map[string]struct{}{
"deepseek-r1": {
"": {},
"latest": {},
"1.5b": {},
"7b": {},
"8b": {},
"14b": {},
"32b": {},
},
}
var errDeprecatedLaunchModelDeclined = fmt.Errorf("%w: deprecated launch model declined", ErrCancelled)
func isDeprecatedLaunchModel(name string) bool {
family, tag := normalizedLaunchModelRef(name)
if _, ok := deprecatedLaunchModels[family]; ok {
return true
}
tags, ok := deprecatedLaunchModelTags[family]
if !ok {
return false
}
_, ok = tags[tag]
return ok
}
func deprecatedLaunchModelPrompt(name, label, commandName, cloudRec, localRec string) string {
if !isDeprecatedLaunchModel(name) {
return ""
}
if label = strings.TrimSpace(label); label == "" {
label = "ollama launch"
}
var b strings.Builder
fmt.Fprintf(&b, "%s does not work well with %s. ", name, label)
switch {
case cloudRec != "" && localRec != "":
fmt.Fprintf(&b, "Try an agent-capable model like %s or %s instead", cloudRec, localRec)
case cloudRec != "":
fmt.Fprintf(&b, "Try an agent-capable model like %s instead", cloudRec)
case localRec != "":
fmt.Fprintf(&b, "Try an agent-capable model like %s instead", localRec)
default:
b.WriteString("Try a newer recommended agent-capable model instead")
}
if command := launchReplacementCommand(commandName, firstNonEmpty(cloudRec, localRec)); command != "" {
fmt.Fprintf(&b, ":\n %s", command)
} else {
b.WriteString(".")
}
fmt.Fprintf(&b, "\n\nLaunch with %s anyway?", name)
return b.String()
}
func launchReplacementCommand(commandName, model string) string {
commandName = strings.TrimSpace(commandName)
model = strings.TrimSpace(model)
if commandName == "" || model == "" {
return ""
}
return fmt.Sprintf("ollama launch %s --model %s", commandName, model)
}
func firstNonEmpty(values ...string) string {
for _, value := range values {
if value != "" {
return value
}
}
return ""
}
func normalizedLaunchModelRef(name string) (string, string) {
name = strings.TrimSpace(strings.ToLower(name))
if name == "" {
return "", ""
}
if base, stripped := modelref.StripCloudSourceTag(name); stripped {
name = base
}
if idx := strings.LastIndex(name, "/"); idx >= 0 {
name = name[idx+1:]
}
tag := ""
if idx := strings.Index(name, ":"); idx >= 0 {
tag = strings.TrimSpace(name[idx+1:])
name = name[:idx]
}
return strings.TrimSpace(name), tag
}
func filterDeprecatedLaunchModelItems(items []ModelItem) []ModelItem {
filtered := items[:0]
for _, item := range items {
if !isDeprecatedLaunchModel(item.Name) {
filtered = append(filtered, item)
}
}
return filtered
}
func filterDeprecatedLaunchModelNames(models []string) []string {
filtered := models[:0]
for _, model := range models {
if !isDeprecatedLaunchModel(model) {
filtered = append(filtered, model)
}
}
return filtered
}
+68
View File
@@ -0,0 +1,68 @@
package launch
import (
"strings"
"testing"
)
func TestLaunchModelDeprecation(t *testing.T) {
tests := []struct {
name string
deprecated bool
}{
{name: "qwen2.5", deprecated: true},
{name: "qwen2.5:14b", deprecated: true},
{name: "qwen2.5-coder:32b", deprecated: true},
{name: "library/qwen2.5-coder:7b", deprecated: true},
{name: "llama3", deprecated: true},
{name: "llama3.1:8b", deprecated: true},
{name: "llama3.2:latest", deprecated: true},
{name: "llama3.3:70b", deprecated: true},
{name: "llama3.2:cloud", deprecated: true},
{name: "codellama", deprecated: true},
{name: "codellama:13b-code", deprecated: true},
{name: "library/codellama:7b", deprecated: true},
{name: "starcoder", deprecated: true},
{name: "starcoder:15b", deprecated: true},
{name: "mistral", deprecated: true},
{name: "mistral:7b", deprecated: true},
{name: "deepseek-r1", deprecated: true},
{name: "deepseek-r1:latest", deprecated: true},
{name: "deepseek-r1:1.5b", deprecated: true},
{name: "deepseek-r1:7b", deprecated: true},
{name: "deepseek-r1:8b", deprecated: true},
{name: "deepseek-r1:14b", deprecated: true},
{name: "deepseek-r1:32b", deprecated: true},
{name: "deepseek-r1:32b-cloud", deprecated: true},
{name: "qwen3.5", deprecated: false},
{name: "qwen3-coder:30b", deprecated: false},
{name: "gemma4", deprecated: false},
{name: "my-qwen2.5-coder", deprecated: false},
{name: "llama3.2-inspired", deprecated: false},
{name: "codellama-inspired", deprecated: false},
{name: "starcoder2:15b", deprecated: false},
{name: "mixtral:8x7b", deprecated: false},
{name: "deepseek-r1:70b", deprecated: false},
{name: "deepseek-r1:671b", deprecated: false},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := isDeprecatedLaunchModel(tt.name); got != tt.deprecated {
t.Fatalf("isDeprecatedLaunchModel(%q) = %v, want %v", tt.name, got, tt.deprecated)
}
})
}
}
func TestDeprecatedLaunchModelErrorMentionsRecommendedModels(t *testing.T) {
prompt := deprecatedLaunchModelPrompt("qwen2.5-coder:32b", "Codex", "codex", "recommended-cloud:cloud", "recommended-local")
if prompt == "" {
t.Fatal("expected deprecated model prompt")
}
for _, want := range []string{"qwen2.5-coder:32b does not work well with Codex", "recommended-cloud:cloud", "recommended-local", "ollama launch codex --model recommended-cloud:cloud", "Launch with qwen2.5-coder:32b anyway?"} {
if !strings.Contains(prompt, want) {
t.Fatalf("prompt %q does not contain %q", prompt, want)
}
}
}
+256 -21
View File
@@ -14,6 +14,7 @@ import (
"strconv"
"strings"
"golang.org/x/mod/semver"
"gopkg.in/yaml.v3"
"github.com/ollama/ollama/api"
@@ -23,7 +24,11 @@ import (
)
const (
hermesInstallScript = "curl -fsSL https://raw.githubusercontent.com/NousResearch/hermes-agent/main/scripts/install.sh | bash -s -- --skip-setup"
// https://github.com/NousResearch/hermes-agent/releases/tag/v2026.6.5
hermesDesktopMinVersion = "v0.16.0"
hermesInstallScript = "curl -fsSL https://hermes-agent.nousresearch.com/install.sh | bash -s -- --skip-setup"
hermesWindowsInstallURL = "https://hermes-agent.nousresearch.com/install.ps1"
hermesWindowsInstallCmd = "& ([scriptblock]::Create((irm " + hermesWindowsInstallURL + "))) -SkipSetup"
hermesProviderName = "Ollama"
hermesProviderKey = "ollama-launch"
hermesLegacyKey = "ollama"
@@ -81,6 +86,177 @@ func (h *Hermes) Run(_ string, _ []LaunchModel, args []string) error {
return hermesAttachedCommand(bin, args...).Run()
}
type HermesDesktop struct {
Hermes
}
func (h *HermesDesktop) String() string { return "Hermes Desktop" }
func (h *HermesDesktop) Run(_ string, _ []LaunchModel, args []string) error {
bin, err := h.binary()
if err != nil {
return err
}
if err := h.ensureHermesDesktopMinVersion(bin); err != nil {
return err
}
return hermesAttachedCommand(bin, h.launchArgs(args)...).Run()
}
func (h *HermesDesktop) ensureHermesDesktopMinVersion(bin string) error {
if hermesGOOS == "windows" {
return nil
}
version := hermesVersionOf(bin)
if version == "" {
return nil
}
if semver.Compare(version, hermesDesktopMinVersion) >= 0 {
return nil
}
fmt.Fprintf(os.Stderr, "%sHermes %s is older than the minimum version (%s) for `hermes desktop`; updating...%s\n", ansiGray, version, hermesDesktopMinVersion, ansiReset)
if err := hermesAttachedCommand(bin, "update").Run(); err != nil {
return fmt.Errorf("failed to update hermes to %s or newer: %w", hermesDesktopMinVersion, err)
}
return nil
}
func hermesVersionOf(bin string) string {
out, err := hermesCommand(bin, "--version").Output()
if err != nil {
return ""
}
firstLine := strings.SplitN(strings.TrimSpace(string(out)), "\n", 2)[0]
return parseHermesVersion(firstLine)
}
func parseHermesVersion(firstLine string) string {
for _, field := range strings.Fields(firstLine) {
if semver.IsValid(field) {
return field
}
}
return ""
}
func (h *HermesDesktop) Onboard() error {
return config.MarkIntegrationOnboarded("hermes-desktop")
}
func (h *HermesDesktop) launchArgs(args []string) []string {
launchArgs := []string{"desktop"}
if h.shouldSkipDesktopBuild(args) {
launchArgs = append(launchArgs, "--skip-build")
}
return append(launchArgs, args...)
}
func (h *HermesDesktop) shouldSkipDesktopBuild(args []string) bool {
if hermesDesktopHasFlag(args, "--skip-build", "--source", "--build-only", "--help", "-h") {
return false
}
return h.packagedAppExists()
}
func (h *HermesDesktop) packagedAppExists() bool {
for _, root := range hermesDesktopReleaseRoots() {
for _, candidate := range hermesDesktopPackagedExecutableCandidates(root) {
if _, err := os.Stat(candidate); err == nil {
return true
}
}
}
return false
}
// These roots mirror Hermes' own install layout:
// install.sh uses ~/.hermes/hermes-agent for user installs and
// /usr/local/lib/hermes-agent for new Linux root installs; install.ps1
// and the bootstrap installer use %LOCALAPPDATA%\hermes\hermes-agent on
// Windows. HERMES_HOME and HERMES_INSTALL_DIR are installer-supported
// overrides.
func hermesDesktopReleaseRoots() []string {
var installRoots []string
add := func(path string) {
path = strings.TrimSpace(path)
if path == "" {
return
}
installRoots = append(installRoots, filepath.Clean(path))
}
if installDir := strings.TrimSpace(os.Getenv("HERMES_INSTALL_DIR")); installDir != "" {
add(installDir)
}
if hermesHome := strings.TrimSpace(os.Getenv("HERMES_HOME")); hermesHome != "" {
add(filepath.Join(hermesHome, "hermes-agent"))
}
home, err := hermesUserHome()
if err == nil {
switch hermesGOOS {
case "windows":
if localAppData := strings.TrimSpace(os.Getenv("LOCALAPPDATA")); localAppData != "" {
add(filepath.Join(localAppData, "hermes", "hermes-agent"))
}
add(filepath.Join(home, ".hermes", "hermes-agent"))
default:
add(filepath.Join(home, ".hermes", "hermes-agent"))
if hermesGOOS == "linux" {
add(filepath.Join(string(filepath.Separator), "usr", "local", "lib", "hermes-agent"))
}
}
}
seen := make(map[string]bool, len(installRoots))
releaseRoots := make([]string, 0, len(installRoots))
for _, root := range installRoots {
releaseRoot := filepath.Join(root, "apps", "desktop", "release")
if seen[releaseRoot] {
continue
}
seen[releaseRoot] = true
releaseRoots = append(releaseRoots, releaseRoot)
}
return releaseRoots
}
func hermesDesktopPackagedExecutableCandidates(releaseRoot string) []string {
switch hermesGOOS {
case "darwin":
matches, err := filepath.Glob(filepath.Join(releaseRoot, "mac*", "Hermes.app", "Contents", "MacOS", "Hermes"))
if err != nil {
return nil
}
return matches
case "windows":
return []string{
filepath.Join(releaseRoot, "win-unpacked", "Hermes.exe"),
filepath.Join(releaseRoot, "win-ia32-unpacked", "Hermes.exe"),
filepath.Join(releaseRoot, "win-arm64-unpacked", "Hermes.exe"),
}
default:
return []string{
filepath.Join(releaseRoot, "linux-unpacked", "hermes"),
filepath.Join(releaseRoot, "linux-unpacked", "Hermes"),
}
}
}
func hermesDesktopHasFlag(args []string, names ...string) bool {
for _, arg := range args {
if arg == "--" {
return false
}
for _, name := range names {
if arg == name {
return true
}
}
}
return false
}
func (h *Hermes) Paths() []string {
configPath, err := hermesConfigPath()
if err != nil {
@@ -183,22 +359,24 @@ func (h *Hermes) installed() bool {
}
func (h *Hermes) ensureInstalled() error {
return h.ensureInstalledFor("hermes")
}
func (h *Hermes) ensureInstalledFor(command string) error {
if h.installed() {
return nil
}
if hermesGOOS == "windows" {
return hermesWindowsHint()
}
var missing []string
for _, dep := range []string{"bash", "curl", "git"} {
if _, err := hermesLookPath(dep); err != nil {
missing = append(missing, dep)
if hermesGOOS != "windows" {
for _, dep := range []string{"bash", "curl", "git"} {
if _, err := hermesLookPath(dep); err != nil {
missing = append(missing, dep)
}
}
}
if len(missing) > 0 {
return fmt.Errorf("Hermes is not installed and required dependencies are missing\n\nInstall the following first:\n %s\n\nThen re-run:\n ollama launch hermes", strings.Join(missing, "\n "))
return fmt.Errorf("Hermes is not installed and required dependencies are missing\n\nInstall the following first:\n %s\n\nThen re-run:\n ollama launch %s", strings.Join(missing, "\n "), command)
}
ok, err := ConfirmPrompt("Hermes is not installed. Install now?")
@@ -210,7 +388,7 @@ func (h *Hermes) ensureInstalled() error {
}
fmt.Fprintf(os.Stderr, "\nInstalling Hermes...\n")
if err := hermesAttachedCommand("bash", "-lc", hermesInstallScript).Run(); err != nil {
if err := h.runInstallScript(); err != nil {
return fmt.Errorf("failed to install hermes: %w", err)
}
@@ -222,6 +400,13 @@ func (h *Hermes) ensureInstalled() error {
return nil
}
func (h *Hermes) runInstallScript() error {
if hermesGOOS == "windows" {
return hermesAttachedCommand("powershell.exe", "-NoProfile", "-ExecutionPolicy", "Bypass", "-Command", hermesWindowsInstallCmd).Run()
}
return hermesAttachedCommand("bash", "-lc", hermesInstallScript).Run()
}
func (h *Hermes) listModels(defaultModel string) []string {
client := hermesOllamaClient()
resp, err := client.List(context.Background())
@@ -259,7 +444,12 @@ func (h *Hermes) binary() (string, error) {
}
if hermesGOOS == "windows" {
return "", hermesWindowsHint()
for _, fallback := range hermesWindowsBinaryFallbacks() {
if _, err := os.Stat(fallback); err == nil {
return fallback, nil
}
}
return "", fmt.Errorf("hermes is not installed")
}
home, err := hermesUserHome()
@@ -274,12 +464,63 @@ func (h *Hermes) binary() (string, error) {
return "", fmt.Errorf("hermes is not installed")
}
func hermesConfigPath() (string, error) {
func hermesWindowsBinaryFallbacks() []string {
var roots []string
add := func(root string) {
root = strings.TrimSpace(root)
if root != "" {
roots = append(roots, filepath.Clean(root))
}
}
add(os.Getenv("HERMES_HOME"))
add(os.Getenv("LOCALAPPDATA"))
if home, err := hermesUserHome(); err == nil {
add(filepath.Join(home, "AppData", "Local"))
}
seen := make(map[string]bool, len(roots))
var fallbacks []string
for _, root := range roots {
if seen[root] {
continue
}
seen[root] = true
fallbacks = append(fallbacks, filepath.Join(root, "hermes-agent", "venv", "Scripts", "hermes.exe"))
if filepath.Base(root) != "hermes" {
fallbacks = append(fallbacks, filepath.Join(root, "hermes", "hermes-agent", "venv", "Scripts", "hermes.exe"))
}
}
return fallbacks
}
func hermesHomePath() (string, error) {
if hermesHome := strings.TrimSpace(os.Getenv("HERMES_HOME")); hermesHome != "" {
return filepath.Clean(hermesHome), nil
}
if hermesGOOS == "windows" {
if localAppData := strings.TrimSpace(os.Getenv("LOCALAPPDATA")); localAppData != "" {
return filepath.Join(localAppData, "hermes"), nil
}
home, err := hermesUserHome()
if err != nil {
return "", err
}
return filepath.Join(home, "AppData", "Local", "hermes"), nil
}
home, err := hermesUserHome()
if err != nil {
return "", err
}
return filepath.Join(home, ".hermes", "config.yaml"), nil
return filepath.Join(home, ".hermes"), nil
}
func hermesConfigPath() (string, error) {
home, err := hermesHomePath()
if err != nil {
return "", err
}
return filepath.Join(home, "config.yaml"), nil
}
func hermesBaseURL() string {
@@ -287,11 +528,11 @@ func hermesBaseURL() string {
}
func hermesEnvPath() (string, error) {
home, err := hermesUserHome()
home, err := hermesHomePath()
if err != nil {
return "", err
}
return filepath.Join(home, ".hermes", ".env"), nil
return filepath.Join(home, ".env"), nil
}
func (h *Hermes) runGatewaySetupPreflight(args []string, runSetup func() error) error {
@@ -671,9 +912,3 @@ func hermesAttachedCommand(name string, args ...string) *exec.Cmd {
cmd.Stderr = os.Stderr
return cmd
}
func hermesWindowsHint() error {
return fmt.Errorf("Hermes on Windows requires WSL2. Install WSL with: wsl --install\n" +
"Then run 'ollama launch hermes' from inside your WSL shell.\n" +
"Docs: https://hermes-agent.nousresearch.com/docs/getting-started/installation/")
}
+386 -13
View File
@@ -8,6 +8,7 @@ import (
"os"
"path/filepath"
"runtime"
"slices"
"strings"
"testing"
@@ -65,6 +66,20 @@ func clearHermesMessagingEnvVars(t *testing.T) {
}
}
func clearHermesDesktopPackageEnvVars(t *testing.T) {
t.Helper()
for _, key := range []string{"HERMES_INSTALL_DIR", "HERMES_HOME", "LOCALAPPDATA"} {
if value, ok := os.LookupEnv(key); ok {
t.Setenv(key, value)
} else {
t.Setenv(key, "")
}
if err := os.Unsetenv(key); err != nil {
t.Fatalf("unset %s: %v", key, err)
}
}
}
func TestHermesIntegration(t *testing.T) {
h := &Hermes{}
@@ -408,19 +423,36 @@ func TestHermesConfigureMigratesLegacyManagedAliases(t *testing.T) {
func TestHermesPathsUsesLocalConfigPathForNativeWindowsHermes(t *testing.T) {
tmpDir := t.TempDir()
winHome := filepath.Join(tmpDir, "winhome")
localAppData := filepath.Join(tmpDir, "LocalAppData")
setTestHome(t, winHome)
withHermesPlatform(t, "windows")
withHermesUserHome(t, winHome)
t.Setenv("PATH", tmpDir)
t.Setenv("LOCALAPPDATA", localAppData)
writeFakeBinary(t, tmpDir, "hermes")
got := (&Hermes{}).Paths()
want := filepath.Join(winHome, ".hermes", "config.yaml")
want := filepath.Join(localAppData, "hermes", "config.yaml")
if len(got) != 1 || got[0] != want {
t.Fatalf("expected local config path %q, got %v", want, got)
}
}
func TestHermesPathsUsesHermesHomeOverride(t *testing.T) {
tmpDir := t.TempDir()
hermesHome := filepath.Join(tmpDir, "custom-hermes-home")
setTestHome(t, filepath.Join(tmpDir, "home"))
withHermesPlatform(t, "windows")
t.Setenv("HERMES_HOME", hermesHome)
t.Setenv("LOCALAPPDATA", filepath.Join(tmpDir, "LocalAppData"))
got := (&Hermes{}).Paths()
want := filepath.Join(hermesHome, "config.yaml")
if len(got) != 1 || got[0] != want {
t.Fatalf("expected HERMES_HOME config path %q, got %v", want, got)
}
}
func TestHermesCurrentModelRequiresHealthyManagedConfig(t *testing.T) {
tmpDir := t.TempDir()
setTestHome(t, tmpDir)
@@ -565,6 +597,314 @@ func TestHermesRunPassthroughArgs(t *testing.T) {
}
}
func writeHermesDesktopPackage(t *testing.T, home string) {
t.Helper()
writeHermesDesktopExecutable(t,
filepath.Join(home, ".hermes", "hermes-agent", "apps", "desktop", "release"),
hermesDesktopTestExecutableRelativePath(hermesGOOS),
)
}
func writeHermesDesktopExecutable(t *testing.T, releaseRoot, relative string) {
t.Helper()
path := filepath.Join(releaseRoot, relative)
if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(path, []byte("#!/bin/sh\nexit 0\n"), 0o755); err != nil {
t.Fatal(err)
}
}
func hermesDesktopTestExecutableRelativePath(goos string) string {
switch goos {
case "darwin":
return filepath.Join("mac-arm64", "Hermes.app", "Contents", "MacOS", "Hermes")
case "windows":
return filepath.Join("win-unpacked", "Hermes.exe")
default:
return filepath.Join("linux-unpacked", "hermes")
}
}
func writeHermesDesktopTestBinary(t *testing.T, dir string) {
t.Helper()
bin := filepath.Join(dir, "hermes")
if err := os.WriteFile(bin, []byte("#!/bin/sh\nif [ \"$1\" = \"--version\" ]; then\n printf 'Hermes Agent v0.16.0 (2026.6.5)\\n'\n exit 0\nfi\nprintf '[%s]\\n' \"$*\" >> \"$HOME/hermes-invocations.log\"\n"), 0o755); err != nil {
t.Fatal(err)
}
}
func readHermesDesktopInvocations(t *testing.T, home string) string {
t.Helper()
data, err := os.ReadFile(filepath.Join(home, "hermes-invocations.log"))
if err != nil {
t.Fatal(err)
}
return strings.TrimSpace(string(data))
}
func TestHermesDesktopRun(t *testing.T) {
if runtime.GOOS == "windows" {
t.Skip("uses a POSIX shell test binary")
}
tests := []struct {
name string
goos string
args []string
hasPackage bool
clearPkgEnv bool
want string
}{
{
name: "desktop subcommand",
goos: "darwin",
args: []string{"--foreground"},
clearPkgEnv: true,
want: "[desktop --foreground]",
},
{
name: "skip build when packaged app exists",
goos: runtime.GOOS,
args: []string{"--cwd", "/tmp/project"},
hasPackage: true,
want: "[desktop --skip-build --cwd /tmp/project]",
},
{
name: "explicit skip build",
goos: runtime.GOOS,
args: []string{"--skip-build"},
hasPackage: true,
want: "[desktop --skip-build]",
},
{
name: "source mode",
goos: runtime.GOOS,
args: []string{"--source"},
hasPackage: true,
want: "[desktop --source]",
},
{
name: "build only",
goos: runtime.GOOS,
args: []string{"--build-only"},
hasPackage: true,
want: "[desktop --build-only]",
},
{
name: "help",
goos: runtime.GOOS,
args: []string{"--help"},
hasPackage: true,
want: "[desktop --help]",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
tmpDir := t.TempDir()
setTestHome(t, tmpDir)
withLauncherHooks(t)
withInteractiveSession(t, true)
withHermesPlatform(t, tt.goos)
clearHermesMessagingEnvVars(t)
if tt.clearPkgEnv {
clearHermesDesktopPackageEnvVars(t)
}
t.Setenv("PATH", tmpDir+string(os.PathListSeparator)+os.Getenv("PATH"))
if tt.hasPackage {
writeHermesDesktopPackage(t, tmpDir)
}
writeHermesDesktopTestBinary(t, tmpDir)
DefaultConfirmPrompt = func(prompt string, options ConfirmOptions) (bool, error) {
t.Fatalf("did not expect messaging prompt during desktop launch: %s", prompt)
return false, nil
}
if err := (&HermesDesktop{}).Run("", nil, tt.args); err != nil {
t.Fatalf("Run returned error: %v", err)
}
if got := readHermesDesktopInvocations(t, tmpDir); got != tt.want {
t.Fatalf("expected %q, got %q", tt.want, got)
}
})
}
}
func writeHermesVersionedTestBinary(t *testing.T, dir, version string) {
t.Helper()
script := "#!/bin/sh\n" +
"case \"$1\" in\n" +
" --version)\n" +
" printf 'Hermes Agent " + version + " (test)\\n'\n" +
" ;;\n" +
" update)\n" +
" printf 'update\\n' >> \"$HOME/hermes-update.log\"\n" +
" ;;\n" +
" *)\n" +
" printf '[%s]\\n' \"$*\" >> \"$HOME/hermes-invocations.log\"\n" +
" ;;\n" +
"esac\n"
bin := filepath.Join(dir, "hermes")
if err := os.WriteFile(bin, []byte(script), 0o755); err != nil {
t.Fatal(err)
}
}
func TestParseHermesVersion(t *testing.T) {
tests := []struct {
name string
input string
want string
}{
{"standard release", "Hermes Agent v0.16.0 (2026.6.5)", "v0.16.0"},
{"newer release", "Hermes Agent v0.17.0 (2026.6.19)", "v0.17.0"},
{"older release", "Hermes Agent v0.15.1 (2026.5.29)", "v0.15.1"},
{"prerelease", "Hermes Agent v0.16.0-rc1 (2026.6.5)", "v0.16.0-rc1"},
{"no version token", "Hermes Agent", ""},
{"empty", "", ""},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := parseHermesVersion(tt.input); got != tt.want {
t.Fatalf("parseHermesVersion(%q) = %q, want %q", tt.input, got, tt.want)
}
})
}
}
func TestHermesDesktopRun_UpdatesCliOlderThanMinVersion(t *testing.T) {
if runtime.GOOS == "windows" {
t.Skip("uses a POSIX shell test binary")
}
tmpDir := t.TempDir()
setTestHome(t, tmpDir)
withLauncherHooks(t)
withInteractiveSession(t, true)
withHermesPlatform(t, runtime.GOOS)
clearHermesMessagingEnvVars(t)
clearHermesDesktopPackageEnvVars(t)
t.Setenv("PATH", tmpDir+string(os.PathListSeparator)+os.Getenv("PATH"))
writeHermesVersionedTestBinary(t, tmpDir, "v0.15.1")
DefaultConfirmPrompt = func(prompt string, options ConfirmOptions) (bool, error) {
t.Fatalf("did not expect messaging prompt during desktop launch: %s", prompt)
return false, nil
}
if err := (&HermesDesktop{}).Run("", nil, []string{"--foreground"}); err != nil {
t.Fatalf("Run returned error: %v", err)
}
updateLog, err := os.ReadFile(filepath.Join(tmpDir, "hermes-update.log"))
if err != nil {
t.Fatalf("expected hermes update to run for an older CLI: %v", err)
}
if strings.TrimSpace(string(updateLog)) != "update" {
t.Fatalf("expected update log 'update', got %q", updateLog)
}
if got := readHermesDesktopInvocations(t, tmpDir); got != "[desktop --foreground]" {
t.Fatalf("expected desktop launch after update, got %q", got)
}
}
func TestHermesDesktopRun_SkipsMinVersionCheckOnWindows(t *testing.T) {
if runtime.GOOS == "windows" {
t.Skip("uses a POSIX shell test binary")
}
tmpDir := t.TempDir()
setTestHome(t, tmpDir)
withLauncherHooks(t)
withInteractiveSession(t, true)
withHermesPlatform(t, "windows")
clearHermesMessagingEnvVars(t)
clearHermesDesktopPackageEnvVars(t)
t.Setenv("PATH", tmpDir+string(os.PathListSeparator)+os.Getenv("PATH"))
writeHermesVersionedTestBinary(t, tmpDir, "v0.15.1")
DefaultConfirmPrompt = func(prompt string, options ConfirmOptions) (bool, error) {
t.Fatalf("did not expect messaging prompt during desktop launch: %s", prompt)
return false, nil
}
if err := (&HermesDesktop{}).Run("", nil, []string{"--foreground"}); err != nil {
t.Fatalf("Run returned error: %v", err)
}
if _, err := os.Stat(filepath.Join(tmpDir, "hermes-update.log")); err == nil {
t.Fatal("expected hermes update NOT to run on Windows, but hermes-update.log exists")
}
if got := readHermesDesktopInvocations(t, tmpDir); got != "[desktop --foreground]" {
t.Fatalf("expected desktop launch without update, got %q", got)
}
}
func TestHermesDesktopRun_DoesNotUpdateCliAtOrAboveMinVersion(t *testing.T) {
if runtime.GOOS == "windows" {
t.Skip("uses a POSIX shell test binary")
}
tmpDir := t.TempDir()
setTestHome(t, tmpDir)
withLauncherHooks(t)
withInteractiveSession(t, true)
withHermesPlatform(t, runtime.GOOS)
clearHermesMessagingEnvVars(t)
clearHermesDesktopPackageEnvVars(t)
t.Setenv("PATH", tmpDir+string(os.PathListSeparator)+os.Getenv("PATH"))
writeHermesVersionedTestBinary(t, tmpDir, "v0.17.0")
DefaultConfirmPrompt = func(prompt string, options ConfirmOptions) (bool, error) {
t.Fatalf("did not expect messaging prompt during desktop launch: %s", prompt)
return false, nil
}
if err := (&HermesDesktop{}).Run("", nil, []string{"--foreground"}); err != nil {
t.Fatalf("Run returned error: %v", err)
}
if _, err := os.Stat(filepath.Join(tmpDir, "hermes-update.log")); err == nil {
t.Fatal("expected hermes update NOT to run for a current CLI, but hermes-update.log exists")
}
if got := readHermesDesktopInvocations(t, tmpDir); got != "[desktop --foreground]" {
t.Fatalf("expected desktop launch without update, got %q", got)
}
}
func TestHermesDesktopRunUsesWindowsLocalAppDataPackage(t *testing.T) {
tmpDir := t.TempDir()
setTestHome(t, tmpDir)
withHermesPlatform(t, "windows")
t.Setenv("LOCALAPPDATA", filepath.Join(tmpDir, "LocalAppData"))
writeHermesDesktopExecutable(t,
filepath.Join(tmpDir, "LocalAppData", "hermes", "hermes-agent", "apps", "desktop", "release"),
hermesDesktopTestExecutableRelativePath("windows"),
)
got := (&HermesDesktop{}).launchArgs([]string{"--cwd", `C:\Users\me\project`})
want := []string{"desktop", "--skip-build", "--cwd", `C:\Users\me\project`}
if diff := compareStrings(got, want); diff != "" {
t.Fatalf("Hermes Desktop launch args mismatch: %s", diff)
}
}
func TestHermesDesktopReleaseRootsIncludeLinuxRootInstall(t *testing.T) {
tmpDir := t.TempDir()
setTestHome(t, tmpDir)
withHermesPlatform(t, "linux")
got := hermesDesktopReleaseRoots()
want := filepath.Join(string(filepath.Separator), "usr", "local", "lib", "hermes-agent", "apps", "desktop", "release")
if !slices.Contains(got, want) {
t.Fatalf("expected Linux root install release path %q in %v", want, got)
}
}
func TestHermesRun_PromptsForMessagingSetupBeforeDefaultLaunch(t *testing.T) {
if runtime.GOOS == "windows" {
t.Skip("uses a POSIX shell test binary")
@@ -943,26 +1283,59 @@ func TestHermesMessagingConfiguredRecognizesSupportedGatewayVars(t *testing.T) {
}
}
func TestHermesEnsureInstalledWindowsShowsWSLGuidance(t *testing.T) {
func TestHermesEnsureInstalledWindowsRunsPowerShellInstaller(t *testing.T) {
if runtime.GOOS == "windows" {
t.Skip("uses a POSIX shell test binary")
}
tmpDir := t.TempDir()
setTestHome(t, tmpDir)
withLauncherHooks(t)
withHermesPlatform(t, "windows")
t.Setenv("PATH", tmpDir)
t.Setenv("LOCALAPPDATA", filepath.Join(tmpDir, "AppData", "Local"))
powershell := filepath.Join(tmpDir, "powershell.exe")
script := fmt.Sprintf(`#!/bin/sh
printf '%%s\n' "$*" >> %q
/bin/mkdir -p %q
/bin/cat > %q <<'EOS'
#!/bin/sh
exit 0
EOS
/bin/chmod +x %q
exit 0
`,
filepath.Join(tmpDir, "powershell.log"),
filepath.Dir(filepath.Join(tmpDir, "AppData", "Local", "hermes", "hermes-agent", "venv", "Scripts", "hermes.exe")),
filepath.Join(tmpDir, "AppData", "Local", "hermes", "hermes-agent", "venv", "Scripts", "hermes.exe"),
filepath.Join(tmpDir, "AppData", "Local", "hermes", "hermes-agent", "venv", "Scripts", "hermes.exe"),
)
if err := os.WriteFile(powershell, []byte(script), 0o755); err != nil {
t.Fatal(err)
}
DefaultConfirmPrompt = func(prompt string, options ConfirmOptions) (bool, error) {
if prompt != "Hermes is not installed. Install now?" {
t.Fatalf("unexpected install prompt %q", prompt)
}
return true, nil
}
h := &Hermes{}
err := h.ensureInstalled()
if err == nil {
t.Fatal("expected WSL guidance error")
if err := h.ensureInstalled(); err != nil {
t.Fatalf("ensureInstalled returned error: %v", err)
}
msg := err.Error()
if !strings.Contains(msg, "wsl --install") {
t.Fatalf("expected install command in guidance, got %v", err)
data, err := os.ReadFile(filepath.Join(tmpDir, "powershell.log"))
if err != nil {
t.Fatal(err)
}
if !strings.Contains(msg, "hermes-agent.nousresearch.com") {
t.Fatalf("expected docs link in guidance, got %v", err)
}
if strings.Contains(msg, "hermes is not installed") {
t.Fatalf("guidance should not lead with 'hermes is not installed', got %v", err)
logs := string(data)
for _, want := range []string{"-NoProfile", "-ExecutionPolicy", "Bypass", "-Command", hermesWindowsInstallURL, "-SkipSetup"} {
if !strings.Contains(logs, want) {
t.Fatalf("expected PowerShell installer args to contain %q, got logs:\n%s", want, logs)
}
}
}
+98 -14
View File
@@ -58,12 +58,15 @@ func TestIntegrationLookup(t *testing.T) {
{"claude desktop", "claude-desktop", true, "Claude Desktop"},
{"claude desktop alias", "claude-app", true, "Claude Desktop"},
{"codex", "codex", true, "Codex"},
{"codex app", "codex-app", true, "Codex App"},
{"codex app desktop alias", "codex-desktop", true, "Codex App"},
{"codex app gui alias", "codex-gui", true, "Codex App"},
{"chatgpt", "chatgpt", true, "ChatGPT"},
{"codex app legacy alias", "codex-app", true, "ChatGPT"},
{"codex app desktop alias", "codex-desktop", true, "ChatGPT"},
{"codex app gui alias", "codex-gui", true, "ChatGPT"},
{"hermes desktop", "hermes-desktop", true, "Hermes Desktop"},
{"kimi", "kimi", true, "Kimi Code CLI"},
{"droid", "droid", true, "Droid"},
{"opencode", "opencode", true, "OpenCode"},
{"omp", "omp", true, "OMP"},
{"pool", "pool", true, "Pool"},
{"unknown integration", "unknown", false, ""},
{"empty string", "", false, ""},
@@ -83,7 +86,7 @@ func TestIntegrationLookup(t *testing.T) {
}
func TestIntegrationRegistry(t *testing.T) {
expectedIntegrations := []string{"claude", "claude-desktop", "cline", "codex", "codex-app", "kimi", "droid", "opencode", "hermes", "pool"}
expectedIntegrations := []string{"claude", "claude-desktop", "cline", "codex", "chatgpt", "kimi", "droid", "opencode", "omp", "hermes", "hermes-desktop", "pool", "qwen"}
for _, name := range expectedIntegrations {
t.Run(name, func(t *testing.T) {
r, ok := integrations[name]
@@ -97,6 +100,30 @@ func TestIntegrationRegistry(t *testing.T) {
}
}
func TestChatGPTMigratesLegacyCodexAppLaunchConfig(t *testing.T) {
setTestHome(t, t.TempDir())
if err := config.SaveIntegration(codexAppIntegrationName, []string{"qwen3.5"}); err != nil {
t.Fatal(err)
}
if err := config.MarkIntegrationOnboarded(codexAppIntegrationName); err != nil {
t.Fatal(err)
}
got, err := loadStoredIntegrationConfig(chatGPTIntegrationName)
if err != nil {
t.Fatalf("loadStoredIntegrationConfig returned error: %v", err)
}
if diff := compareStrings(got.Models, []string{"qwen3.5"}); diff != "" {
t.Fatalf("migrated models mismatch: %s", diff)
}
if !got.Onboarded {
t.Fatal("migrated integration should remain onboarded")
}
if _, err := config.LoadIntegration(chatGPTIntegrationName); err != nil {
t.Fatalf("canonical ChatGPT config was not written: %v", err)
}
}
func TestHiddenIntegrationsExcludedFromVisibleLists(t *testing.T) {
for _, info := range ListIntegrationInfos() {
switch info.Name {
@@ -1079,6 +1106,51 @@ func TestShowOrPullWithPolicy_CloudModelNotFound_FailsEarlyForAllPolicies(t *tes
}
}
func TestShowOrPullWithPolicy_CloudModelShowUnavailableAllowsSelection(t *testing.T) {
oldHook := DefaultConfirmPrompt
DefaultConfirmPrompt = func(prompt string, options ConfirmOptions) (bool, error) {
t.Fatal("confirm prompt should not be called for explicit cloud models")
return false, nil
}
defer func() { DefaultConfirmPrompt = oldHook }()
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
switch r.URL.Path {
case "/api/show":
w.WriteHeader(http.StatusInternalServerError)
fmt.Fprintf(w, `{"error":"temporary failure"}`)
case "/api/status":
w.WriteHeader(http.StatusInternalServerError)
fmt.Fprintf(w, `{"error":"temporary failure"}`)
case "/api/pull":
t.Fatal("pull should not be called for explicit cloud models")
default:
w.WriteHeader(http.StatusNotFound)
}
}))
defer srv.Close()
u, _ := url.Parse(srv.URL)
client := api.NewClient(u, srv.Client())
if err := showOrPullWithPolicy(context.Background(), client, "glm-5.1:cloud", missingModelFail, true); err != nil {
t.Fatalf("showOrPullWithPolicy returned error: %v", err)
}
}
func TestShowOrPullWithPolicy_CloudModelShowUnreachableAllowsSelection(t *testing.T) {
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
t.Fatalf("unexpected request after server close: %s %s", r.Method, r.URL.Path)
}))
u, _ := url.Parse(srv.URL)
client := api.NewClient(u, srv.Client())
srv.Close()
if err := showOrPullWithPolicy(context.Background(), client, "glm-5.1:cloud", missingModelFail, true); err != nil {
t.Fatalf("showOrPullWithPolicy returned error: %v", err)
}
}
func TestShowOrPullWithPolicy_CloudModelDisabled_FailsWithCloudDisabledError(t *testing.T) {
oldHook := DefaultConfirmPrompt
DefaultConfirmPrompt = func(prompt string, options ConfirmOptions) (bool, error) {
@@ -1741,9 +1813,9 @@ func TestIntegration_InstallHint(t *testing.T) {
wantURL: "https://developers.openai.com/codex/cli/",
},
{
name: "codex app has hint",
input: "codex-app",
wantURL: "https://developers.openai.com/codex/quickstart",
name: "chatgpt has hint",
input: "chatgpt",
wantURL: "https://chatgpt.com/download",
},
{
name: "openclaw has hint",
@@ -1829,7 +1901,7 @@ func TestListIntegrationInfos(t *testing.T) {
if codexAppSupported() != nil {
filtered := make([]string, 0, len(want))
for _, name := range want {
if name != "codex-app" {
if name != "chatgpt" {
filtered = append(filtered, name)
}
}
@@ -1846,9 +1918,9 @@ func TestListIntegrationInfos(t *testing.T) {
for _, info := range infos {
got = append(got, info.Name)
}
wantPrefix := []string{"claude", "codex-app", "hermes", "openclaw"}
wantPrefix := []string{"claude", "chatgpt", "hermes", "openclaw", "opencode", "hermes-desktop", "codex", "copilot", "omp"}
if codexAppSupported() != nil {
wantPrefix = []string{"claude", "hermes", "openclaw", "opencode"}
wantPrefix = []string{"claude", "hermes", "openclaw", "opencode", "hermes-desktop", "codex", "copilot", "omp"}
}
if len(got) < len(wantPrefix) {
t.Fatalf("expected at least %d integrations, got %v", len(wantPrefix), got)
@@ -1870,9 +1942,9 @@ func TestListIntegrationInfos(t *testing.T) {
})
t.Run("includes known integrations", func(t *testing.T) {
known := map[string]bool{"claude": false, "cline": false, "codex": false, "opencode": false}
known := map[string]bool{"claude": false, "cline": false, "codex": false, "opencode": false, "omp": false}
if codexAppSupported() == nil {
known["codex-app"] = false
known["chatgpt"] = false
}
if poolsideGOOS != "windows" {
known["pool"] = false
@@ -1898,6 +1970,15 @@ func TestListIntegrationInfos(t *testing.T) {
t.Fatal("expected hermes to be included in ListIntegrationInfos")
})
t.Run("includes hermes desktop", func(t *testing.T) {
for _, info := range infos {
if info.Name == "hermes-desktop" {
return
}
}
t.Fatal("expected hermes-desktop to be included in ListIntegrationInfos")
})
t.Run("hermes still resolves explicitly", func(t *testing.T) {
name, runner, err := LookupIntegration("hermes")
if err != nil {
@@ -1996,6 +2077,7 @@ func TestIntegration_Editor(t *testing.T) {
{"claude", false},
{"claude-desktop", false},
{"codex", false},
{"omp", false},
{"nonexistent", false},
}
for _, tt := range tests {
@@ -2020,12 +2102,14 @@ func TestIntegration_AutoInstallable(t *testing.T) {
{"openclaw", true},
{"pi", true},
{"hermes", true},
{"hermes-desktop", true},
{"cline", true},
{"qwen", true},
{"claude", false},
{"claude", true},
{"claude-desktop", false},
{"codex", false},
{"opencode", false},
{"opencode", true},
{"omp", false},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
+99 -22
View File
@@ -287,12 +287,14 @@ Flags and extra arguments require an integration name.
Supported integrations:
claude Claude Code
codex-app Codex App (aliases: codex-desktop, codex-gui)
chatgpt ChatGPT (aliases: codex-app, codex-desktop, codex-gui)
hermes Hermes Agent
openclaw OpenClaw (aliases: clawdbot, moltbot)
opencode OpenCode
codex Codex
hermes-desktop Hermes Desktop
copilot Copilot CLI (aliases: copilot-cli)
omp OMP
droid Droid
kimi Kimi Code CLI
pi Pi
@@ -305,9 +307,10 @@ Examples:
ollama launch
ollama launch claude
ollama launch claude --model <model>
ollama launch codex-app
ollama launch codex-app --restore
ollama launch chatgpt
ollama launch chatgpt --restore
ollama launch hermes
ollama launch hermes-desktop
ollama launch droid --config (does not auto-launch)
ollama launch codex --restore
ollama launch codex -- --sandbox workspace-write`,
@@ -704,9 +707,12 @@ func (c *launcherClient) resolveRunModel(ctx context.Context, req RunModelReques
}
if usable {
if err := c.ensureModelsReady(ctx, []string{current}); err != nil {
return "", err
if !errors.Is(err, errDeprecatedLaunchModelDeclined) {
return "", err
}
} else {
return current, nil
}
return current, nil
}
}
@@ -723,7 +729,7 @@ func (c *launcherClient) resolveRunModel(ctx context.Context, req RunModelReques
}
func (c *launcherClient) launchSingleIntegration(ctx context.Context, name string, runner Runner, saved *config.IntegrationConfig, req IntegrationLaunchRequest) error {
target, _, err := c.resolveSingleIntegrationTarget(ctx, runner, primaryModelFromConfig(saved), req)
target, _, err := c.resolveSingleIntegrationTarget(ctx, name, runner, primaryModelFromConfig(saved), req)
if err != nil {
return err
}
@@ -745,14 +751,22 @@ func (c *launcherClient) launchEditorIntegration(ctx context.Context, name strin
models, needsConfigure := c.resolveEditorLaunchModels(ctx, saved, req)
if needsConfigure {
selected, err := c.selectMultiModelsForIntegration(ctx, runner, models)
selected, err := c.selectMultiModelsForIntegration(ctx, name, runner, models)
if err != nil {
return err
}
models = selected
} else if len(models) > 0 {
if err := c.ensureModelsReady(ctx, models[:1]); err != nil {
return err
if err := c.ensureModelsReadyFor(ctx, models[:1], runner.String(), name); err != nil {
if !errors.Is(err, errDeprecatedLaunchModelDeclined) || req.ModelOverride != "" {
return err
}
selected, err := c.selectMultiModelsForIntegration(ctx, name, runner, models)
if err != nil {
return err
}
models = selected
needsConfigure = true
}
}
@@ -761,7 +775,8 @@ func (c *launcherClient) launchEditorIntegration(ctx context.Context, name strin
}
var launchModels []LaunchModel
if (needsConfigure || req.ModelOverride != "") && !savedMatchesModels(saved, models) {
liveConfigMatches := slices.Equal(editor.Models(), models)
if needsConfigure || req.ModelOverride != "" || !savedMatchesModels(saved, models) || !liveConfigMatches {
launchModels = c.modelInventory().Resolve(ctx, models)
if err := prepareEditorIntegration(name, editor, launchModels); err != nil {
return err
@@ -780,7 +795,7 @@ func (c *launcherClient) launchManagedSingleIntegration(ctx context.Context, nam
selectionCurrent = primaryModelFromConfig(saved)
}
target, needsConfigure, err := c.resolveSingleIntegrationTarget(ctx, runner, selectionCurrent, req)
target, needsConfigure, err := c.resolveSingleIntegrationTarget(ctx, name, runner, selectionCurrent, req)
if err != nil {
return err
}
@@ -950,7 +965,7 @@ func (c *launcherClient) managedSingleConfigureModels(ctx context.Context, manag
return dedupeModelList(models), nil
}
func (c *launcherClient) resolveSingleIntegrationTarget(ctx context.Context, runner Runner, current string, req IntegrationLaunchRequest) (string, bool, error) {
func (c *launcherClient) resolveSingleIntegrationTarget(ctx context.Context, name string, runner Runner, current string, req IntegrationLaunchRequest) (string, bool, error) {
target := req.ModelOverride
needsConfigure := req.ForceConfigure
skipReadiness := false
@@ -974,14 +989,24 @@ func (c *launcherClient) resolveSingleIntegrationTarget(ctx context.Context, run
}
if needsConfigure && req.ModelOverride == "" {
selected, err := c.selectSingleModelWithSelectorReady(ctx, fmt.Sprintf("Select model for %s:", runner), target, DefaultSingleSelector, !skipReadiness)
selected, err := c.selectSingleModelWithSelectorReady(ctx, fmt.Sprintf("Select model for %s:", runner), target, DefaultSingleSelector, !skipReadiness, runner.String(), name)
if err != nil {
return "", false, err
}
target = selected
} else if !skipReadiness {
if err := c.ensureModelsReady(ctx, []string{target}); err != nil {
return "", false, err
if err := c.ensureModelsReadyFor(ctx, []string{target}, runner.String(), name); err != nil {
if !errors.Is(err, errDeprecatedLaunchModelDeclined) {
return "", false, err
}
// "Pick another model" is an interactive recovery path, including
// when --model supplied the initial target.
selected, err := c.selectSingleModelWithSelectorReady(ctx, fmt.Sprintf("Select model for %s:", runner), target, DefaultSingleSelector, true, runner.String(), name)
if err != nil {
return "", false, err
}
target = selected
needsConfigure = true
}
}
@@ -1019,7 +1044,7 @@ func managedRequiresInteractiveOnboarding(managed any) bool {
}
func (c *launcherClient) selectSingleModelWithSelector(ctx context.Context, title, current string, selector SingleSelector) (string, error) {
return c.selectSingleModelWithSelectorReady(ctx, title, current, selector, true)
return c.selectSingleModelWithSelectorReady(ctx, title, current, selector, true, "ollama launch", "")
}
func (c *launcherClient) latestAccountState() *AccountState {
@@ -1029,7 +1054,7 @@ func (c *launcherClient) latestAccountState() *AccountState {
return c.accountState
}
func (c *launcherClient) selectSingleModelWithSelectorReady(ctx context.Context, title, current string, selector SingleSelector, ensureReady bool) (string, error) {
func (c *launcherClient) selectSingleModelWithSelectorReady(ctx context.Context, title, current string, selector SingleSelector, ensureReady bool, label, commandName string) (string, error) {
if selector == nil && DefaultSingleSelectorWithUpdates == nil {
return "", fmt.Errorf("no selector configured")
}
@@ -1054,11 +1079,15 @@ func (c *launcherClient) selectSingleModelWithSelectorReady(ctx context.Context,
return "", ErrCancelled
}
if ensureReady {
if err := c.ensureModelsReady(ctx, []string{selected}); err != nil {
if err := c.ensureModelsReadyFor(ctx, []string{selected}, label, commandName); err != nil {
if errors.Is(err, errUpgradeCancelled) {
current = selected
continue
}
if errors.Is(err, errDeprecatedLaunchModelDeclined) {
current = selected
continue
}
return "", err
}
}
@@ -1066,7 +1095,7 @@ func (c *launcherClient) selectSingleModelWithSelectorReady(ctx context.Context,
}
}
func (c *launcherClient) selectMultiModelsForIntegration(ctx context.Context, runner Runner, preChecked []string) ([]string, error) {
func (c *launcherClient) selectMultiModelsForIntegration(ctx context.Context, name string, runner Runner, preChecked []string) ([]string, error) {
if DefaultMultiSelector == nil && DefaultMultiSelectorWithUpdates == nil {
return nil, fmt.Errorf("no selector configured")
}
@@ -1088,12 +1117,16 @@ func (c *launcherClient) selectMultiModelsForIntegration(ctx context.Context, ru
if err != nil {
return nil, err
}
accepted, skipped, err := c.selectReadyModelsForSave(ctx, selected)
accepted, skipped, err := c.selectReadyModelsForSave(ctx, selected, runner.String(), name)
if err != nil {
if errors.Is(err, errUpgradeCancelled) {
orderedChecked = append([]string(nil), selected...)
continue
}
if errors.Is(err, errDeprecatedLaunchModelDeclined) {
orderedChecked = append([]string(nil), selected...)
continue
}
return nil, err
}
for _, skip := range skipped {
@@ -1132,6 +1165,8 @@ func (c *launcherClient) loadSelectableModels(ctx context.Context, preChecked []
cloudDisabled, _ := cloudStatusDisabled(ctx, c.apiClient)
items, orderedChecked, _, _ := buildModelListWithRecommendations(inventory, recommendations, preChecked, current)
items = filterDeprecatedLaunchModelItems(items)
orderedChecked = filterDeprecatedLaunchModelNames(orderedChecked)
if cloudDisabled {
items = filterCloudItems(items)
orderedChecked = c.filterDisabledCloudModels(ctx, orderedChecked)
@@ -1208,13 +1243,31 @@ func (c *launcherClient) requestRecommendations(ctx context.Context) ([]ModelIte
}
func (c *launcherClient) ensureModelsReady(ctx context.Context, models []string) error {
return c.ensureModelsReadyFor(ctx, models, "ollama launch", "")
}
func (c *launcherClient) ensureModelsReadyFor(ctx context.Context, models []string, label, commandName string) error {
models = dedupeModelList(models)
if len(models) == 0 {
return nil
}
cloudRec, localRec := c.agentCapableRecommendations(ctx)
cloudModels := make(map[string]bool, len(models))
for _, model := range models {
if prompt := deprecatedLaunchModelPrompt(model, label, commandName, cloudRec, localRec); prompt != "" {
ok, err := ConfirmPromptWithOptions(prompt, ConfirmOptions{
YesLabel: "Launch anyway",
NoLabel: "Pick another model",
Default: ConfirmDefaultNo,
})
if err != nil {
return err
}
if !ok {
return errDeprecatedLaunchModelDeclined
}
}
isCloudModel := isCloudModelName(model)
if isCloudModel {
cloudModels[model] = true
@@ -1229,6 +1282,27 @@ func (c *launcherClient) ensureModelsReady(ctx context.Context, models []string)
return ensureAuth(ctx, c.apiClient, cloudModels, models)
}
func (c *launcherClient) agentCapableRecommendations(ctx context.Context) (cloud, local string) {
recs := c.recommendations(ctx)
cloudDisabled, known := cloudStatusDisabled(ctx, c.apiClient)
for _, rec := range recs {
if rec.Name == "" || isDeprecatedLaunchModel(rec.Name) {
continue
}
if isCloudModelName(rec.Name) {
if cloud == "" && !(known && cloudDisabled) {
cloud = rec.Name
}
} else if local == "" {
local = rec.Name
}
if cloud != "" && local != "" {
break
}
}
return cloud, local
}
func dedupeModelList(models []string) []string {
deduped := make([]string, 0, len(models))
seen := make(map[string]bool, len(models))
@@ -1247,16 +1321,19 @@ type skippedModel struct {
reason string
}
func (c *launcherClient) selectReadyModelsForSave(ctx context.Context, selected []string) ([]string, []skippedModel, error) {
func (c *launcherClient) selectReadyModelsForSave(ctx context.Context, selected []string, label, commandName string) ([]string, []skippedModel, error) {
selected = dedupeModelList(selected)
accepted := make([]string, 0, len(selected))
skipped := make([]skippedModel, 0, len(selected))
for _, model := range selected {
if err := c.ensureModelsReady(ctx, []string{model}); err != nil {
if err := c.ensureModelsReadyFor(ctx, []string{model}, label, commandName); err != nil {
if errors.Is(err, errUpgradeCancelled) {
return nil, nil, err
}
if errors.Is(err, errDeprecatedLaunchModelDeclined) {
return nil, nil, err
}
if errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) {
return nil, nil, err
}
+608 -91
View File
File diff suppressed because it is too large. Load diff
+15 -11
View File
@@ -194,10 +194,10 @@ func ensureCloudAuth(ctx context.Context, client *api.Client, modelList string)
}
var aErr api.AuthorizationError
if !errors.As(err, &aErr) || aErr.SigninURL == "" {
if err != nil {
return err
}
if err != nil && !errors.As(err, &aErr) {
return nil
}
if err == nil || aErr.SigninURL == "" {
return fmt.Errorf("%s requires sign in", modelList)
}
@@ -258,19 +258,23 @@ func showOrPullWithPolicy(ctx context.Context, client *api.Client, model string,
if _, err := client.Show(ctx, &api.ShowRequest{Model: model}); err == nil {
return nil
} else {
if isCloudModel {
if disabled, known := cloudStatusDisabled(ctx, client); known && disabled {
return errors.New(internalcloud.DisabledError("remote inference is unavailable"))
}
var statusErr api.StatusError
if errors.As(err, &statusErr) && statusErr.StatusCode == http.StatusNotFound {
return fmt.Errorf("model %q not found", model)
}
return nil
}
var statusErr api.StatusError
if !errors.As(err, &statusErr) || statusErr.StatusCode != http.StatusNotFound {
return err
}
}
if isCloudModel {
if disabled, known := cloudStatusDisabled(ctx, client); known && disabled {
return errors.New(internalcloud.DisabledError("remote inference is unavailable"))
}
return fmt.Errorf("model %q not found", model)
}
switch policy {
case missingModelAutoPull:
return pullMissingModel(ctx, client, model)
+454
View File
@@ -0,0 +1,454 @@
package launch
import (
"fmt"
"os"
"os/exec"
"path/filepath"
"runtime"
"slices"
"strings"
"github.com/ollama/ollama/cmd/config"
"github.com/ollama/ollama/cmd/internal/fileutil"
"github.com/ollama/ollama/envconfig"
"github.com/ollama/ollama/types/model"
"gopkg.in/yaml.v3"
)
const (
ompIntegrationName = "omp"
ompProviderName = "ollama"
ompSetupVersion = 1
ompWebSearchPlugin = "@ollama/pi-web-search"
)
// OMP implements Runner for the OMP coding-agent integration.
type OMP struct{}
func (o *OMP) String() string { return "OMP" }
func (o *OMP) Paths() []string {
var paths []string
for _, pathFn := range []func() (string, error){ompModelsPath, ompConfigPath} {
path, err := pathFn()
if err != nil {
continue
}
if _, err := os.Stat(path); err == nil {
paths = append(paths, path)
}
}
return paths
}
func (o *OMP) Configure(model string) error {
return o.ConfigureWithModels(model, []LaunchModel{fallbackLaunchModel(model)})
}
func (o *OMP) ConfigureWithModels(primary string, models []LaunchModel) error {
if primary == "" {
return nil
}
if len(models) == 0 {
models = []LaunchModel{fallbackLaunchModel(primary)}
}
if err := writeOMPModelsConfig(primary, models); err != nil {
return err
}
return writeOMPAgentConfig()
}
func (o *OMP) CurrentModel() string {
cfg, err := readOMPModelsConfig()
if err != nil {
return ""
}
provider, ok := ompProvider(cfg)
if !ok {
return ""
}
if !ompProviderHealthy(provider) {
return ""
}
models, _ := provider["models"].([]any)
for _, raw := range models {
entry, ok := raw.(map[string]any)
if !ok {
continue
}
if id, _ := entry["id"].(string); id != "" {
return id
}
}
return ""
}
func (o *OMP) Onboard() error {
return config.MarkIntegrationOnboarded(ompIntegrationName)
}
func (o *OMP) RequiresInteractiveOnboarding() bool { return false }
func (o *OMP) args(model string, extra []string) []string {
var args []string
if model != "" {
args = append(args, "--model", ompModelName(model))
}
args = append(args, extra...)
return args
}
func ompModelName(model string) string {
if strings.HasPrefix(model, "ollama/") {
return model
}
return "ollama/" + model
}
func (o *OMP) findPath() (string, error) {
if p, err := exec.LookPath("omp"); err == nil {
return p, nil
}
home, err := os.UserHomeDir()
if err != nil {
return "", err
}
for _, dir := range []string{
filepath.Join(home, ".local", "bin"),
filepath.Join(home, ".bun", "bin"),
} {
for _, name := range ompExecutableNames() {
fallback := filepath.Join(dir, name)
if _, err := os.Stat(fallback); err == nil {
return fallback, nil
}
}
}
return "", exec.ErrNotFound
}
func ompExecutableNames() []string {
if runtime.GOOS == "windows" {
return []string{"omp.exe", "omp.cmd", "omp.bat"}
}
return []string{"omp"}
}
func (o *OMP) Run(model string, _ []LaunchModel, args []string) error {
ompPath, err := o.findPath()
if err != nil {
return fmt.Errorf("omp is not installed, install from https://omp.sh")
}
ensureOMPWebSearchPlugin(ompPath)
cmd := exec.Command(ompPath, o.args(model, args)...)
cmd.Stdin = os.Stdin
cmd.Stdout = os.Stdout
cmd.Stderr = os.Stderr
cmd.Env = os.Environ()
return cmd.Run()
}
func ensureOMPWebSearchPlugin(bin string) {
if !shouldManageOllamaWebSearch() {
fmt.Fprintf(os.Stderr, "%sCloud is disabled; skipping %s setup.%s\n", ansiGray, ompWebSearchPlugin, ansiReset)
return
}
fmt.Fprintf(os.Stderr, "%sChecking OMP web search plugin...%s\n", ansiGray, ansiReset)
installed, err := ompPluginInstalled(bin, ompWebSearchPlugin)
if err != nil {
fmt.Fprintf(os.Stderr, "%s Warning: could not check %s installation: %v%s\n", ansiYellow, ompWebSearchPlugin, err, ansiReset)
return
}
verb := "Installing"
warnVerb := "install"
doneVerb := "Installed"
if installed {
verb = "Updating"
warnVerb = "update"
doneVerb = "Updated"
}
fmt.Fprintf(os.Stderr, "%s%s %s...%s\n", ansiGray, verb, ompWebSearchPlugin, ansiReset)
cmd := exec.Command(bin, "plugin", "install", ompWebSearchPlugin)
cmd.Stdout = os.Stdout
cmd.Stderr = os.Stderr
if err := cmd.Run(); err != nil {
fmt.Fprintf(os.Stderr, "%s Warning: could not %s %s: %v%s\n", ansiYellow, warnVerb, ompWebSearchPlugin, err, ansiReset)
return
}
fmt.Fprintf(os.Stderr, "%s ✓ %s %s%s\n", ansiGreen, doneVerb, ompWebSearchPlugin, ansiReset)
}
func ompPluginInstalled(bin, plugin string) (bool, error) {
cmd := exec.Command(bin, "plugin", "list")
out, err := cmd.CombinedOutput()
if err != nil {
msg := strings.TrimSpace(string(out))
if msg == "" {
return false, err
}
return false, fmt.Errorf("%w: %s", err, msg)
}
versioned := plugin + "@"
for _, line := range strings.Split(string(out), "\n") {
trimmed := strings.TrimSpace(line)
if strings.Contains(trimmed, versioned) || trimmed == plugin {
return true, nil
}
}
return false, nil
}
func ompModelsPath() (string, error) {
dir, err := ompAgentDir()
if err != nil {
return "", err
}
return filepath.Join(dir, "models.yml"), nil
}
func ompConfigPath() (string, error) {
dir, err := ompAgentDir()
if err != nil {
return "", err
}
return filepath.Join(dir, "config.yml"), nil
}
func ompAgentDir() (string, error) {
if dir := strings.TrimSpace(os.Getenv("PI_CODING_AGENT_DIR")); dir != "" {
return dir, nil
}
home, err := os.UserHomeDir()
if err != nil {
return "", err
}
configDir := strings.TrimSpace(os.Getenv("PI_CONFIG_DIR"))
if configDir == "" {
configDir = ".omp"
}
if filepath.IsAbs(configDir) {
return filepath.Join(configDir, "agent"), nil
}
return filepath.Join(home, configDir, "agent"), nil
}
func readOMPModelsConfig() (map[string]any, error) {
path, err := ompModelsPath()
if err != nil {
return nil, err
}
data, err := os.ReadFile(path)
if err != nil {
return nil, err
}
var cfg map[string]any
if err := yaml.Unmarshal(data, &cfg); err != nil {
return nil, err
}
if cfg == nil {
cfg = make(map[string]any)
}
return cfg, nil
}
func writeOMPModelsConfig(primary string, models []LaunchModel) error {
path, err := ompModelsPath()
if err != nil {
return err
}
if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil {
return err
}
cfg := make(map[string]any)
if existing, err := readOMPModelsConfig(); err == nil {
cfg = existing
}
provider := ensureOMPProvider(cfg)
existingByID := ompModelEntriesByID(provider)
ordered := append([]LaunchModel(nil), models...)
if model, ok := findLaunchModel(ordered, primary); ok {
ordered = append([]LaunchModel{model}, removeLaunchModel(ordered, primary)...)
} else {
ordered = append([]LaunchModel{fallbackLaunchModel(primary)}, ordered...)
}
var merged []any
seen := make(map[string]bool, len(ordered))
for _, model := range ordered {
if model.Name == "" || seen[model.Name] {
continue
}
seen[model.Name] = true
entry := ompModelConfig(model)
if existing, ok := existingByID[model.Name]; ok {
for key, value := range existing {
if _, overridden := entry[key]; !overridden {
entry[key] = value
}
}
}
merged = append(merged, entry)
}
for _, raw := range ompProviderModels(provider) {
entry, ok := raw.(map[string]any)
if !ok {
merged = append(merged, raw)
continue
}
id, _ := entry["id"].(string)
if id == "" || seen[id] {
continue
}
merged = append(merged, entry)
}
provider["models"] = merged
data, err := yaml.Marshal(cfg)
if err != nil {
return err
}
return fileutil.WriteWithBackup(path, data, ompIntegrationName)
}
func writeOMPAgentConfig() error {
path, err := ompConfigPath()
if err != nil {
return err
}
if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil {
return err
}
cfg := make(map[string]any)
if data, err := os.ReadFile(path); err == nil {
if err := yaml.Unmarshal(data, &cfg); err != nil {
return err
}
if cfg == nil {
cfg = make(map[string]any)
}
}
cfg["setupVersion"] = ompSetupVersion
data, err := yaml.Marshal(cfg)
if err != nil {
return err
}
return fileutil.WriteWithBackup(path, data, ompIntegrationName)
}
func ensureOMPProvider(cfg map[string]any) map[string]any {
providers, _ := cfg["providers"].(map[string]any)
if providers == nil {
providers = make(map[string]any)
cfg["providers"] = providers
}
provider, _ := providers[ompProviderName].(map[string]any)
if provider == nil {
provider = make(map[string]any)
providers[ompProviderName] = provider
}
provider["baseUrl"] = ompBaseURL()
provider["api"] = "openai-responses"
provider["auth"] = "none"
provider["discovery"] = map[string]any{"type": "ollama"}
return provider
}
func ompBaseURL() string {
return strings.TrimRight(envconfig.ConnectableHost().String(), "/") + "/v1"
}
func ompProviderHealthy(provider map[string]any) bool {
baseURL, _ := provider["baseUrl"].(string)
if strings.TrimRight(baseURL, "/") != strings.TrimRight(ompBaseURL(), "/") {
return false
}
api, _ := provider["api"].(string)
if api != "openai-responses" {
return false
}
auth, _ := provider["auth"].(string)
if auth != "none" {
return false
}
discovery, _ := provider["discovery"].(map[string]any)
if discovery == nil {
return false
}
discoveryType, _ := discovery["type"].(string)
return discoveryType == "ollama"
}
func ompProvider(cfg map[string]any) (map[string]any, bool) {
providers, ok := cfg["providers"].(map[string]any)
if !ok {
return nil, false
}
provider, ok := providers[ompProviderName].(map[string]any)
return provider, ok
}
func ompProviderModels(provider map[string]any) []any {
models, _ := provider["models"].([]any)
return models
}
func ompModelEntriesByID(provider map[string]any) map[string]map[string]any {
out := make(map[string]map[string]any)
for _, raw := range ompProviderModels(provider) {
entry, ok := raw.(map[string]any)
if !ok {
continue
}
if id, _ := entry["id"].(string); id != "" {
out[id] = entry
}
}
return out
}
func ompModelConfig(modelInfo LaunchModel) map[string]any {
entry := map[string]any{
"id": modelInfo.Name,
"name": modelInfo.Name,
}
input := []string{"text"}
if slices.Contains(modelInfo.Capabilities, model.CapabilityVision) {
input = append(input, "image")
}
entry["input"] = input
if modelInfo.ContextLength > 0 {
entry["contextWindow"] = modelInfo.ContextLength
}
if modelInfo.MaxOutputTokens > 0 {
entry["maxTokens"] = modelInfo.MaxOutputTokens
}
return entry
}
func removeLaunchModel(models []LaunchModel, name string) []LaunchModel {
out := make([]LaunchModel, 0, len(models))
for _, model := range models {
if launchModelMatches(model.Name, name) || launchModelMatches(name, model.Name) {
continue
}
out = append(out, model)
}
return out
}
+687
View File
@@ -0,0 +1,687 @@
package launch
import (
"fmt"
"net/http"
"net/http/httptest"
"os"
"path/filepath"
"runtime"
"slices"
"strings"
"testing"
modelpkg "github.com/ollama/ollama/types/model"
"gopkg.in/yaml.v3"
)
func TestMain(m *testing.M) {
if os.Getenv("OLLAMA_LAUNCH_OMP_TEST_HELPER") == "1" {
runOMPTestHelper()
return
}
os.Exit(m.Run())
}
func runOMPTestHelper() {
logPath := os.Getenv("OLLAMA_LAUNCH_OMP_TEST_LOG")
if logPath != "" {
f, err := os.OpenFile(logPath, os.O_CREATE|os.O_APPEND|os.O_WRONLY, 0o644)
if err == nil {
_, _ = fmt.Fprintln(f, strings.Join(os.Args[1:], " "))
_ = f.Close()
}
}
if len(os.Args) >= 3 && os.Args[1] == "plugin" && os.Args[2] == "list" {
fmt.Print(os.Getenv("OLLAMA_LAUNCH_OMP_TEST_PLUGIN_LIST"))
os.Exit(0)
}
if len(os.Args) >= 4 && os.Args[1] == "plugin" && os.Args[2] == "install" {
if os.Getenv("OLLAMA_LAUNCH_OMP_TEST_FAIL_INSTALL") == "1" {
_, _ = fmt.Fprintln(os.Stderr, "install failed")
os.Exit(1)
}
os.Exit(0)
}
os.Exit(0)
}
func setOMPTestHome(t *testing.T, dir string) {
t.Helper()
setTestHome(t, dir)
t.Setenv("PI_CONFIG_DIR", "")
t.Setenv("PI_CODING_AGENT_DIR", "")
}
func TestOMPIntegration(t *testing.T) {
o := &OMP{}
t.Run("String", func(t *testing.T) {
if got := o.String(); got != "OMP" {
t.Errorf("String() = %q, want %q", got, "OMP")
}
})
t.Run("implements Runner", func(t *testing.T) {
var _ Runner = o
})
t.Run("implements ManagedSingleModel", func(t *testing.T) {
var _ ManagedSingleModel = o
})
t.Run("implements ManagedModelListConfigurer", func(t *testing.T) {
var _ ManagedModelListConfigurer = o
})
t.Run("does not require interactive onboarding", func(t *testing.T) {
var _ ManagedInteractiveOnboarding = o
if o.RequiresInteractiveOnboarding() {
t.Fatal("OMP onboarding should not require an interactive terminal")
}
})
}
func TestOMPArgs(t *testing.T) {
o := &OMP{}
tests := []struct {
name string
model string
args []string
want []string
}{
{"with model", "gemma4", nil, []string{"--model", "ollama/gemma4"}},
{"with cloud model", "kimi-k2.6:cloud", nil, []string{"--model", "ollama/kimi-k2.6:cloud"}},
{"empty model", "", nil, nil},
{"with model and extra", "gemma4", []string{"--help"}, []string{"--model", "ollama/gemma4", "--help"}},
{"already qualified", "ollama/gemma4", nil, []string{"--model", "ollama/gemma4"}},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got := o.args(tt.model, tt.args)
if !slices.Equal(got, tt.want) {
t.Errorf("args(%q, %v) = %v, want %v", tt.model, tt.args, got, tt.want)
}
})
}
}
func TestOMPRun_WebSearchPluginLifecycle(t *testing.T) {
seedOMPHelperBinary := func(t *testing.T, dir string) {
t.Helper()
src, err := os.Executable()
if err != nil {
t.Fatal(err)
}
data, err := os.ReadFile(src)
if err != nil {
t.Fatal(err)
}
dst := filepath.Join(dir, ompExecutableNames()[0])
if err := os.WriteFile(dst, data, 0o755); err != nil {
t.Fatal(err)
}
}
setCloudStatus := func(t *testing.T, disabled bool) {
t.Helper()
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.URL.Path == "/api/status" {
fmt.Fprintf(w, `{"cloud":{"disabled":%t,"source":"config"}}`, disabled)
return
}
http.NotFound(w, r)
}))
t.Cleanup(srv.Close)
t.Setenv("OLLAMA_HOST", srv.URL)
}
setup := func(t *testing.T, pluginList string, cloudDisabled bool) (string, *OMP) {
t.Helper()
tmpDir := t.TempDir()
setOMPTestHome(t, tmpDir)
t.Setenv("PATH", tmpDir)
t.Setenv("OLLAMA_LAUNCH_OMP_TEST_HELPER", "1")
t.Setenv("OLLAMA_LAUNCH_OMP_TEST_PLUGIN_LIST", pluginList)
logPath := filepath.Join(tmpDir, "omp.log")
t.Setenv("OLLAMA_LAUNCH_OMP_TEST_LOG", logPath)
setCloudStatus(t, cloudDisabled)
seedOMPHelperBinary(t, tmpDir)
return logPath, &OMP{}
}
t.Run("web search missing installs before launch", func(t *testing.T) {
logPath, o := setup(t, "No plugins installed\n", false)
if err := o.Run("kimi-k2.6:cloud", nil, []string{"session"}); err != nil {
t.Fatalf("Run() error = %v", err)
}
calls, err := os.ReadFile(logPath)
if err != nil {
t.Fatal(err)
}
got := string(calls)
if !strings.Contains(got, "plugin list\n") {
t.Fatalf("expected plugin list call, got:\n%s", got)
}
if !strings.Contains(got, "plugin install "+ompWebSearchPlugin+"\n") {
t.Fatalf("expected plugin install call, got:\n%s", got)
}
if !strings.Contains(got, "--model ollama/kimi-k2.6:cloud session\n") {
t.Fatalf("expected final omp launch call, got:\n%s", got)
}
})
t.Run("web search present refreshes before launch", func(t *testing.T) {
logPath, o := setup(t, "npm Plugins:\n\n● "+ompWebSearchPlugin+"@0.0.5\n", false)
if err := o.Run("gemma4", nil, []string{"chat"}); err != nil {
t.Fatalf("Run() error = %v", err)
}
calls, err := os.ReadFile(logPath)
if err != nil {
t.Fatal(err)
}
got := string(calls)
if !strings.Contains(got, "plugin install "+ompWebSearchPlugin+"\n") {
t.Fatalf("expected plugin refresh install call, got:\n%s", got)
}
if !strings.Contains(got, "--model ollama/gemma4 chat\n") {
t.Fatalf("expected final omp launch call, got:\n%s", got)
}
})
t.Run("web search install failure warns and continues", func(t *testing.T) {
logPath, o := setup(t, "No plugins installed\n", false)
t.Setenv("OLLAMA_LAUNCH_OMP_TEST_FAIL_INSTALL", "1")
stderr := captureStderr(t, func() {
if err := o.Run("gemma4", nil, []string{"chat"}); err != nil {
t.Fatalf("Run() should continue after plugin install failure, got %v", err)
}
})
if !strings.Contains(stderr, "Warning: could not install "+ompWebSearchPlugin) {
t.Fatalf("expected install warning, got:\n%s", stderr)
}
calls, err := os.ReadFile(logPath)
if err != nil {
t.Fatal(err)
}
if !strings.Contains(string(calls), "--model ollama/gemma4 chat\n") {
t.Fatalf("expected final omp launch call, got:\n%s", calls)
}
})
t.Run("cloud disabled skips web search plugin management", func(t *testing.T) {
logPath, o := setup(t, "No plugins installed\n", true)
stderr := captureStderr(t, func() {
if err := o.Run("gemma4", nil, []string{"chat"}); err != nil {
t.Fatalf("Run() error = %v", err)
}
})
if !strings.Contains(stderr, "Cloud is disabled; skipping "+ompWebSearchPlugin+" setup.") {
t.Fatalf("expected cloud-disabled skip message, got:\n%s", stderr)
}
calls, err := os.ReadFile(logPath)
if err != nil {
t.Fatal(err)
}
got := string(calls)
if strings.Contains(got, "plugin list\n") || strings.Contains(got, "plugin install "+ompWebSearchPlugin+"\n") {
t.Fatalf("did not expect plugin management calls, got:\n%s", got)
}
if !strings.Contains(got, "--model ollama/gemma4 chat\n") {
t.Fatalf("expected final omp launch call, got:\n%s", got)
}
})
}
func TestOMPFindPath(t *testing.T) {
o := &OMP{}
t.Run("finds omp in PATH", func(t *testing.T) {
tmpDir := t.TempDir()
name := "omp"
if runtime.GOOS == "windows" {
name = "omp.exe"
}
fakeBin := filepath.Join(tmpDir, name)
os.WriteFile(fakeBin, []byte("#!/bin/sh\n"), 0o755)
t.Setenv("PATH", tmpDir)
got, err := o.findPath()
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if got != fakeBin {
t.Errorf("findPath() = %q, want %q", got, fakeBin)
}
})
t.Run("falls back to ~/.local/bin/omp", func(t *testing.T) {
home := t.TempDir()
setOMPTestHome(t, home)
t.Setenv("PATH", t.TempDir())
fallback := filepath.Join(home, ".local", "bin", ompExecutableNames()[0])
os.MkdirAll(filepath.Dir(fallback), 0o755)
os.WriteFile(fallback, []byte("#!/bin/sh\n"), 0o755)
got, err := o.findPath()
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if got != fallback {
t.Errorf("findPath() = %q, want %q", got, fallback)
}
})
t.Run("falls back to ~/.bun/bin/omp", func(t *testing.T) {
home := t.TempDir()
setOMPTestHome(t, home)
t.Setenv("PATH", t.TempDir())
fallback := filepath.Join(home, ".bun", "bin", ompExecutableNames()[0])
os.MkdirAll(filepath.Dir(fallback), 0o755)
os.WriteFile(fallback, []byte("#!/bin/sh\n"), 0o755)
got, err := o.findPath()
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if got != fallback {
t.Errorf("findPath() = %q, want %q", got, fallback)
}
})
t.Run("returns error when not found", func(t *testing.T) {
home := t.TempDir()
setOMPTestHome(t, home)
t.Setenv("PATH", t.TempDir())
if _, err := o.findPath(); err == nil {
t.Fatal("expected error, got nil")
}
})
}
func TestOMPConfigureWithModelsWritesModelsYML(t *testing.T) {
home := t.TempDir()
setOMPTestHome(t, home)
t.Setenv("OLLAMA_HOST", "http://0.0.0.0:11434")
o := &OMP{}
models := []LaunchModel{
{
Name: "glm-5.1:cloud",
ContextLength: 202_752,
MaxOutputTokens: 131_072,
},
{
Name: "qwen3.6",
Capabilities: []modelpkg.Capability{modelpkg.CapabilityVision},
},
}
if err := o.ConfigureWithModels("glm-5.1:cloud", models); err != nil {
t.Fatalf("ConfigureWithModels returned error: %v", err)
}
path := filepath.Join(home, ".omp", "agent", "models.yml")
data, err := os.ReadFile(path)
if err != nil {
t.Fatalf("failed to read models.yml: %v", err)
}
cfg := parseOMPConfigYAML(t, data)
provider := ompProviderFromYAML(t, cfg)
if provider["baseUrl"] != "http://127.0.0.1:11434/v1" {
t.Fatalf("baseUrl = %v, want connectable OpenAI-compatible host", provider["baseUrl"])
}
if provider["api"] != "openai-responses" {
t.Fatalf("api = %v, want openai-responses", provider["api"])
}
if provider["auth"] != "none" {
t.Fatalf("auth = %v, want none", provider["auth"])
}
discovery, _ := provider["discovery"].(map[string]any)
if discovery["type"] != "ollama" {
t.Fatalf("discovery = %v, want type ollama", discovery)
}
entries := ompModelEntriesFromYAML(t, provider)
if len(entries) != 2 {
t.Fatalf("models length = %d, want 2", len(entries))
}
if entries[0]["id"] != "glm-5.1:cloud" {
t.Fatalf("first model id = %v, want primary first", entries[0]["id"])
}
if got := numericYAMLValue(entries[0]["contextWindow"]); got != 202_752 {
t.Fatalf("contextWindow = %d, want 202752", got)
}
if got := numericYAMLValue(entries[0]["maxTokens"]); got != 131_072 {
t.Fatalf("maxTokens = %d, want 131072", got)
}
if input := stringSliceYAMLValue(entries[1]["input"]); !slices.Equal(input, []string{"text", "image"}) {
t.Fatalf("vision input = %v, want [text image]", input)
}
if got := o.CurrentModel(); got != "glm-5.1:cloud" {
t.Fatalf("CurrentModel = %q, want glm-5.1:cloud", got)
}
configPath := filepath.Join(home, ".omp", "agent", "config.yml")
configData, err := os.ReadFile(configPath)
if err != nil {
t.Fatalf("failed to read config.yml: %v", err)
}
config := parseOMPConfigYAML(t, configData)
if got := numericYAMLValue(config["setupVersion"]); got != ompSetupVersion {
t.Fatalf("setupVersion = %d, want %d", got, ompSetupVersion)
}
if paths := o.Paths(); !slices.Equal(paths, []string{path, configPath}) {
t.Fatalf("Paths = %v, want [%s %s]", paths, path, configPath)
}
}
func TestOMPConfigureWithModelsPreservesExistingConfig(t *testing.T) {
home := t.TempDir()
setOMPTestHome(t, home)
modelsPath := filepath.Join(home, ".omp", "agent", "models.yml")
if err := os.MkdirAll(filepath.Dir(modelsPath), 0o755); err != nil {
t.Fatal(err)
}
existing := []byte(`
providers:
anthropic:
baseUrl: https://example.com/anthropic
ollama:
baseUrl: http://old-host:11434
api: openai-responses
auth: none
models:
- id: old-model
name: Old Model
customField: keep-me
`)
if err := os.WriteFile(modelsPath, existing, 0o644); err != nil {
t.Fatal(err)
}
configPath := filepath.Join(home, ".omp", "agent", "config.yml")
existingConfig := []byte(`
lastChangelogVersion: 15.7.6
setupVersion: 0
theme: monochrome
`)
if err := os.WriteFile(configPath, existingConfig, 0o644); err != nil {
t.Fatal(err)
}
o := &OMP{}
if err := o.ConfigureWithModels("new-model", []LaunchModel{{Name: "new-model"}, {Name: "old-model"}}); err != nil {
t.Fatalf("ConfigureWithModels returned error: %v", err)
}
data, err := os.ReadFile(modelsPath)
if err != nil {
t.Fatal(err)
}
cfg := parseOMPConfigYAML(t, data)
providers, _ := cfg["providers"].(map[string]any)
if _, ok := providers["anthropic"]; !ok {
t.Fatalf("expected non-Ollama provider to be preserved: %v", providers)
}
provider := ompProviderFromYAML(t, cfg)
if provider["baseUrl"] != "http://127.0.0.1:11434/v1" {
t.Fatalf("baseUrl = %v, want repaired OpenAI-compatible host", provider["baseUrl"])
}
entries := ompModelEntriesFromYAML(t, provider)
if len(entries) != 2 {
t.Fatalf("models length = %d, want 2", len(entries))
}
if entries[0]["id"] != "new-model" {
t.Fatalf("first model id = %v, want new-model", entries[0]["id"])
}
if entries[1]["id"] != "old-model" {
t.Fatalf("second model id = %v, want old-model", entries[1]["id"])
}
if entries[1]["customField"] != "keep-me" {
t.Fatalf("custom field was not preserved: %v", entries[1])
}
configData, err := os.ReadFile(configPath)
if err != nil {
t.Fatal(err)
}
config := parseOMPConfigYAML(t, configData)
if got := numericYAMLValue(config["setupVersion"]); got != ompSetupVersion {
t.Fatalf("setupVersion = %d, want %d", got, ompSetupVersion)
}
if config["theme"] != "monochrome" {
t.Fatalf("theme was not preserved: %v", config)
}
if config["lastChangelogVersion"] != "15.7.6" {
t.Fatalf("lastChangelogVersion was not preserved: %v", config)
}
}
func TestOMPConfigureWithModelsAlwaysMarksSetupComplete(t *testing.T) {
home := t.TempDir()
setOMPTestHome(t, home)
configPath := filepath.Join(home, ".omp", "agent", "config.yml")
if err := os.MkdirAll(filepath.Dir(configPath), 0o755); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(configPath, []byte("setupVersion: 2\n"), 0o644); err != nil {
t.Fatal(err)
}
o := &OMP{}
if err := o.ConfigureWithModels("new-model", []LaunchModel{{Name: "new-model"}}); err != nil {
t.Fatalf("ConfigureWithModels returned error: %v", err)
}
configData, err := os.ReadFile(configPath)
if err != nil {
t.Fatal(err)
}
config := parseOMPConfigYAML(t, configData)
if got := numericYAMLValue(config["setupVersion"]); got != ompSetupVersion {
t.Fatalf("setupVersion = %d, want %d", got, ompSetupVersion)
}
}
func TestOMPConfigureWithModelsRespectsPiConfigDir(t *testing.T) {
home := t.TempDir()
setOMPTestHome(t, home)
t.Setenv("PI_CONFIG_DIR", ".custom-omp")
o := &OMP{}
if err := o.ConfigureWithModels("new-model", []LaunchModel{{Name: "new-model"}}); err != nil {
t.Fatalf("ConfigureWithModels returned error: %v", err)
}
modelsPath := filepath.Join(home, ".custom-omp", "agent", "models.yml")
configPath := filepath.Join(home, ".custom-omp", "agent", "config.yml")
for _, path := range []string{modelsPath, configPath} {
if _, err := os.Stat(path); err != nil {
t.Fatalf("expected %s to be written: %v", path, err)
}
}
if _, err := os.Stat(filepath.Join(home, ".omp", "agent", "models.yml")); !os.IsNotExist(err) {
t.Fatalf("expected default OMP models path to be untouched, got err %v", err)
}
if paths := o.Paths(); !slices.Equal(paths, []string{modelsPath, configPath}) {
t.Fatalf("Paths = %v, want [%s %s]", paths, modelsPath, configPath)
}
}
func TestOMPConfigureWithModelsRespectsPiCodingAgentDir(t *testing.T) {
home := t.TempDir()
setOMPTestHome(t, home)
agentDir := filepath.Join(home, "agent-override")
t.Setenv("PI_CONFIG_DIR", ".ignored-omp")
t.Setenv("PI_CODING_AGENT_DIR", agentDir)
o := &OMP{}
if err := o.ConfigureWithModels("new-model", []LaunchModel{{Name: "new-model"}}); err != nil {
t.Fatalf("ConfigureWithModels returned error: %v", err)
}
modelsPath := filepath.Join(agentDir, "models.yml")
configPath := filepath.Join(agentDir, "config.yml")
for _, path := range []string{modelsPath, configPath} {
if _, err := os.Stat(path); err != nil {
t.Fatalf("expected %s to be written: %v", path, err)
}
}
if _, err := os.Stat(filepath.Join(home, ".ignored-omp", "agent", "models.yml")); !os.IsNotExist(err) {
t.Fatalf("expected PI_CONFIG_DIR path to be ignored when PI_CODING_AGENT_DIR is set, got err %v", err)
}
if got := o.CurrentModel(); got != "new-model" {
t.Fatalf("CurrentModel = %q, want new-model", got)
}
}
func TestOMPCurrentModelRequiresHealthyProvider(t *testing.T) {
home := t.TempDir()
setOMPTestHome(t, home)
t.Setenv("OLLAMA_HOST", "http://127.0.0.1:11434")
modelsPath := filepath.Join(home, ".omp", "agent", "models.yml")
if err := os.MkdirAll(filepath.Dir(modelsPath), 0o755); err != nil {
t.Fatal(err)
}
tests := []struct {
name string
provider string
}{
{
name: "wrong base url",
provider: "" +
" baseUrl: http://127.0.0.1:9999/v1\n" +
" api: openai-responses\n" +
" auth: none\n" +
" discovery:\n" +
" type: ollama\n",
},
{
name: "wrong api",
provider: "" +
" baseUrl: http://127.0.0.1:11434/v1\n" +
" api: openai-chat\n" +
" auth: none\n" +
" discovery:\n" +
" type: ollama\n",
},
{
name: "wrong auth",
provider: "" +
" baseUrl: http://127.0.0.1:11434/v1\n" +
" api: openai-responses\n" +
" auth: api-key\n" +
" discovery:\n" +
" type: ollama\n",
},
{
name: "wrong discovery",
provider: "" +
" baseUrl: http://127.0.0.1:11434/v1\n" +
" api: openai-responses\n" +
" auth: none\n" +
" discovery:\n" +
" type: static\n",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
cfg := "providers:\n" +
" ollama:\n" +
tt.provider +
" models:\n" +
" - id: gemma4\n"
if err := os.WriteFile(modelsPath, []byte(cfg), 0o644); err != nil {
t.Fatal(err)
}
if got := (&OMP{}).CurrentModel(); got != "" {
t.Fatalf("expected stale config to return empty current model, got %q", got)
}
})
}
}
func parseOMPConfigYAML(t *testing.T, data []byte) map[string]any {
t.Helper()
var cfg map[string]any
if err := yaml.Unmarshal(data, &cfg); err != nil {
t.Fatalf("generated YAML did not parse: %v\n%s", err, data)
}
return cfg
}
func ompProviderFromYAML(t *testing.T, cfg map[string]any) map[string]any {
t.Helper()
providers, ok := cfg["providers"].(map[string]any)
if !ok {
t.Fatalf("providers missing from config: %v", cfg)
}
provider, ok := providers["ollama"].(map[string]any)
if !ok {
t.Fatalf("ollama provider missing from config: %v", providers)
}
return provider
}
func ompModelEntriesFromYAML(t *testing.T, provider map[string]any) []map[string]any {
t.Helper()
rawModels, ok := provider["models"].([]any)
if !ok {
t.Fatalf("provider models missing: %v", provider)
}
models := make([]map[string]any, 0, len(rawModels))
for _, raw := range rawModels {
entry, ok := raw.(map[string]any)
if !ok {
t.Fatalf("model entry has unexpected type %T: %v", raw, raw)
}
models = append(models, entry)
}
return models
}
func numericYAMLValue(value any) int {
switch v := value.(type) {
case int:
return v
case int64:
return int(v)
case float64:
return int(v)
default:
return 0
}
}
func stringSliceYAMLValue(value any) []string {
raw, _ := value.([]any)
out := make([]string, 0, len(raw))
for _, item := range raw {
if s, ok := item.(string); ok {
out = append(out, s)
}
}
return out
}
+117 -4
View File
@@ -8,11 +8,16 @@ import (
"path/filepath"
"runtime"
"slices"
"strings"
"github.com/ollama/ollama/cmd/internal/fileutil"
"github.com/ollama/ollama/envconfig"
)
const openCodeInstallScript = "curl -fsSL https://opencode.ai/install | bash"
var openCodeGOOS = runtime.GOOS
// OpenCode implements Runner and Editor for OpenCode integration.
// Config is passed via OPENCODE_CONFIG_CONTENT env var at launch time
// instead of writing to opencode's config files.
@@ -33,7 +38,7 @@ func findOpenCode() (string, bool) {
return "", false
}
name := "opencode"
if runtime.GOOS == "windows" {
if openCodeGOOS == "windows" {
name = "opencode.exe"
}
fallback := filepath.Join(home, ".opencode", "bin", name)
@@ -44,9 +49,9 @@ func findOpenCode() (string, bool) {
}
func (o *OpenCode) Run(model string, models []LaunchModel, args []string) error {
opencodePath, ok := findOpenCode()
if !ok {
return fmt.Errorf("opencode is not installed, install from https://opencode.ai")
opencodePath, err := ensureOpenCodeInstalled()
if err != nil {
return err
}
cmd := exec.Command(opencodePath, args...)
@@ -60,6 +65,78 @@ func (o *OpenCode) Run(model string, models []LaunchModel, args []string) error
return cmd.Run()
}
func ensureOpenCodeInstalled() (string, error) {
if opencodePath, ok := findOpenCode(); ok {
return opencodePath, nil
}
if err := checkOpenCodeInstallerDependencies(); err != nil {
return "", err
}
ok, err := ConfirmPrompt("OpenCode is not installed. Install now?")
if err != nil {
return "", err
}
if !ok {
return "", fmt.Errorf("opencode installation cancelled")
}
bin, args, err := openCodeInstallerCommand(openCodeGOOS)
if err != nil {
return "", err
}
fmt.Fprintf(os.Stderr, "\nInstalling OpenCode...\n")
cmd := exec.Command(bin, args...)
cmd.Stdin = os.Stdin
cmd.Stdout = os.Stdout
cmd.Stderr = os.Stderr
if err := cmd.Run(); err != nil {
return "", fmt.Errorf("failed to install opencode: %w", err)
}
opencodePath, ok := findOpenCode()
if !ok {
return "", fmt.Errorf("opencode was installed but the binary was not found on PATH\n\nYou may need to restart your shell")
}
fmt.Fprintf(os.Stderr, "%sOpenCode installed successfully%s\n\n", ansiGreen, ansiReset)
return opencodePath, nil
}
func checkOpenCodeInstallerDependencies() error {
switch openCodeGOOS {
case "windows":
if _, err := exec.LookPath("npm"); err != nil {
return fmt.Errorf("opencode is not installed and required dependencies are missing\n\nInstall the following first:\n npm (Node.js): https://nodejs.org/\n\nThen re-run:\n ollama launch opencode")
}
default:
var missing []string
if _, err := exec.LookPath("curl"); err != nil {
missing = append(missing, "curl: https://curl.se/")
}
if _, err := exec.LookPath("bash"); err != nil {
missing = append(missing, "bash: https://www.gnu.org/software/bash/")
}
if len(missing) > 0 {
return fmt.Errorf("opencode is not installed and required dependencies are missing\n\nInstall the following first:\n %s\n\nThen re-run:\n ollama launch opencode", strings.Join(missing, "\n "))
}
}
return nil
}
func openCodeInstallerCommand(goos string) (string, []string, error) {
switch goos {
case "windows":
return "npm", []string{"install", "-g", "opencode-ai@latest"}, nil
case "darwin", "linux":
return "bash", []string{"-c", "set -o pipefail; " + openCodeInstallScript}, nil
default:
return "", nil, fmt.Errorf("unsupported platform for opencode install: %s", goos)
}
}
// resolveContent returns the inline config to send via OPENCODE_CONFIG_CONTENT.
// Returns content built by Edit if available, otherwise builds from model.json
// with the requested model as primary (e.g. re-launch with saved config).
@@ -278,6 +355,25 @@ func buildModelEntries(modelList []LaunchModel) map[string]any {
"output": []string{"text"},
}
}
if model.HasCapability("thinking") {
entry["reasoning"] = true
if openCodeModelSupportsThinkingLevels(model) {
entry["options"] = map[string]any{"reasoningEffort": "medium"}
entry["variants"] = map[string]any{
"low": map[string]any{"reasoningEffort": "low"},
"medium": map[string]any{"reasoningEffort": "medium"},
"high": map[string]any{"reasoningEffort": "high"},
"max": map[string]any{"reasoningEffort": "max"},
}
} else {
entry["variants"] = map[string]any{
"none": map[string]any{"reasoningEffort": "none"},
"low": map[string]any{"disabled": true},
"medium": map[string]any{"disabled": true},
"high": map[string]any{"disabled": true},
}
}
}
if model.MaxOutputTokens > 0 {
limit := make(map[string]any)
if model.ContextLength > 0 {
@@ -290,3 +386,20 @@ func buildModelEntries(modelList []LaunchModel) map[string]any {
}
return models
}
func openCodeModelSupportsThinkingLevels(model LaunchModel) bool {
for _, family := range append([]string{model.Details.Family}, model.Details.Families...) {
if normalizeOpenCodeModelFamily(family) == "gptoss" {
return true
}
}
return strings.Contains(normalizeOpenCodeModelFamily(model.Name), "gptoss")
}
func normalizeOpenCodeModelFamily(s string) string {
s = strings.ToLower(s)
s = strings.ReplaceAll(s, "-", "")
s = strings.ReplaceAll(s, "_", "")
return s
}
+307
View File
@@ -6,8 +6,10 @@ import (
"os"
"path/filepath"
"runtime"
"strings"
"testing"
"github.com/ollama/ollama/api"
"github.com/ollama/ollama/types/model"
)
@@ -174,6 +176,54 @@ func TestOpenCodeEdit(t *testing.T) {
t.Fatalf("modalities.output = %v, want [text]", output)
}
})
t.Run("thinking model gets on off reasoning variants", func(t *testing.T) {
models := buildModelEntries([]LaunchModel{{Name: "thinking-model", Capabilities: []model.Capability{model.CapabilityThinking}}})
entry, _ := models["thinking-model"].(map[string]any)
if entry["reasoning"] != true {
t.Fatalf("reasoning = %v, want true", entry["reasoning"])
}
variants, _ := entry["variants"].(map[string]any)
none, _ := variants["none"].(map[string]any)
if none["reasoningEffort"] != "none" {
t.Fatalf("variants.none.reasoningEffort = %v, want none", none["reasoningEffort"])
}
for _, level := range []string{"low", "medium", "high"} {
variant, _ := variants[level].(map[string]any)
if variant["disabled"] != true {
t.Fatalf("variants.%s.disabled = %v, want true", level, variant["disabled"])
}
}
})
t.Run("gpt oss gets reasoning level variants", func(t *testing.T) {
models := buildModelEntries([]LaunchModel{{Name: "gpt-oss:120b-cloud", Capabilities: []model.Capability{model.CapabilityThinking}}})
entry, _ := models["gpt-oss:120b-cloud"].(map[string]any)
options, _ := entry["options"].(map[string]any)
if options["reasoningEffort"] != "medium" {
t.Fatalf("options.reasoningEffort = %v, want medium", options["reasoningEffort"])
}
variants, _ := entry["variants"].(map[string]any)
for _, level := range []string{"low", "medium", "high", "max"} {
variant, _ := variants[level].(map[string]any)
if variant["reasoningEffort"] != level {
t.Fatalf("variants.%s.reasoningEffort = %v, want %s", level, variant["reasoningEffort"], level)
}
}
})
t.Run("gpt oss family gets reasoning level variants", func(t *testing.T) {
models := buildModelEntries([]LaunchModel{{Name: "reasoning-model", Capabilities: []model.Capability{model.CapabilityThinking}, Details: api.ModelDetails{Families: []string{"gptoss"}}}})
entry, _ := models["reasoning-model"].(map[string]any)
variants, _ := entry["variants"].(map[string]any)
max, _ := variants["max"].(map[string]any)
if max["reasoningEffort"] != "max" {
t.Fatalf("variants.max.reasoningEffort = %v, want max", max["reasoningEffort"])
}
})
}
func TestBuildModelEntries(t *testing.T) {
@@ -289,12 +339,16 @@ func TestLookupCloudModelLimit(t *testing.T) {
}
func TestFindOpenCode(t *testing.T) {
oldGOOS := openCodeGOOS
t.Cleanup(func() { openCodeGOOS = oldGOOS })
t.Run("fallback to ~/.opencode/bin", func(t *testing.T) {
tmpDir := t.TempDir()
setTestHome(t, tmpDir)
// Ensure opencode is not on PATH
t.Setenv("PATH", tmpDir)
openCodeGOOS = runtime.GOOS
// Without the fallback binary, findOpenCode should fail
if _, ok := findOpenCode(); ok {
@@ -322,6 +376,259 @@ func TestFindOpenCode(t *testing.T) {
})
}
func TestEnsureOpenCodeInstalled(t *testing.T) {
oldGOOS := openCodeGOOS
t.Cleanup(func() { openCodeGOOS = oldGOOS })
withConfirm := func(t *testing.T, fn func(prompt string) (bool, error)) {
t.Helper()
oldConfirm := DefaultConfirmPrompt
DefaultConfirmPrompt = func(prompt string, options ConfirmOptions) (bool, error) {
return fn(prompt)
}
t.Cleanup(func() { DefaultConfirmPrompt = oldConfirm })
}
t.Run("already installed", func(t *testing.T) {
setTestHome(t, t.TempDir())
tmpDir := t.TempDir()
t.Setenv("PATH", tmpDir)
openCodeGOOS = runtime.GOOS
writeFakeBinary(t, tmpDir, "opencode")
withConfirm(t, func(prompt string) (bool, error) {
t.Fatalf("did not expect prompt, got %q", prompt)
return false, nil
})
bin, err := ensureOpenCodeInstalled()
if err != nil {
t.Fatalf("ensureOpenCodeInstalled() error = %v", err)
}
if filepath.Base(bin) == "" {
t.Fatalf("expected opencode binary path, got %q", bin)
}
})
t.Run("missing dependencies", func(t *testing.T) {
setTestHome(t, t.TempDir())
t.Setenv("PATH", t.TempDir())
openCodeGOOS = "linux"
withConfirm(t, func(prompt string) (bool, error) {
t.Fatalf("did not expect prompt, got %q", prompt)
return false, nil
})
_, err := ensureOpenCodeInstalled()
if err == nil || !strings.Contains(err.Error(), "required dependencies are missing") {
t.Fatalf("expected missing dependency error, got %v", err)
}
})
t.Run("missing and user declines install", func(t *testing.T) {
setTestHome(t, t.TempDir())
tmpDir := t.TempDir()
t.Setenv("PATH", tmpDir)
openCodeGOOS = "linux"
writeFakeBinary(t, tmpDir, "curl")
writeFakeBinary(t, tmpDir, "bash")
withConfirm(t, func(prompt string) (bool, error) {
if !strings.Contains(prompt, "OpenCode is not installed.") {
t.Fatalf("unexpected prompt: %q", prompt)
}
return false, nil
})
_, err := ensureOpenCodeInstalled()
if err == nil || !strings.Contains(err.Error(), "installation cancelled") {
t.Fatalf("expected cancellation error, got %v", err)
}
})
t.Run("missing and user confirms unix install succeeds", func(t *testing.T) {
if runtime.GOOS == "windows" {
t.Skip("uses POSIX shell fake binaries")
}
homeDir := t.TempDir()
setTestHome(t, homeDir)
tmpDir := t.TempDir()
t.Setenv("PATH", tmpDir)
openCodeGOOS = "linux"
writeFakeBinary(t, tmpDir, "curl")
installLog := filepath.Join(tmpDir, "bash.log")
opencodePath := filepath.Join(homeDir, ".opencode", "bin", "opencode")
bashScript := fmt.Sprintf(`#!/bin/sh
echo "$@" >> %q
if [ "$1" = "-c" ]; then
/bin/mkdir -p %q
/bin/cat > %q <<'EOS'
#!/bin/sh
exit 0
EOS
/bin/chmod +x %q
fi
exit 0
`, installLog, filepath.Dir(opencodePath), opencodePath, opencodePath)
if err := os.WriteFile(filepath.Join(tmpDir, "bash"), []byte(bashScript), 0o755); err != nil {
t.Fatalf("failed to write fake bash: %v", err)
}
withConfirm(t, func(prompt string) (bool, error) {
return true, nil
})
bin, err := ensureOpenCodeInstalled()
if err != nil {
t.Fatalf("ensureOpenCodeInstalled() error = %v", err)
}
if bin != opencodePath {
t.Fatalf("bin = %q, want %q", bin, opencodePath)
}
logData, err := os.ReadFile(installLog)
if err != nil {
t.Fatalf("failed to read install log: %v", err)
}
if !strings.Contains(string(logData), openCodeInstallScript) {
t.Fatalf("expected opencode install script in log, got:\n%s", string(logData))
}
})
t.Run("missing and user confirms windows install succeeds", func(t *testing.T) {
if runtime.GOOS == "windows" {
t.Skip("uses POSIX shell fake binaries")
}
homeDir := t.TempDir()
setTestHome(t, homeDir)
tmpDir := t.TempDir()
t.Setenv("PATH", tmpDir)
openCodeGOOS = "windows"
installLog := filepath.Join(tmpDir, "npm.log")
opencodePath := filepath.Join(homeDir, ".opencode", "bin", "opencode.exe")
npmScript := fmt.Sprintf(`#!/bin/sh
echo "$@" >> %q
/bin/mkdir -p %q
/bin/cat > %q <<'EOS'
@echo off
exit /b 0
EOS
/bin/chmod +x %q
exit 0
`, installLog, filepath.Dir(opencodePath), opencodePath, opencodePath)
if err := os.WriteFile(filepath.Join(tmpDir, "npm"), []byte(npmScript), 0o755); err != nil {
t.Fatalf("failed to write fake npm: %v", err)
}
withConfirm(t, func(prompt string) (bool, error) {
return true, nil
})
bin, err := ensureOpenCodeInstalled()
if err != nil {
t.Fatalf("ensureOpenCodeInstalled() error = %v", err)
}
if bin != opencodePath {
t.Fatalf("bin = %q, want %q", bin, opencodePath)
}
logData, err := os.ReadFile(installLog)
if err != nil {
t.Fatalf("failed to read install log: %v", err)
}
if !strings.Contains(string(logData), "install -g opencode-ai@latest") {
t.Fatalf("expected npm install command in log, got:\n%s", string(logData))
}
})
t.Run("install command fails", func(t *testing.T) {
if runtime.GOOS == "windows" {
t.Skip("uses POSIX shell fake binaries")
}
setTestHome(t, t.TempDir())
tmpDir := t.TempDir()
t.Setenv("PATH", tmpDir)
openCodeGOOS = "linux"
writeFakeBinary(t, tmpDir, "curl")
if err := os.WriteFile(filepath.Join(tmpDir, "bash"), []byte("#!/bin/sh\nexit 1\n"), 0o755); err != nil {
t.Fatalf("failed to write fake bash: %v", err)
}
withConfirm(t, func(prompt string) (bool, error) {
return true, nil
})
_, err := ensureOpenCodeInstalled()
if err == nil || !strings.Contains(err.Error(), "failed to install opencode") {
t.Fatalf("expected install failure error, got %v", err)
}
})
}
func TestOpenCodeInstallerCommand(t *testing.T) {
tests := []struct {
name string
goos string
wantBin string
wantParts []string
wantErr bool
}{
{
name: "linux",
goos: "linux",
wantBin: "bash",
wantParts: []string{"-c", "set -o pipefail", "https://opencode.ai/install"},
},
{
name: "darwin",
goos: "darwin",
wantBin: "bash",
wantParts: []string{"-c", "set -o pipefail", "https://opencode.ai/install"},
},
{
name: "windows",
goos: "windows",
wantBin: "npm",
wantParts: []string{"install", "-g", "opencode-ai@latest"},
},
{
name: "unsupported",
goos: "plan9",
wantErr: true,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
bin, args, err := openCodeInstallerCommand(tt.goos)
if tt.wantErr {
if err == nil {
t.Fatal("expected error")
}
return
}
if err != nil {
t.Fatalf("openCodeInstallerCommand() error = %v", err)
}
if bin != tt.wantBin {
t.Fatalf("bin = %q, want %q", bin, tt.wantBin)
}
joined := strings.Join(args, " ")
for _, want := range tt.wantParts {
if !strings.Contains(joined, want) {
t.Fatalf("args %q missing %q", joined, want)
}
}
})
}
}
// Verify that the BackfillsCloudModelLimitOnExistingEntry test from the old
// file-based approach is covered by the new inline config approach.
func TestOpenCodeEdit_CloudModelLimitStructure(t *testing.T) {
+2 -2
View File
@@ -351,7 +351,7 @@ func npmArgs(prefix string, args ...string) []string {
}
func ensurePiWebSearchPackage(bin string) {
if !shouldManagePiWebSearch() {
if !shouldManageOllamaWebSearch() {
fmt.Fprintf(os.Stderr, "%sCloud is disabled; skipping %s setup.%s\n", ansiGray, piWebSearchPkg, ansiReset)
return
}
@@ -395,7 +395,7 @@ func ensurePiWebSearchPackage(bin string) {
fmt.Fprintf(os.Stderr, "%s ✓ Updated %s%s\n", ansiGreen, piWebSearchPkg, ansiReset)
}
func shouldManagePiWebSearch() bool {
func shouldManageOllamaWebSearch() bool {
client, err := api.ClientFromEnvironment()
if err != nil {
return true
+39 -5
View File
@@ -33,7 +33,7 @@ type IntegrationInfo struct {
Description string
}
var launcherIntegrationOrder = []string{"claude", "codex-app", "hermes", "openclaw", "opencode", "codex", "copilot", "cline", "droid", "pi", "pool", "qwen"}
var launcherIntegrationOrder = []string{"claude", "chatgpt", "hermes", "openclaw", "opencode", "hermes-desktop", "codex", "copilot", "omp", "cline", "droid", "pi", "pool", "qwen"}
var integrationSpecs = []*IntegrationSpec{
{
@@ -45,6 +45,10 @@ var integrationSpecs = []*IntegrationSpec{
_, err := (&Claude{}).findPath()
return err == nil
},
EnsureInstalled: func() error {
_, err := ensureClaudeInstalled()
return err
},
URL: "https://code.claude.com/docs/en/quickstart",
},
},
@@ -91,15 +95,15 @@ var integrationSpecs = []*IntegrationSpec{
},
},
{
Name: "codex-app",
Name: chatGPTIntegrationName,
Runner: &CodexApp{},
Aliases: []string{"codex-desktop", "codex-gui"},
Description: "An AI agent you can delegate real work to, by OpenAI",
Aliases: []string{codexAppIntegrationName, "codex-desktop", "codex-gui"},
Description: "Complete work with ChatGPT",
Install: IntegrationInstallSpec{
CheckInstalled: func() bool {
return codexAppInstalled()
},
URL: "https://developers.openai.com/codex/quickstart",
URL: "https://chatgpt.com/download",
},
},
{
@@ -153,9 +157,25 @@ var integrationSpecs = []*IntegrationSpec{
_, ok := findOpenCode()
return ok
},
EnsureInstalled: func() error {
_, err := ensureOpenCodeInstalled()
return err
},
URL: "https://opencode.ai",
},
},
{
Name: "omp",
Runner: &OMP{},
Description: "AI coding agent with IDE integration",
Install: IntegrationInstallSpec{
CheckInstalled: func() bool {
_, err := (&OMP{}).findPath()
return err == nil
},
URL: "https://omp.sh",
},
},
{
Name: "openclaw",
Runner: &Openclaw{},
@@ -220,6 +240,20 @@ var integrationSpecs = []*IntegrationSpec{
URL: "https://hermes-agent.nousresearch.com/docs/getting-started/installation/",
},
},
{
Name: "hermes-desktop",
Runner: &HermesDesktop{},
Description: "Desktop app for Hermes Agent by Nous Research",
Install: IntegrationInstallSpec{
CheckInstalled: func() bool {
return (&Hermes{}).installed()
},
EnsureInstalled: func() error {
return (&Hermes{}).ensureInstalledFor("hermes-desktop")
},
URL: "https://hermes-agent.nousresearch.com/docs/getting-started/installation/",
},
},
{
Name: "vscode",
Runner: &VSCode{},
+8
View File
@@ -61,6 +61,14 @@ func TestEditorRunsDoNotRewriteConfig(t *testing.T) {
return filepath.Join(home, ".kimi", "config.toml")
},
},
{
name: "omp",
binary: "omp",
runner: &OMP{},
checkPath: func(home string) string {
return filepath.Join(home, ".omp", "agent", "models.yml")
},
},
}
for _, tt := range tests {
+22 -2
View File
@@ -27,10 +27,18 @@ var errCancelled = ErrCancelled
// When set, ConfirmPrompt delegates to it instead of using raw terminal I/O.
var DefaultConfirmPrompt func(prompt string, options ConfirmOptions) (bool, error)
type ConfirmDefault int
const (
ConfirmDefaultYes ConfirmDefault = iota
ConfirmDefaultNo
)
// ConfirmOptions customizes labels for confirmation prompts.
type ConfirmOptions struct {
YesLabel string
NoLabel string
Default ConfirmDefault
}
// SingleSelector is a function type for single item selection.
@@ -111,7 +119,12 @@ func ConfirmPromptWithOptions(prompt string, options ConfirmOptions) (bool, erro
}
defer term.Restore(fd, oldState)
fmt.Fprintf(os.Stderr, "%s (\033[1my\033[0m/n) ", prompt)
defaultNo := options.Default == ConfirmDefaultNo
if defaultNo {
fmt.Fprintf(os.Stderr, "%s (y/\033[1mN\033[0m) ", prompt)
} else {
fmt.Fprintf(os.Stderr, "%s (\033[1my\033[0m/n) ", prompt)
}
buf := make([]byte, 1)
for {
@@ -120,7 +133,14 @@ func ConfirmPromptWithOptions(prompt string, options ConfirmOptions) (bool, erro
}
switch buf[0] {
case 'Y', 'y', 13:
case 'Y', 'y':
fmt.Fprintf(os.Stderr, "yes\r\n")
return true, nil
case 13:
if defaultNo {
fmt.Fprintf(os.Stderr, "no\r\n")
return false, nil
}
fmt.Fprintf(os.Stderr, "yes\r\n")
return true, nil
case 'N', 'n', 27, 3:
+109
View File
@@ -0,0 +1,109 @@
package launch
import (
"fmt"
"os"
"sync"
"time"
)
// SpinnerFrames are the braille spinner frames used by the bubbletea TUIs in
// this codebase (sign-in, upgrade). StartSpinner uses the same frames for its
// fallback so the restart spinner matches the look of those flows.
var SpinnerFrames = []string{"⠋", "⠙", "⠹", "⠸", "⠼", "⠴", "⠦", "⠧", "⠇", "⠏"}
// DefaultSpinner, when set, starts an animated spinner displaying message and
// returns a *Spinner. cmd/cmd.go registers a bubbletea implementation from
// cmd/tui; when unset (or when it returns nil, e.g. no TTY) StartSpinner falls
// back to a simple ANSI spinner using SpinnerFrames.
var DefaultSpinner func(message string) *Spinner
// Spinner is a handle on a running animated spinner. Stop halts the spinner
// and clears its line (it blocks until the spinner has fully stopped and is
// safe to call multiple times). Cancelled returns a channel that is closed if
// the user interrupts the spinner (e.g. with Ctrl+C); wait loops can select on
// it to abort early. For the non-interactive ANSI fallback the channel is never
// closed because Ctrl+C raises SIGINT and terminates the process directly.
type Spinner struct {
stop func()
cancelled chan struct{}
}
// NewSpinner builds a Spinner from a stop function and a cancellation channel.
// It is intended for implementations of DefaultSpinner (e.g. the bubbletea
// spinner in cmd/tui). stop must be safe to call multiple times; cancelled is
// closed by the implementation when the user interrupts the spinner, or left
// open when interruption is handled another way (e.g. SIGINT).
func NewSpinner(stop func(), cancelled chan struct{}) *Spinner {
return &Spinner{stop: stop, cancelled: cancelled}
}
// Stop halts the spinner and clears its line. It is a no-op when the spinner
// already stopped (for example after the user cancelled it).
func (s *Spinner) Stop() {
if s != nil && s.stop != nil {
s.stop()
}
}
// Cancelled returns a channel that is closed when the user interrupts the
// spinner. Callers may select on it to abort a blocking wait.
func (s *Spinner) Cancelled() <-chan struct{} {
if s == nil {
return nil
}
return s.cancelled
}
// StartSpinner begins an animated spinner displaying message and returns a
// *Spinner handle. It uses DefaultSpinner when available, otherwise a simple
// ANSI fallback that renders SpinnerFrames to stderr without requiring a TTY.
func StartSpinner(message string) *Spinner {
if DefaultSpinner != nil {
if s := DefaultSpinner(message); s != nil {
return s
}
}
return defaultSpinner(message)
}
// defaultSpinner renders SpinnerFrames to stderr without requiring a TTY. It
// runs in its own goroutine so it can animate while a caller polls; Stop
// signals the goroutine to exit, waits for it, and clears the spinner line.
func defaultSpinner(message string) *Spinner {
frames := SpinnerFrames
frame := 0
fmt.Fprintf(os.Stderr, "\r\033[90m%s %s\033[0m", message, frames[0])
done := make(chan struct{})
exited := make(chan struct{})
var once sync.Once
go func() {
ticker := time.NewTicker(100 * time.Millisecond)
defer ticker.Stop()
for {
select {
case <-done:
close(exited)
return
case <-ticker.C:
frame++
fmt.Fprintf(os.Stderr, "\r\033[90m%s %s\033[0m", message, frames[frame%len(frames)])
}
}
}()
stop := func() {
once.Do(func() {
close(done)
<-exited
fmt.Fprintf(os.Stderr, "\r\033[K")
})
}
// Ctrl+C in non-raw mode raises SIGINT and terminates the process by
// default (the launch flow installs no SIGINT handler), so this cancelled
// channel is intentionally never closed.
return &Spinner{stop: stop, cancelled: make(chan struct{})}
}
+406
View File
@@ -0,0 +1,406 @@
package chat
import (
"context"
"fmt"
"slices"
"strings"
"time"
tea "github.com/charmbracelet/bubbletea"
coreagent "github.com/ollama/ollama/agent"
)
type chatApprovalChoice struct {
label string
key string
allow bool
allowTools bool
allowAll bool
reason string
}
var chatApprovalChoices = []chatApprovalChoice{
{label: "Approve once", key: "1", allow: true},
{label: "Always allow tool", key: "2", allow: true, allowTools: true},
{label: "Deny", key: "3", reason: "Tool execution denied."},
}
type chatApprovalPrompt struct {
request coreagent.ApprovalRequest
reply chan<- coreagent.Approval
cursor int
}
func (m chatModel) approvalPrompterForRun(controller *chatApprovalController) coreagent.ApprovalPrompter {
if m.opts.ApprovalPrompter != nil {
return m.opts.ApprovalPrompter
}
return controller
}
func (m *chatModel) ensureApprovalState() *coreagent.ApprovalState {
if m.approvalState == nil {
m.approvalState = &coreagent.ApprovalState{}
m.approvalState.Set(m.defaultAllowAll, nil)
}
return m.approvalState
}
func (m *chatModel) resetApprovalState() {
m.approvalState = &coreagent.ApprovalState{}
m.approvalState.Set(m.defaultAllowAll, nil)
}
func (m chatModel) allowAllToolsEnabled() bool {
if m.approvalState == nil {
return m.defaultAllowAll
}
return m.approvalState.AllGranted()
}
func (m *chatModel) setAllowAllTools(allowAll bool) {
if allowAll {
m.ensureApprovalState().GrantAll()
} else {
m.ensureApprovalState().Set(false, nil)
}
m.opts.AllowAllTools = allowAll
}
func (m *chatModel) openApprovalPrompt(msg chatApprovalPromptMsg) {
m.approvalPrompt = &chatApprovalPrompt{request: msg.request, reply: msg.reply}
m.status = "approval required"
m.thinking = false
m.thinkingTokens = 0
m.upsertApprovalToolEntries(msg.request)
}
func (m *chatModel) togglePermissionMode() (tea.Model, tea.Cmd) {
m.setAllowAllTools(!m.allowAllToolsEnabled())
if m.allowAllToolsEnabled() {
m.permissionNotice = "full access enabled"
m.status = "full access enabled"
if m.approvalPrompt != nil {
updated, cmd := m.resolveApprovalPrompt(chatApprovalChoice{allow: true, allowAll: true})
if model, ok := updated.(chatModel); ok {
model.permissionNotice = "full access enabled"
model.status = "full access enabled"
return model, cmd
}
return updated, cmd
}
return *m, nil
}
m.permissionNotice = "review mode enabled"
m.status = "review mode enabled"
return *m, nil
}
func (m *chatModel) upsertApprovalToolEntries(request coreagent.ApprovalRequest) {
for _, call := range request.Calls {
idx := m.findToolEntry(call.ToolCallID)
if idx < 0 {
m.groupCompletedToolHistory()
m.entries = append(m.entries, newChatEntry(chatEntry{role: "tool"}))
idx = len(m.entries) - 1
}
m.entries[idx].detail = call.ToolName
m.entries[idx].label = toolInvocationLabel(call.ToolName, call.Args)
m.entries[idx].status = "approval"
m.entries[idx].toolID = call.ToolCallID
m.entries[idx].args = call.Args
m.entries[idx].startedAt = time.Now()
m.applyToolOutputModeTo(idx)
m.markEntryDirty(idx)
}
}
func (m chatModel) updateApprovalPrompt(msg tea.KeyMsg) (tea.Model, tea.Cmd) {
switch msg.Type {
case tea.KeyLeft, tea.KeyUp:
m.moveApprovalChoice(-1)
case tea.KeyRight, tea.KeyDown, tea.KeyTab:
m.moveApprovalChoice(1)
case tea.KeyRunes:
switch string(msg.Runes) {
case "1", "2", "3":
choice := chatApprovalChoices[int(msg.Runes[0]-'1')]
return m.resolveApprovalPrompt(choice)
}
case tea.KeyEnter:
choice := chatApprovalChoices[clamp(m.approvalPrompt.cursor, 0, len(chatApprovalChoices)-1)]
return m.resolveApprovalPrompt(choice)
case tea.KeyEsc, tea.KeyCtrlC:
return m.resolveApprovalPrompt(chatApprovalChoice{reason: "Tool execution denied."})
}
return m, nil
}
func (m *chatModel) moveApprovalChoice(delta int) {
if m.approvalPrompt == nil {
return
}
m.approvalPrompt.cursor = (m.approvalPrompt.cursor + delta) % len(chatApprovalChoices)
if m.approvalPrompt.cursor < 0 {
m.approvalPrompt.cursor += len(chatApprovalChoices)
}
m.markApprovalPromptEntryDirty()
}
func (m *chatModel) markApprovalPromptEntryDirty() {
if m.approvalPrompt == nil {
return
}
for _, call := range m.approvalPrompt.request.Calls {
if idx := m.findToolEntry(call.ToolCallID); idx >= 0 {
m.markEntryDirty(idx)
}
}
}
func (m chatModel) resolveApprovalPrompt(choice chatApprovalChoice) (tea.Model, tea.Cmd) {
if m.approvalPrompt == nil {
return m, nil
}
printedLines := m.flowPrintedLines
var printedTranscript []string
if printedLines > 0 {
printedTranscript = slices.Clone(m.transcriptLines(m.viewWidth()))
}
prompt := m.approvalPrompt
m.approvalPrompt = nil
m.status = "running"
if !choice.allow {
m.status = "denied"
}
if choice.allowAll {
m.setAllowAllTools(true)
}
allowScopes := approvalScopes(prompt.request)
if choice.allowTools {
m.ensureApprovalState().GrantScopes(allowScopes)
}
for _, call := range prompt.request.Calls {
if idx := m.findToolEntry(call.ToolCallID); idx >= 0 && m.entries[idx].status == "approval" {
if !choice.allow {
m.entries[idx].status = "error"
m.entries[idx].err = choice.reason
if m.entries[idx].err == "" {
m.entries[idx].err = "Tool execution denied."
}
} else {
m.entries[idx].status = "queued"
}
m.markEntryDirty(idx)
}
}
result := coreagent.Approval{Allow: choice.allow, AllowAll: choice.allowAll, Reason: choice.reason}
if choice.allowTools {
result.AllowScopes = allowScopes
}
prompt.reply <- result
return m.withFlowTranscriptRefreshAfter(printedTranscript, printedLines, waitForChatMsg(m.events))
}
func (m chatModel) renderApprovalPromptLines(width int) []string {
prompt := m.approvalPrompt
if prompt == nil {
return nil
}
if width <= 0 {
width = 80
}
bodyWidth := max(20, width-2)
var lines []string
if len(prompt.request.Calls) <= 1 {
detail := approvalRequestDetail(prompt.request, bodyWidth)
if detail == "" {
label := "Tool request"
if len(prompt.request.Calls) == 1 {
label = toolDisplayName(prompt.request.Calls[0].ToolName)
}
lines = append(lines, wrapChatText(fmt.Sprintf("%s wants to run", label), width)...)
} else {
lines = append(lines, indentLines(splitRenderedBody(detail), " ")...)
}
lines = append(lines, "")
}
lines = append(lines, indentLines(renderApprovalChoices(prompt.request, prompt.cursor, bodyWidth), " ")...)
return lines
}
func approvalRequestDetail(request coreagent.ApprovalRequest, width int) string {
if len(request.Calls) == 0 {
return ""
}
if len(request.Calls) == 1 {
return approvalToolCallDetail(request.Calls[0], width)
}
lines := make([]string, 0, len(request.Calls))
for _, call := range request.Calls {
lines = append(lines, toolInvocationLabel(call.ToolName, call.Args))
}
return chatMetaStyle.Render(strings.Join(lines, "\n"))
}
func approvalToolCallDetail(call coreagent.ApprovalToolCall, width int) string {
if isShellToolName(call.ToolName) {
command, ok := rawStringArg(call.Args, "command")
if !ok {
return ""
}
return strings.Join(wrapChatText(shellPromptPrefix(call.ToolName)+command, width), "\n")
}
switch call.ToolName {
case "edit":
path, ok := rawStringArg(call.Args, "path")
if !ok {
return ""
}
var lines []string
lines = append(lines, "path: "+path)
if oldText, ok := rawStringArg(call.Args, "old_text"); ok {
lines = append(lines, fmt.Sprintf("old_text: %d chars", len([]rune(oldText))))
}
if newText, ok := rawStringArg(call.Args, "new_text"); ok {
lines = append(lines, fmt.Sprintf("new_text: %d chars", len([]rune(newText))))
}
return chatMetaStyle.Render(strings.Join(lines, "\n"))
default:
if len(call.Args) == 0 {
return ""
}
return strings.Join(renderToolCallArgs(call.Args, width), "\n")
}
}
func renderApprovalChoices(request coreagent.ApprovalRequest, cursor int, width int) []string {
var lines []string
for i, choice := range chatApprovalChoices {
label := choice.key + ". " + approvalChoiceLabel(choice, request)
wrapped := wrapChatText(label, max(20, width-2))
if i == clamp(cursor, 0, len(chatApprovalChoices)-1) {
for j, line := range wrapped {
if j == 0 {
lines = append(lines, chatPickerSelectedStyle.Render("> "+line))
} else {
lines = append(lines, chatPickerSelectedStyle.Render(" "+line))
}
}
} else {
for _, line := range wrapped {
lines = append(lines, chatPickerTextStyle.Render(" "+line))
}
}
}
return lines
}
func approvalChoiceLabel(choice chatApprovalChoice, request coreagent.ApprovalRequest) string {
if !choice.allowTools {
return choice.label
}
scopes := approvalScopes(request)
if len(scopes) == 1 {
call := approvalCallForScope(request, scopes[0])
if isShellToolName(call.ToolName) {
if command, ok := rawStringArg(call.Args, "command"); ok && strings.TrimSpace(command) != "" {
return "Always allow this command"
}
}
return "Always allow " + toolDisplayName(call.ToolName)
}
return "Always allow these requests"
}
func approvalScopes(request coreagent.ApprovalRequest) []string {
seen := make(map[string]bool, len(request.Calls))
var scopes []string
for _, call := range request.Calls {
scope := approvalScope(call)
if scope == "" || seen[scope] {
continue
}
seen[scope] = true
scopes = append(scopes, scope)
}
return scopes
}
func approvalCallForScope(request coreagent.ApprovalRequest, scope string) coreagent.ApprovalToolCall {
for _, call := range request.Calls {
if approvalScope(call) == scope {
return call
}
}
return coreagent.ApprovalToolCall{}
}
func approvalScope(call coreagent.ApprovalToolCall) string {
if scope := strings.TrimSpace(call.ApprovalScope); scope != "" {
return scope
}
return strings.TrimSpace(call.ToolName)
}
type chatApprovalPrompter struct {
ch chan<- tea.Msg
}
func (p chatApprovalPrompter) PromptApproval(ctx context.Context, request coreagent.ApprovalRequest) (coreagent.Approval, error) {
reply := make(chan coreagent.Approval, 1)
select {
case p.ch <- chatApprovalPromptMsg{request: request, reply: reply}:
case <-ctx.Done():
return coreagent.Approval{Reason: "Tool approval canceled."}, nil
}
select {
case result := <-reply:
return result, nil
case <-ctx.Done():
return coreagent.Approval{Reason: "Tool approval canceled."}, nil
}
}
type chatApprovalController struct {
ch chan<- tea.Msg
state *coreagent.ApprovalState
}
func newChatApprovalController(ch chan<- tea.Msg, state *coreagent.ApprovalState) *chatApprovalController {
return &chatApprovalController{
ch: ch,
state: state,
}
}
func (c *chatApprovalController) PromptApproval(ctx context.Context, request coreagent.ApprovalRequest) (coreagent.Approval, error) {
if result, ok := c.preapproved(request); ok {
return result, nil
}
return chatApprovalPrompter{ch: c.ch}.PromptApproval(ctx, request)
}
func (c *chatApprovalController) preapproved(request coreagent.ApprovalRequest) (coreagent.Approval, bool) {
if c == nil {
return coreagent.Approval{}, false
}
if c.state.AllGranted() {
return coreagent.Approval{Allow: true, AllowAll: true}, true
}
scopes := approvalScopes(request)
if len(scopes) == 0 {
return coreagent.Approval{}, false
}
for _, scope := range scopes {
if !c.state.Allows(scope) {
return coreagent.Approval{}, false
}
}
return coreagent.Approval{Allow: true, AllowScopes: scopes}, true
}
+495
View File
@@ -0,0 +1,495 @@
package chat
import (
"context"
"strings"
"testing"
"time"
tea "github.com/charmbracelet/bubbletea"
coreagent "github.com/ollama/ollama/agent"
)
func testApprovalRequest() coreagent.ApprovalRequest {
return coreagent.ApprovalRequest{
WorkingDir: "/repo",
Calls: []coreagent.ApprovalToolCall{{
ToolCallID: "call-1",
ToolName: "edit",
Args: map[string]any{"path": "note.txt"},
ApprovalScope: "edit",
}},
}
}
func testApprovalState(allowAll bool, scopes map[string]bool) *coreagent.ApprovalState {
state := &coreagent.ApprovalState{}
state.Set(allowAll, scopes)
return state
}
func TestChatApprovalApprovesOnce(t *testing.T) {
reply := make(chan coreagent.Approval, 1)
m := chatModel{
approvalPrompt: &chatApprovalPrompt{
request: testApprovalRequest(),
reply: reply,
},
events: make(chan tea.Msg),
}
updated, cmd := m.updateApprovalPrompt(tea.KeyMsg{Type: tea.KeyEnter})
if cmd == nil {
t.Fatal("approval should resume waiting for agent events")
}
fm := updated.(chatModel)
if fm.approvalPrompt != nil {
t.Fatal("approval prompt should close")
}
result := <-reply
if !result.Allow || result.AllowAll {
t.Fatalf("approval = %#v, want allow once", result)
}
}
func TestChatApprovalAllowsTool(t *testing.T) {
reply := make(chan coreagent.Approval, 1)
m := chatModel{
approvalPrompt: &chatApprovalPrompt{
request: testApprovalRequest(),
reply: reply,
cursor: 1,
},
events: make(chan tea.Msg),
}
updated, _ := m.updateApprovalPrompt(tea.KeyMsg{Type: tea.KeyEnter})
fm := updated.(chatModel)
if fm.allowAllToolsEnabled() {
t.Fatal("allowing a tool should not enable full access")
}
if !fm.approvalState.Allows("edit") {
t.Fatal("edit scope was not saved")
}
result := <-reply
if !result.Allow || result.AllowAll || len(result.AllowScopes) != 1 || result.AllowScopes[0] != "edit" {
t.Fatalf("approval = %#v, want per-tool approval", result)
}
}
func TestChatApprovalLabelsSecondChoiceAsPerTool(t *testing.T) {
lines := stripANSI(strings.Join(renderApprovalChoices(testApprovalRequest(), 1, 80), "\n"))
if !strings.Contains(lines, "2. Always allow Edit") {
t.Fatalf("approval choices = %q, want per-tool option", lines)
}
if strings.Contains(lines, "Approve all") {
t.Fatalf("approval choices = %q, should not offer approve all as option 2", lines)
}
}
func TestChatApprovalLabelsShellChoiceAsCommandScoped(t *testing.T) {
request := coreagent.ApprovalRequest{
WorkingDir: "/repo",
Calls: []coreagent.ApprovalToolCall{{
ToolCallID: "call-1",
ToolName: "bash",
Args: map[string]any{"command": "pwd"},
ApprovalScope: "bash\x00pwd",
}},
}
lines := stripANSI(strings.Join(renderApprovalChoices(request, 1, 80), "\n"))
if !strings.Contains(lines, "2. Always allow this command") {
t.Fatalf("approval choices = %q, want command-scoped option", lines)
}
if strings.Contains(lines, "Always allow Bash") {
t.Fatalf("approval choices = %q, should not offer top-level Bash approval", lines)
}
}
func TestChatApprovalUsesShellNameForPermissionPrompt(t *testing.T) {
request := coreagent.ApprovalRequest{
WorkingDir: "/repo",
Calls: []coreagent.ApprovalToolCall{{
ToolCallID: "call-1",
ToolName: "bash",
Args: map[string]any{"command": "pwd"},
ApprovalScope: "bash\x00pwd",
}},
}
detail := stripANSI(approvalRequestDetail(request, 80))
if !strings.Contains(detail, "$ pwd") {
t.Fatalf("approval detail should show command prompt, got %q", detail)
}
m := chatModel{}
m.upsertApprovalToolEntries(request)
if len(m.entries) != 1 {
t.Fatalf("entries = %#v", m.entries)
}
line := stripANSI(toolStatusLine(m.entries[0]))
if !strings.Contains(line, `Bash("pwd")`) || !strings.Contains(line, "needs approval") {
t.Fatalf("approval status line = %q", line)
}
}
func TestChatApprovalPromptOmitsDuplicateBatchDetails(t *testing.T) {
request := coreagent.ApprovalRequest{
WorkingDir: "/repo",
Calls: []coreagent.ApprovalToolCall{
{
ToolCallID: "call-1",
ToolName: "bash",
Args: map[string]any{"command": "git rev-parse --abbrev-ref HEAD"},
ApprovalScope: "bash\x00git rev-parse --abbrev-ref HEAD",
},
{
ToolCallID: "call-2",
ToolName: "bash",
Args: map[string]any{"command": "git branch -a"},
ApprovalScope: "bash\x00git branch -a",
},
},
}
m := chatModel{
approvalPrompt: &chatApprovalPrompt{request: request},
}
lines := stripANSI(strings.Join(m.renderApprovalPromptLines(120), "\n"))
if strings.Contains(lines, `Bash("git rev-parse --abbrev-ref HEAD")`) || strings.Contains(lines, `Bash("git branch -a")`) {
t.Fatalf("batched approval prompt should not duplicate visible tool rows:\n%s", lines)
}
for _, want := range []string{"1. Approve once", "2. Always allow these requests", "3. Deny"} {
if !strings.Contains(lines, want) {
t.Fatalf("batched approval prompt missing %q:\n%s", want, lines)
}
}
}
func TestChatApprovalKeepsQueuedBatchCallsVisible(t *testing.T) {
reply := make(chan coreagent.Approval, 1)
request := coreagent.ApprovalRequest{
WorkingDir: "/repo",
Calls: []coreagent.ApprovalToolCall{
{
ToolCallID: "call-1",
ToolName: "bash",
Args: map[string]any{"command": "git rev-parse --abbrev-ref HEAD"},
ApprovalScope: "bash\x00git rev-parse --abbrev-ref HEAD",
},
{
ToolCallID: "call-2",
ToolName: "bash",
Args: map[string]any{"command": "git branch -a"},
ApprovalScope: "bash\x00git branch -a",
},
},
}
m := chatModel{
running: true,
events: make(chan tea.Msg),
}
m.openApprovalPrompt(chatApprovalPromptMsg{request: request, reply: reply})
updated, _ := m.resolveApprovalPrompt(chatApprovalChoice{allow: true})
m = updated.(chatModel)
m.applyAgentEvent(coreagent.Event{
Type: coreagent.EventToolStarted,
ToolCallID: "call-1",
ToolName: "bash",
Args: request.Calls[0].Args,
})
m.applyAgentEvent(coreagent.Event{
Type: coreagent.EventToolFinished,
ToolCallID: "call-1",
ToolName: "bash",
Args: request.Calls[0].Args,
Content: "parth-agent-tui\n",
})
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventMessageDelta, Content: "I'm on branch parth-agent-tui."})
transcript := stripANSI(m.renderTranscript(180))
for _, want := range []string{
`Bash("git rev-parse --abbrev-ref HEAD")`,
`Bash("git branch -a")`,
"I'm on branch parth-agent-tui.",
} {
if !strings.Contains(transcript, want) {
t.Fatalf("transcript missing %q:\n%s", want, transcript)
}
}
}
func TestChatApprovalPromptRepaintsFlowTranscript(t *testing.T) {
reply := make(chan coreagent.Approval, 1)
request := coreagent.ApprovalRequest{
WorkingDir: "/repo",
Calls: []coreagent.ApprovalToolCall{
{
ToolCallID: "call-1",
ToolName: "web_fetch",
Args: map[string]any{"url": "https://parthsareen.com/"},
ApprovalScope: "web_fetch",
},
{
ToolCallID: "call-2",
ToolName: "web_fetch",
Args: map[string]any{"url": "https://github.com/ParthSareen"},
ApprovalScope: "web_fetch",
},
},
}
m := chatModel{
running: true,
width: 160,
flowPrintedLines: 1,
entries: []chatEntry{
{role: "user", content: "research parth"},
},
}
updated, cmd := m.Update(chatApprovalPromptMsg{request: request, reply: reply})
if cmd == nil {
t.Fatal("opening approval should repaint flow transcript")
}
fm := updated.(chatModel)
transcript := stripANSI(fm.renderTranscript(160))
for _, want := range []string{
`Web Fetch("https://parthsareen.com/") needs approval`,
`Web Fetch("https://github.com/ParthSareen") needs approval`,
} {
if !strings.Contains(transcript, want) {
t.Fatalf("transcript missing %q:\n%s", want, transcript)
}
}
}
func TestChatApprovalResolutionRepaintsFlowTranscript(t *testing.T) {
reply := make(chan coreagent.Approval, 1)
request := coreagent.ApprovalRequest{
WorkingDir: "/repo",
Calls: []coreagent.ApprovalToolCall{
{
ToolCallID: "call-1",
ToolName: "web_fetch",
Args: map[string]any{"url": "https://parthsareen.com/"},
ApprovalScope: "web_fetch",
},
{
ToolCallID: "call-2",
ToolName: "web_fetch",
Args: map[string]any{"url": "https://github.com/ParthSareen"},
ApprovalScope: "web_fetch",
},
},
}
m := chatModel{
running: true,
width: 160,
events: make(chan tea.Msg),
entries: []chatEntry{
{role: "user", content: "research parth"},
},
}
m.openApprovalPrompt(chatApprovalPromptMsg{request: request, reply: reply})
printed := len(m.transcriptLines(160))
m.flowPrintedLines = printed
updated, cmd := m.resolveApprovalPrompt(chatApprovalChoice{allow: true})
if cmd == nil {
t.Fatal("approval resolution should keep waiting for agent events")
}
fm := updated.(chatModel)
if fm.flowPrintedLines >= printed {
t.Fatalf("approval resolution should repaint and hold queued rows, flowPrintedLines = %d, was %d", fm.flowPrintedLines, printed)
}
if result := <-reply; !result.Allow {
t.Fatalf("approval = %#v, want allow", result)
}
}
func TestChatApprovalBatchCollapsesAtNextToolBoundary(t *testing.T) {
reply := make(chan coreagent.Approval, 1)
request := coreagent.ApprovalRequest{
WorkingDir: "/repo",
Calls: []coreagent.ApprovalToolCall{
{
ToolCallID: "call-1",
ToolName: "bash",
Args: map[string]any{"command": "git rev-parse --abbrev-ref HEAD"},
ApprovalScope: "bash\x00git rev-parse --abbrev-ref HEAD",
},
{
ToolCallID: "call-2",
ToolName: "bash",
Args: map[string]any{"command": "git branch -a"},
ApprovalScope: "bash\x00git branch -a",
},
},
}
m := chatModel{
running: true,
events: make(chan tea.Msg),
}
m.openApprovalPrompt(chatApprovalPromptMsg{request: request, reply: reply})
updated, _ := m.resolveApprovalPrompt(chatApprovalChoice{allow: true})
m = updated.(chatModel)
for _, call := range request.Calls {
m.applyAgentEvent(coreagent.Event{
Type: coreagent.EventToolStarted,
ToolCallID: call.ToolCallID,
ToolName: call.ToolName,
Args: call.Args,
})
m.applyAgentEvent(coreagent.Event{
Type: coreagent.EventToolFinished,
ToolCallID: call.ToolCallID,
ToolName: call.ToolName,
Args: call.Args,
Content: "ok\n",
})
}
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventMessageDelta, Content: "I'm on branch parth-agent-tui."})
transcript := stripANSI(m.renderTranscript(180))
if strings.Contains(transcript, "Ran 2 commands") {
t.Fatalf("completed batch should stay expanded until the next tool boundary:\n%s", transcript)
}
if !strings.Contains(transcript, `Bash("git branch -a")`) {
t.Fatalf("completed batch should keep concrete command rows before the next boundary:\n%s", transcript)
}
m.applyAgentEvent(coreagent.Event{
Type: coreagent.EventToolStarted,
ToolCallID: "call-3",
ToolName: "bash",
Args: map[string]any{"command": "git status --short"},
})
transcript = stripANSI(m.renderTranscript(180))
if !strings.Contains(transcript, "Ran 2 commands") {
t.Fatalf("completed batch should collapse when a new tool starts:\n%s", transcript)
}
if !strings.Contains(transcript, `Bash("git status --short")`) {
t.Fatalf("new running command should remain concrete after previous batch collapses:\n%s", transcript)
}
}
func TestChatApprovalPrompterCancels(t *testing.T) {
ctx, cancel := context.WithCancel(context.Background())
cancel()
result, err := (chatApprovalPrompter{ch: make(chan tea.Msg)}).PromptApproval(ctx, testApprovalRequest())
if err != nil {
t.Fatal(err)
}
if result.Allow || result.Reason == "" {
t.Fatalf("approval = %#v, want canceled denial", result)
}
}
func TestChatApprovalControllerAutoApprovesAfterFullAccessToggle(t *testing.T) {
events := make(chan tea.Msg, 1)
state := testApprovalState(false, nil)
controller := newChatApprovalController(events, state)
state.GrantAll()
result, err := controller.PromptApproval(context.Background(), testApprovalRequest())
if err != nil {
t.Fatal(err)
}
if !result.Allow || !result.AllowAll {
t.Fatalf("approval = %#v, want full-access approval", result)
}
select {
case msg := <-events:
t.Fatalf("approval UI event should not be sent after full access toggle: %#v", msg)
default:
}
}
func TestChatPermissionToggleSyncsRunningApprovalController(t *testing.T) {
events := make(chan tea.Msg, 1)
state := testApprovalState(false, nil)
m := chatModel{
approvalState: state,
approvalController: newChatApprovalController(events, state),
}
updated, _ := m.togglePermissionMode()
fm := updated.(chatModel)
result, err := fm.approvalController.PromptApproval(context.Background(), testApprovalRequest())
if err != nil {
t.Fatal(err)
}
if !result.Allow || !result.AllowAll {
t.Fatalf("approval = %#v, want full-access approval", result)
}
}
func TestChatPermissionToggleFromFullAccessRequiresReviewInRunningController(t *testing.T) {
events := make(chan tea.Msg, 1)
state := testApprovalState(true, nil)
m := chatModel{
approvalState: state,
approvalController: newChatApprovalController(events, state),
}
updated, _ := m.togglePermissionMode()
fm := updated.(chatModel)
if fm.allowAllToolsEnabled() {
t.Fatal("full access should be disabled")
}
resultCh := make(chan coreagent.Approval, 1)
go func() {
result, err := fm.approvalController.PromptApproval(context.Background(), testApprovalRequest())
if err != nil {
resultCh <- coreagent.Approval{Reason: err.Error()}
return
}
resultCh <- result
}()
select {
case msg := <-events:
prompt, ok := msg.(chatApprovalPromptMsg)
if !ok {
t.Fatalf("event = %#v, want approval prompt", msg)
}
prompt.reply <- coreagent.Approval{Reason: "denied"}
case <-time.After(time.Second):
t.Fatal("expected approval prompt after toggling from full access to review")
}
result := <-resultCh
if result.Allow {
t.Fatalf("approval = %#v, want review prompt result", result)
}
}
func TestChatApprovalPromptSkippedWhenFullAccessEnabledInFlight(t *testing.T) {
reply := make(chan coreagent.Approval, 1)
// Full access is on by the time the buffered approval request reaches the
// UI (toggled after the agent sent the request but before Update ran).
// The stale prompt must not surface; the request is auto-approved.
m := chatModel{approvalState: testApprovalState(true, nil), running: true}
updated, _ := m.Update(chatApprovalPromptMsg{request: testApprovalRequest(), reply: reply})
fm := updated.(chatModel)
if fm.approvalPrompt != nil {
t.Fatalf("approval prompt = %#v, want nil (full access on)", fm.approvalPrompt)
}
if got := fm.status; got == "approval required" {
t.Fatalf("status = %q, should not show approval required", got)
}
select {
case result := <-reply:
if !result.Allow || !result.AllowAll {
t.Fatalf("approval = %#v, want full-access approval", result)
}
default:
t.Fatal("expected auto-approval sent on the reply channel")
}
}
+1226
View File
File diff suppressed because it is too large. Load diff
+66
View File
@@ -0,0 +1,66 @@
package chat
import (
"context"
"errors"
"fmt"
"os/exec"
"runtime"
"strings"
tea "github.com/charmbracelet/bubbletea"
)
type chatClipboardErrorMsg struct {
err error
}
var writeClipboard = writeSystemClipboard
func copyTextCmd(ctx context.Context, text string) tea.Cmd {
return func() tea.Msg {
if err := writeClipboard(ctx, text); err != nil {
return chatClipboardErrorMsg{err: err}
}
return nil
}
}
func writeSystemClipboard(ctx context.Context, text string) error {
if ctx == nil {
ctx = context.Background()
}
switch runtime.GOOS {
case "darwin":
return runClipboardCommand(ctx, text, "pbcopy")
case "windows":
return runClipboardCommand(ctx, text, "clip")
default:
for _, candidate := range []struct {
name string
args []string
}{
{name: "wl-copy"},
{name: "xclip", args: []string{"-selection", "clipboard"}},
{name: "xsel", args: []string{"--clipboard", "--input"}},
} {
if _, err := exec.LookPath(candidate.name); err != nil {
continue
}
return runClipboardCommand(ctx, text, candidate.name, candidate.args...)
}
return errors.New("no clipboard command found")
}
}
func runClipboardCommand(ctx context.Context, text, name string, args ...string) error {
cmd := exec.CommandContext(ctx, name, args...)
cmd.Stdin = strings.NewReader(text)
if output, err := cmd.CombinedOutput(); err != nil {
if len(output) > 0 {
return fmt.Errorf("%s: %w: %s", name, err, strings.TrimSpace(string(output)))
}
return fmt.Errorf("%s: %w", name, err)
}
return nil
}
+434
View File
@@ -0,0 +1,434 @@
package chat
import (
"context"
"errors"
"fmt"
"net/http"
"strings"
"time"
tea "github.com/charmbracelet/bubbletea"
"github.com/ollama/ollama/api"
"github.com/ollama/ollama/cmd/launch"
"github.com/ollama/ollama/internal/modelref"
)
type cloudAuthKind string
const (
cloudAuthSignIn cloudAuthKind = "signin"
cloudAuthUpgrade cloudAuthKind = "upgrade"
cloudAuthChecking cloudAuthKind = "checking"
)
const cloudPlanVerificationUnavailable = "Could not verify Ollama plan. Try again in a moment or use a local model."
// Sign-in/upgrade verification polling bounds. While the check is healthy but
// the user hasn't signed in yet, polling stays prompt so completion is detected
// quickly. When the check itself fails, polling backs off so a down server
// isn't hammered, and gives up after maxPollFailures consecutive errors (or
// pollHardCap elapsed) so the user isn't stuck on a spinner with no recourse
// beyond Esc.
const (
maxPollFailures = 6
pollBackoffBase = 3 * time.Second
pollBackoffCap = 30 * time.Second
pollHardCap = 2 * time.Minute
)
// cloudAuthPrompt is an inline modal that handles sign-in and plan-upgrade
// flows when a user selects a cloud model from the picker.
type cloudAuthPrompt struct {
modelName string
requiredPlan string
signInURL string
upgradeURL string
kind cloudAuthKind
spinner int
openNow bool
polling bool
// pollStarted tracks when sign-in/upgrade verification polling began, for
// the hard-cap timeout. Lazily set on the first poll response.
pollStarted time.Time
// pollFailures counts consecutive verification-check errors; once it
// reaches maxPollFailures the modal gives up and surfaces an error.
pollFailures int
// pollErr holds the last verification error, rendered while retrying.
pollErr string
}
type cloudAuthCheckMsg struct {
err error
signInURL string
}
type cloudModelPreflightMsg struct {
model string
err error
signInURL string
}
type cloudAuthTickMsg struct{}
type cloudAuthPollMsg struct {
done bool
err error
}
func checkCloudModelCmd(ctx context.Context, check func(context.Context, string, string) error, model, requiredPlan string) tea.Cmd {
if check == nil {
return nil
}
return func() tea.Msg {
if ctx == nil {
ctx = context.Background()
}
err := check(ctx, model, requiredPlan)
var signInURL string
if err != nil {
var authErr api.AuthorizationError
if errors.As(err, &authErr) && authErr.SigninURL != "" {
signInURL = authErr.SigninURL
}
}
return cloudAuthCheckMsg{err: err, signInURL: signInURL}
}
}
func cloudModelPreflightCmd(ctx context.Context, opts Options, modelName, requiredPlan string) tea.Cmd {
modelName = strings.TrimSpace(modelName)
if opts.CheckCloudModel == nil || modelName == "" || !modelref.HasExplicitCloudSource(modelName) {
return nil
}
return func() tea.Msg {
if ctx == nil {
ctx = context.Background()
}
plan := strings.TrimSpace(requiredPlan)
if plan == "" && opts.ModelOptions != nil {
models, err := opts.ModelOptions(ctx)
if err == nil {
for _, model := range models {
if strings.EqualFold(strings.TrimSpace(model.Name), modelName) {
plan = strings.TrimSpace(model.RequiredPlan)
break
}
}
}
}
err := opts.CheckCloudModel(ctx, modelName, plan)
return cloudModelPreflightMsg{
model: modelName,
err: err,
signInURL: cloudAuthSignInURL(err),
}
}
}
func cloudAuthSignInURL(err error) string {
if err == nil {
return ""
}
var authErr api.AuthorizationError
if errors.As(err, &authErr) && (authErr.StatusCode == http.StatusUnauthorized || authErr.SigninURL != "") {
return authErr.SigninURL
}
return ""
}
func cloudAuthTickCmd() tea.Cmd {
return tea.Tick(200*time.Millisecond, func(t time.Time) tea.Msg {
return cloudAuthTickMsg{}
})
}
func (m chatModel) updateCloudModelPreflight(msg cloudModelPreflightMsg) (tea.Model, tea.Cmd) {
if msg.model == "" || !strings.EqualFold(strings.TrimSpace(m.opts.Model), strings.TrimSpace(msg.model)) {
return m, nil
}
if msg.err == nil {
if m.status == cloudPlanVerificationUnavailable {
m.status = "ready"
}
return m, nil
}
if msg.signInURL != "" {
return m.startCloudAuthSignIn(msg.model, "", msg.signInURL)
}
m.status = cloudPlanVerificationUnavailable
return m, nil
}
func pollCloudAuthCmd(ctx context.Context, poll func(context.Context) (string, bool, error), delay time.Duration) tea.Cmd {
if poll == nil {
return nil
}
return func() tea.Msg {
if ctx == nil {
ctx = context.Background()
}
// Back off before the next check when the previous one failed. Honor
// context cancellation so an abandoned modal doesn't block on the
// full delay.
if delay > 0 {
timer := time.NewTimer(delay)
defer timer.Stop()
select {
case <-ctx.Done():
case <-timer.C:
}
}
pollCtx, cancel := context.WithTimeout(ctx, 3*time.Second)
defer cancel()
_, done, err := poll(pollCtx)
return cloudAuthPollMsg{done: done, err: err}
}
}
func (m *chatModel) startCloudAuthSignIn(modelName, requiredPlan, signInURL string) (tea.Model, tea.Cmd) {
// When no sign-in URL is available yet, show the "checking" state while
// we verify the plan, rather than rendering a blank "Navigate to:" URL.
kind := cloudAuthSignIn
if signInURL == "" {
kind = cloudAuthChecking
}
m.cloudAuthPrompt = &cloudAuthPrompt{
modelName: modelName,
requiredPlan: requiredPlan,
kind: kind,
signInURL: signInURL,
polling: true,
}
m.status = "cloud-auth"
m.modelPicker = nil
m.modelPickerModels = nil
if m.opts.OpenBrowser != nil && signInURL != "" {
m.opts.OpenBrowser(signInURL)
}
if signInURL == "" {
return m, checkCloudModelCmd(m.ctx, m.opts.CheckCloudModel, modelName, requiredPlan)
}
return m, tea.Batch(cloudAuthTickCmd(), pollCloudAuthCmd(m.ctx, m.opts.PollCloudAuth, 0))
}
func (m *chatModel) startCloudAuthUpgrade(modelName, requiredPlan string) (tea.Model, tea.Cmd) {
m.cloudAuthPrompt = &cloudAuthPrompt{
modelName: modelName,
requiredPlan: requiredPlan,
kind: cloudAuthUpgrade,
upgradeURL: launch.DefaultUpgradeURL,
openNow: true,
}
m.status = "cloud-auth"
m.modelPicker = nil
m.modelPickerModels = nil
return m, nil
}
func (m chatModel) updateCloudAuthPrompt(msg tea.Msg) (tea.Model, tea.Cmd) {
switch msg := msg.(type) {
case cloudAuthCheckMsg:
if msg.err == nil {
// Auth passed — apply the pending model.
return m.completeCloudAuth()
}
// Determine if sign-in or upgrade is needed.
if msg.signInURL != "" {
m.cloudAuthPrompt.kind = cloudAuthSignIn
m.cloudAuthPrompt.signInURL = msg.signInURL
m.cloudAuthPrompt.polling = true
if m.opts.OpenBrowser != nil {
m.opts.OpenBrowser(msg.signInURL)
}
return m, tea.Batch(cloudAuthTickCmd(), pollCloudAuthCmd(m.ctx, m.opts.PollCloudAuth, 0))
}
// Could be a plan upgrade error or unknown error.
m.cloudAuthPrompt = nil
m.openModelOnInit = false
m.status = "ready"
m.entries = append(m.entries, newChatEntry(chatEntry{role: "error", content: fmt.Sprintf("Could not switch model: %v", msg.err), err: msg.err.Error()}))
return m, nil
case cloudAuthTickMsg:
if m.cloudAuthPrompt == nil {
return m, nil
}
m.cloudAuthPrompt.spinner++
return m, cloudAuthTickCmd()
case cloudAuthPollMsg:
if m.cloudAuthPrompt == nil {
return m, nil
}
if msg.done {
// Signed in — re-check auth to see if plan is satisfied.
m.cloudAuthPrompt.polling = false
m.cloudAuthPrompt.pollFailures = 0
m.cloudAuthPrompt.pollErr = ""
return m, checkCloudModelCmd(m.ctx, m.opts.CheckCloudModel, m.cloudAuthPrompt.modelName, m.cloudAuthPrompt.requiredPlan)
}
// Lazily mark the start of the polling window on the first response.
if m.cloudAuthPrompt.pollStarted.IsZero() {
m.cloudAuthPrompt.pollStarted = time.Now()
}
// Hard cap: give up if verification drags on too long for any reason.
if time.Since(m.cloudAuthPrompt.pollStarted) > pollHardCap {
return m.failCloudAuthPoll(errors.New("sign-in is taking longer than expected; check your connection and try again"))
}
if msg.err != nil {
// The verification check itself failed (network down, server 5xx).
// Back off and retry, but give up after a handful of consecutive
// failures so the user isn't stuck on a spinner with no signal.
m.cloudAuthPrompt.pollFailures++
m.cloudAuthPrompt.pollErr = msg.err.Error()
if m.cloudAuthPrompt.pollFailures >= maxPollFailures {
return m.failCloudAuthPoll(fmt.Errorf("couldn't verify sign-in: %w", msg.err))
}
delay := pollBackoffCap
if d := pollBackoffBase << (m.cloudAuthPrompt.pollFailures - 1); d < pollBackoffCap {
delay = d
}
return m, pollCloudAuthCmd(m.ctx, m.opts.PollCloudAuth, delay)
}
// Healthy but not signed in yet — keep polling promptly so sign-in
// completion is detected without added latency.
m.cloudAuthPrompt.pollFailures = 0
m.cloudAuthPrompt.pollErr = ""
return m, pollCloudAuthCmd(m.ctx, m.opts.PollCloudAuth, 0)
case tea.KeyMsg:
if msg.Type == tea.KeyEsc || msg.Type == tea.KeyCtrlC {
m.cloudAuthPrompt = nil
m.pendingModel = ""
m.openModelOnInit = false
m.status = "ready"
return m, nil
}
if m.cloudAuthPrompt.kind == cloudAuthUpgrade && !m.cloudAuthPrompt.polling {
switch msg.Type {
case tea.KeyLeft, tea.KeyRight, tea.KeyTab:
m.cloudAuthPrompt.openNow = !m.cloudAuthPrompt.openNow
case tea.KeyEnter:
if m.cloudAuthPrompt.openNow {
m.cloudAuthPrompt.polling = true
if m.opts.OpenBrowser != nil && m.cloudAuthPrompt.upgradeURL != "" {
m.opts.OpenBrowser(m.cloudAuthPrompt.upgradeURL)
}
return m, tea.Batch(cloudAuthTickCmd(), pollCloudAuthCmd(m.ctx, m.opts.PollCloudAuth, 0))
}
m.cloudAuthPrompt = nil
m.pendingModel = ""
m.openModelOnInit = false
m.status = "ready"
return m, nil
}
}
}
return m, nil
}
// failCloudAuthPoll abandons the sign-in/upgrade verification modal, surfaces
// an error entry to the user, and returns to the ready state so they can
// re-pick a model and retry.
func (m chatModel) failCloudAuthPoll(err error) (tea.Model, tea.Cmd) {
m.cloudAuthPrompt = nil
m.pendingModel = ""
m.openModelOnInit = false
m.status = "ready"
m.entries = append(m.entries, newChatEntry(chatEntry{role: "error", content: fmt.Sprintf("Could not switch model: %v", err), err: err.Error()}))
return m, nil
}
func (m chatModel) completeCloudAuth() (tea.Model, tea.Cmd) {
pending := m.cloudAuthPrompt.modelName
m.cloudAuthPrompt = nil
m.pendingModel = ""
m.modelPicker = nil
m.modelPickerModels = nil
m.openModelOnInit = false
m.status = "ready"
if err := m.applyModelSelection(pending, true); err != nil {
m.entries = append(m.entries, newChatEntry(chatEntry{role: "error", content: fmt.Sprintf("Could not switch model: %v", err), err: err.Error()}))
m.status = "error"
return m, nil
}
return m, m.startModelPreload(pending)
}
func (m chatModel) renderCloudAuthPrompt(width int) string {
if m.cloudAuthPrompt == nil {
return ""
}
if width <= 0 {
width = 80
}
p := m.cloudAuthPrompt
spinnerFrames := []string{"⠋", "⠙", "⠹", "⠸", "⠼", "⠴", "⠦", "⠧", "⠇", "⠏"}
frame := spinnerFrames[p.spinner%len(spinnerFrames)]
var b strings.Builder
switch p.kind {
case cloudAuthChecking:
fmt.Fprintf(&b, "%s Checking %s...\n\n", frame, chatPickerSelectedStyle.Render(p.modelName))
b.WriteString(chatPickerMetaStyle.Render("esc cancel"))
case cloudAuthSignIn:
fmt.Fprintf(&b, "To use %s, please sign in.\n\n", chatPickerSelectedStyle.Render(p.modelName))
b.WriteString("Navigate to:\n")
urlWrap := chatPickerTextStyle
if width > 4 {
urlWrap = chatPickerTextStyle.Width(width - 4)
}
b.WriteString(urlWrap.Render(p.signInURL))
b.WriteString("\n\n")
if p.pollErr != "" {
b.WriteString(chatPickerMetaStyle.Render(frame + " Couldn't verify sign-in: " + p.pollErr + " — retrying..."))
} else {
b.WriteString(chatPickerMetaStyle.Render(frame + " Waiting for sign in to complete..."))
}
b.WriteString("\n\n")
b.WriteString(chatPickerMetaStyle.Render("esc cancel"))
case cloudAuthUpgrade:
fmt.Fprintf(&b, "To use %s, upgrade your Ollama plan.\n\n", chatPickerSelectedStyle.Render(p.modelName))
if !p.polling {
var yesBtn, noBtn string
if p.openNow {
yesBtn = chatPickerSelectedStyle.Render(" Yes ")
noBtn = chatPickerMetaStyle.Render(" No ")
} else {
yesBtn = chatPickerMetaStyle.Render(" Yes ")
noBtn = chatPickerSelectedStyle.Render(" No ")
}
b.WriteString("Open upgrade page now?\n")
b.WriteString(yesBtn + " " + noBtn)
b.WriteString("\n\n")
if !p.openNow {
b.WriteString("Or navigate to:\n")
urlWrap := chatPickerTextStyle
if width > 4 {
urlWrap = chatPickerTextStyle.Width(width - 4)
}
if u := p.upgradeURL; u != "" {
b.WriteString(urlWrap.Render(u))
} else {
b.WriteString(urlWrap.Render(launch.DefaultUpgradeURL))
}
b.WriteString("\n\n")
}
b.WriteString(chatPickerMetaStyle.Render("←/→ navigate • enter confirm • esc cancel"))
} else {
if p.pollErr != "" {
b.WriteString(chatPickerMetaStyle.Render(frame + " Couldn't verify upgrade: " + p.pollErr + " — retrying..."))
} else {
b.WriteString(chatPickerMetaStyle.Render(frame + " Waiting for upgrade to complete..."))
}
b.WriteString("\n\n")
b.WriteString(chatPickerMetaStyle.Render("esc cancel"))
}
}
return b.String()
}
+247
View File
@@ -0,0 +1,247 @@
package chat
import (
"context"
"errors"
"strings"
"testing"
)
func TestCloudAuthTickDoesNotPoll(t *testing.T) {
polls := 0
m := chatModel{
cloudAuthPrompt: &cloudAuthPrompt{polling: true},
opts: Options{
PollCloudAuth: func(context.Context) (string, bool, error) {
polls++
return "", false, nil
},
},
}
updated, cmd := m.updateCloudAuthPrompt(cloudAuthTickMsg{})
m = updated.(chatModel)
if m.cloudAuthPrompt.spinner != 1 {
t.Fatalf("spinner = %d, want 1", m.cloudAuthPrompt.spinner)
}
if polls != 0 {
t.Fatalf("polls = %d, want 0 before running returned tick command", polls)
}
if cmd == nil {
t.Fatal("tick should schedule the next tick")
}
if _, ok := cmd().(cloudAuthTickMsg); !ok {
t.Fatal("tick should schedule another tick, not a poll")
}
if polls != 0 {
t.Fatalf("polls = %d, want 0 after running returned tick command", polls)
}
}
func TestCloudAuthPollSchedulesNextPoll(t *testing.T) {
polls := 0
m := chatModel{
cloudAuthPrompt: &cloudAuthPrompt{polling: true},
opts: Options{
PollCloudAuth: func(context.Context) (string, bool, error) {
polls++
return "", false, nil
},
},
}
_, cmd := m.updateCloudAuthPrompt(cloudAuthPollMsg{})
if cmd == nil {
t.Fatal("poll should schedule the next poll")
}
msg, ok := cmd().(cloudAuthPollMsg)
if !ok {
t.Fatal("poll should schedule another poll, not a tick")
}
if msg.done {
t.Fatal("poll should report not done")
}
if polls != 1 {
t.Fatalf("polls = %d, want 1", polls)
}
}
func TestCloudModelPreflightFailureShowsPlanVerificationNotice(t *testing.T) {
m := chatModel{
opts: Options{
Model: "glm-5.2:cloud",
},
}
updated, cmd := m.updateCloudModelPreflight(cloudModelPreflightMsg{
model: "glm-5.2:cloud",
err: errors.New("temporary network failure"),
})
if cmd != nil {
t.Fatal("transient preflight failure should not start an auth modal")
}
m = updated.(chatModel)
if got := m.status; got != cloudPlanVerificationUnavailable {
t.Fatalf("status = %q", got)
}
if m.cloudAuthPrompt != nil {
t.Fatalf("cloud auth prompt = %#v, want nil", m.cloudAuthPrompt)
}
}
func TestCloudModelPreflightIgnoresStaleModel(t *testing.T) {
m := chatModel{
opts: Options{
Model: "glm-5.2:cloud",
},
status: "ready",
}
updated, _ := m.updateCloudModelPreflight(cloudModelPreflightMsg{
model: "kimi-k2.7-code:cloud",
err: errors.New("temporary network failure"),
})
m = updated.(chatModel)
if got := m.status; got != "ready" {
t.Fatalf("status = %q, want unchanged", got)
}
}
func TestCloudModelPreflightCommandChecksCloudModel(t *testing.T) {
var checkedModel, checkedPlan string
cmd := cloudModelPreflightCmd(context.Background(), Options{
CheckCloudModel: func(_ context.Context, model, requiredPlan string) error {
checkedModel = model
checkedPlan = requiredPlan
return errors.New("temporary network failure")
},
ModelOptions: func(context.Context) ([]ModelOption, error) {
return []ModelOption{{Name: "glm-5.2:cloud", RequiredPlan: "pro", Cloud: true}}, nil
},
}, "glm-5.2:cloud", "")
if cmd == nil {
t.Fatal("cloud preflight command should be scheduled")
}
raw := cmd()
msg, ok := raw.(cloudModelPreflightMsg)
if !ok {
t.Fatalf("message = %T, want cloudModelPreflightMsg", raw)
}
if checkedModel != "glm-5.2:cloud" || checkedPlan != "pro" {
t.Fatalf("checked model/plan = %q/%q", checkedModel, checkedPlan)
}
if msg.model != "glm-5.2:cloud" || msg.err == nil || !strings.Contains(msg.err.Error(), "temporary") {
t.Fatalf("message = %#v", msg)
}
}
func TestCloudAuthPollGivesUpAfterConsecutiveFailures(t *testing.T) {
pollErr := errors.New("whoami: connection refused")
m := chatModel{
cloudAuthPrompt: &cloudAuthPrompt{polling: true, kind: cloudAuthSignIn},
opts: Options{
PollCloudAuth: func(context.Context) (string, bool, error) {
return "", false, pollErr
},
},
}
// The first maxPollFailures-1 failures should keep retrying.
for i := 1; i < maxPollFailures; i++ {
updated, _ := m.updateCloudAuthPrompt(cloudAuthPollMsg{done: false, err: pollErr})
m = updated.(chatModel)
if m.cloudAuthPrompt == nil {
t.Fatalf("failure %d: prompt cleared early", i)
}
if got := m.cloudAuthPrompt.pollFailures; got != i {
t.Fatalf("failure %d: pollFailures = %d, want %d", i, got, i)
}
if m.cloudAuthPrompt.pollErr != pollErr.Error() {
t.Fatalf("failure %d: pollErr = %q, want %q", i, m.cloudAuthPrompt.pollErr, pollErr.Error())
}
}
// The threshold failure gives up: prompt cleared, back to ready, error entry.
updated, cmd := m.updateCloudAuthPrompt(cloudAuthPollMsg{done: false, err: pollErr})
m = updated.(chatModel)
if cmd != nil {
t.Fatalf("threshold failure should not reschedule, got cmd %T", cmd)
}
if m.cloudAuthPrompt != nil {
t.Fatalf("prompt = %#v, want nil after give-up", m.cloudAuthPrompt)
}
if m.status != "ready" {
t.Fatalf("status = %q, want ready", m.status)
}
if len(m.entries) == 0 {
t.Fatal("expected an error entry after give-up")
}
last := m.entries[len(m.entries)-1]
if last.role != "error" || !strings.Contains(last.content, "couldn't verify sign-in") {
t.Fatalf("last entry = %+v, want error containing sign-in failure", last)
}
}
func TestCloudAuthPollResetsFailuresOnHealthyResponse(t *testing.T) {
pollErr := errors.New("whoami: timeout")
m := chatModel{
cloudAuthPrompt: &cloudAuthPrompt{polling: true, kind: cloudAuthSignIn},
opts: Options{
PollCloudAuth: func(context.Context) (string, bool, error) {
return "", false, pollErr
},
},
}
// Accumulate some failures without hitting the threshold.
for range maxPollFailures - 2 {
updated, _ := m.updateCloudAuthPrompt(cloudAuthPollMsg{done: false, err: pollErr})
m = updated.(chatModel)
}
if got := m.cloudAuthPrompt.pollFailures; got != maxPollFailures-2 {
t.Fatalf("pollFailures = %d, want %d", got, maxPollFailures-2)
}
// A healthy (no-error, not-done) response resets the streak so a later
// transient blip isn't counted against a recovered connection.
updated, _ := m.updateCloudAuthPrompt(cloudAuthPollMsg{done: false, err: nil})
m = updated.(chatModel)
if m.cloudAuthPrompt == nil {
t.Fatal("healthy response should keep the prompt open")
}
if got := m.cloudAuthPrompt.pollFailures; got != 0 {
t.Fatalf("pollFailures = %d, want 0 after healthy response", got)
}
if m.cloudAuthPrompt.pollErr != "" {
t.Fatalf("pollErr = %q, want empty after healthy response", m.cloudAuthPrompt.pollErr)
}
}
func TestCloudAuthPollCompletesAfterFailures(t *testing.T) {
pollErr := errors.New("whoami: timeout")
m := chatModel{
cloudAuthPrompt: &cloudAuthPrompt{
modelName: "glm-5.2:cloud",
polling: true,
kind: cloudAuthSignIn,
pollFailures: maxPollFailures - 1,
},
opts: Options{
CheckCloudModel: func(context.Context, string, string) error { return nil },
PollCloudAuth: func(context.Context) (string, bool, error) { return "", false, pollErr },
},
}
// A successful sign-in mid-retry should clear the failure state and re-check.
updated, _ := m.updateCloudAuthPrompt(cloudAuthPollMsg{done: true})
m = updated.(chatModel)
if m.cloudAuthPrompt.polling {
t.Fatal("done should stop polling")
}
if m.cloudAuthPrompt.pollFailures != 0 || m.cloudAuthPrompt.pollErr != "" {
t.Fatalf("failure state not reset: failures=%d err=%q", m.cloudAuthPrompt.pollFailures, m.cloudAuthPrompt.pollErr)
}
}
+101
View File
@@ -0,0 +1,101 @@
package chat
import (
"context"
"slices"
tea "github.com/charmbracelet/bubbletea"
coreagent "github.com/ollama/ollama/agent"
"github.com/ollama/ollama/api"
)
func (m *chatModel) startManualCompaction() (tea.Model, tea.Cmd) {
if m.running || m.compacting {
m.entries = append(m.entries, newChatEntry(chatEntry{role: "system", content: "Wait for the current response to finish before compacting."}))
return *m, nil
}
m.refreshContextWindowTokens(m.opts.Model)
if m.opts.Compactor == nil {
m.entries = append(m.entries, newChatEntry(chatEntry{role: "system", content: coreagent.CompactionSkippedMessage("compaction is unavailable")}))
m.status = "compact skipped"
return *m, nil
}
ctx := m.ctx
if ctx == nil {
ctx = context.Background()
}
runCtx, cancel := context.WithCancel(ctx)
compactor := m.opts.Compactor
events := make(chan tea.Msg, 128)
m.compacting = true
m.compactingTokens = 0
m.cancel = cancel
m.compactEvents = events
m.status = "compacting"
messages := slices.Clone(m.messages)
var tools api.Tools
if m.opts.Tools != nil {
tools = m.opts.Tools.Tools()
}
req := coreagent.CompactionRequest{
ChatID: m.chatID,
Model: m.opts.Model,
SystemPrompt: m.systemPrompt(""),
Messages: messages,
Tools: tools,
Format: m.opts.Format,
Options: m.opts.Options,
KeepAlive: m.opts.KeepAlive,
Force: true,
Progress: func(progress coreagent.CompactionProgress) {
select {
case events <- chatCompactProgressMsg{tokens: progress.Tokens}:
case <-runCtx.Done():
}
},
}
go func() {
defer close(events)
result, err := compactor.MaybeCompact(runCtx, req)
select {
case events <- chatCompactDoneMsg{result: result, err: err}:
case <-runCtx.Done():
}
}()
tickCmd := m.scheduleTick()
return *m, tea.Batch(waitForChatMsg(events), tickCmd)
}
func (m chatModel) finishManualCompaction(msg chatCompactDoneMsg) (tea.Model, tea.Cmd) {
wasCanceling := m.status == "canceling"
m.compacting = false
m.compactEvents = nil
m.cancel = nil
m.compactingTokens = 0
if wasCanceling || isChatContextCanceledError(msg.err) {
m.status = "compact canceled"
return m.withFlowTranscriptFlush(nil)
}
if msg.err != nil {
m.entries = append(m.entries, newChatEntry(chatEntry{role: "system", content: coreagent.CompactionSkippedMessage(msg.err.Error())}))
m.status = "compact skipped"
return m.withFlowTranscriptFlush(nil)
}
if !msg.result.Compacted {
m.entries = append(m.entries, newChatEntry(chatEntry{role: "system", content: coreagent.CompactionSkippedMessage(msg.result.Reason)}))
m.status = "compact skipped"
return m.withFlowTranscriptFlush(nil)
}
m.messages = msg.result.Messages
m.liveMessages = nil
m.entries = entriesFromMessages(m.messages)
m.contextTokens = m.estimatePromptTokens(m.messages, "")
m.contextEstimate = true
m.scroll = 0
m.flowPrintedLines = 0
m.status = "compacted"
return m.withFlowTranscriptFlush(nil)
}
+558
View File
@@ -0,0 +1,558 @@
package chat
import (
"encoding/json"
"fmt"
"os"
"path/filepath"
"slices"
"strings"
tea "github.com/charmbracelet/bubbletea"
"github.com/charmbracelet/lipgloss"
coreagent "github.com/ollama/ollama/agent"
"github.com/ollama/ollama/api"
)
type chatPromptDebug struct {
request api.ChatRequest
tokens int
scroll int
}
const maxPromptDebugToolResultRunes = 400
func (m *chatModel) handleSaveCommand(args string) (tea.Model, tea.Cmd) {
filename, err := saveRequestFilename(args)
if err != nil {
return m.addDebugError(err)
}
raw, err := m.rawRequestJSON()
if err != nil {
return m.addDebugError(err)
}
dir, err := m.debugWorkingDir()
if err != nil {
return m.addDebugError(err)
}
path := filepath.Join(dir, filename)
if err := os.WriteFile(path, []byte(raw+"\n"), 0o644); err != nil {
return m.addDebugError(err)
}
m.entries = append(m.entries, newSlashEntry(fmt.Sprintf("saved as %s", filename)))
m.status = "saved"
return *m, nil
}
func (m *chatModel) handlePromptCommand(args string) (tea.Model, tea.Cmd) {
if strings.TrimSpace(args) != "" {
return m.addDebugError(fmt.Errorf("usage: /prompt"))
}
req, tokens := m.requestPreview()
m.promptDebug = &chatPromptDebug{
request: req,
tokens: tokens,
}
m.flowPrintedLines = 0
m.selection = chatSelection{}
m.status = "prompt"
return *m, tea.Batch(tea.ClearScreen, tea.EnableMouseCellMotion)
}
func (m chatModel) updatePromptDebug(msg tea.KeyMsg) (tea.Model, tea.Cmd) {
if m.promptDebug == nil {
return m, nil
}
switch msg.Type {
case tea.KeyEsc, tea.KeyCtrlC, tea.KeyEnter:
return m.closePromptDebug()
case tea.KeyUp, tea.KeyCtrlP:
m.promptDebug.scroll--
case tea.KeyDown, tea.KeyCtrlN:
m.promptDebug.scroll++
case tea.KeyPgUp:
m.promptDebug.scroll -= max(1, m.promptDebugPageSize())
case tea.KeyPgDown:
m.promptDebug.scroll += max(1, m.promptDebugPageSize())
case tea.KeyHome, tea.KeyCtrlHome:
m.promptDebug.scroll = 0
case tea.KeyEnd, tea.KeyCtrlEnd:
m.promptDebug.scroll = m.promptDebugMaxScroll()
}
if m.promptDebug != nil {
m.promptDebug.scroll = clamp(m.promptDebug.scroll, 0, m.promptDebugMaxScroll())
}
return m, nil
}
func (m chatModel) closePromptDebug() (tea.Model, tea.Cmd) {
m.promptDebug = nil
m.status = "ready"
m.flowPrintedLines = 0
next, printCmd := m.flowTranscriptFlushCmd()
return next, tea.Sequence(tea.DisableMouse, tea.ClearScreen, printCmd)
}
func (m chatModel) renderPromptDebug(width, height int) string {
if width <= 0 {
width = 80
}
if height <= 0 {
height = 24
}
if m.promptDebug == nil {
return renderFullFrame("", width, height)
}
header := []string{
chatPickerTitleStyle.Render("Prompt"),
chatPickerMetaStyle.Render("full request preview • /save <filename> saved as <filename>.json"),
"",
}
footer := chatPickerMetaStyle.Render("↑/↓ scroll • pgup/pgdn page • enter/esc close")
bodyHeight := max(0, height-len(header)-1)
body := m.promptDebugLines(width)
maxScroll := max(0, len(body)-bodyHeight)
scroll := clamp(m.promptDebug.scroll, 0, maxScroll)
if bodyHeight < len(body) {
body = body[scroll:min(len(body), scroll+bodyHeight)]
}
lines := slices.Clone(header)
lines = append(lines, body...)
for len(lines) < height-1 {
lines = append(lines, "")
}
lines = append(lines, footer)
return renderFrameLines(lines, width, height)
}
func (m chatModel) promptDebugPageSize() int {
height := m.height
if height <= 0 {
height = 24
}
return max(1, height-5)
}
func (m chatModel) promptDebugMaxScroll() int {
if m.promptDebug == nil {
return 0
}
width := m.viewWidth()
height := m.height
if height <= 0 {
height = 24
}
bodyHeight := max(0, height-4)
return max(0, len(m.promptDebugLines(width))-bodyHeight)
}
func (m *chatModel) addDebugError(err error) (tea.Model, tea.Cmd) {
m.entries = append(m.entries, newChatEntry(chatEntry{role: "error", content: err.Error(), err: err.Error()}))
m.status = "error"
return *m, nil
}
func (m chatModel) rawRequestJSON() (string, error) {
req, _ := m.requestPreview()
data, err := json.MarshalIndent(req, "", " ")
if err != nil {
return "", err
}
return string(data), nil
}
func (m chatModel) requestPreview() (api.ChatRequest, int) {
opts := m.previewRunOptions()
messages := m.previewMessages()
req := m.previewChatRequest(opts, messages)
return req, m.estimatePromptTokens(messages, opts.SystemPrompt)
}
func (m chatModel) previewRunOptions() coreagent.RunOptions {
return coreagent.RunOptions{
ChatID: m.chatID,
Model: m.opts.Model,
SystemPrompt: m.systemPrompt(""),
Format: m.opts.Format,
Options: m.opts.Options,
Think: m.opts.Think,
KeepAlive: m.opts.KeepAlive,
}
}
func (m chatModel) previewMessages() []api.Message {
if len(m.liveMessages) > 0 {
return slices.Clone(m.liveMessages)
}
return slices.Clone(m.messages)
}
func (m chatModel) previewChatRequest(opts coreagent.RunOptions, messages []api.Message) api.ChatRequest {
requestMessages := slices.Clone(messages)
if strings.TrimSpace(opts.SystemPrompt) != "" {
withSystem := make([]api.Message, 0, len(requestMessages)+1)
withSystem = append(withSystem, api.Message{Role: "system", Content: opts.SystemPrompt})
requestMessages = append(withSystem, requestMessages...)
}
format := opts.Format
if format == "json" {
format = `"` + format + `"`
}
req := api.ChatRequest{
Model: opts.Model,
Messages: requestMessages,
Format: json.RawMessage(format),
Options: opts.Options,
Think: opts.Think,
}
if opts.KeepAlive != nil {
req.KeepAlive = opts.KeepAlive
}
if m.opts.Tools != nil && !m.opts.ToolsDisabled {
req.Tools = m.opts.Tools.Tools()
}
return req
}
func (m chatModel) promptDebugLines(width int) []string {
if m.promptDebug == nil {
return nil
}
req := m.promptDebug.request
innerWidth := max(20, width-2)
lines := []string{
chatHeaderStyle.Render("Request"),
promptDebugFieldLine("model", req.Model, innerWidth),
promptDebugFieldLine("estimated prompt", m.promptTokenText(m.promptDebug.tokens), innerWidth),
promptDebugFieldLine("messages", fmt.Sprint(len(req.Messages)), innerWidth),
promptDebugFieldLine("tools", fmt.Sprint(len(req.Tools)), innerWidth),
}
if len(req.Format) > 0 {
lines = append(lines, promptDebugFieldLine("format", strings.TrimSpace(string(req.Format)), innerWidth))
}
if req.Options != nil {
lines = append(lines, promptDebugMapLines("options", req.Options, innerWidth)...)
}
if req.Think != nil {
lines = append(lines, promptDebugBlockLines("think", req.Think.String(), innerWidth, chatHistoryTextStyle)...)
}
if req.KeepAlive != nil {
lines = append(lines, promptDebugFieldLine("keep_alive", req.KeepAlive.String(), innerWidth))
}
lines = append(lines, "", chatHeaderStyle.Render("Messages"))
if len(req.Messages) == 0 {
lines = append(lines, chatMetaStyle.Render("none"))
} else {
for i, msg := range req.Messages {
if i > 0 {
lines = append(lines, "")
}
lines = append(lines, promptDebugMessageLines(i+1, msg, innerWidth)...)
}
}
lines = append(lines, "", chatHeaderStyle.Render("Tools"))
if len(req.Tools) == 0 {
lines = append(lines, chatMetaStyle.Render("none"))
return lines
}
for i, tool := range req.Tools {
if i > 0 {
lines = append(lines, "")
}
lines = append(lines, promptDebugToolLines(i+1, tool, innerWidth)...)
}
return lines
}
func promptDebugFieldLine(label, value string, width int) string {
labelText := label + ":"
value = strings.TrimSpace(value)
if value == "" {
value = "_empty_"
}
line := chatHistoryLabelStyle.Render(labelText) + " " + chatHistoryTextStyle.Render(value)
return truncateRenderedLine(line, width)
}
func promptDebugMessageLines(index int, msg api.Message, width int) []string {
role := promptMessageLabel(msg)
header := fmt.Sprintf("%d. %s", index, role)
lines := []string{historyRoleStyle(msg.Role).Render(header)}
if strings.TrimSpace(msg.Thinking) != "" {
lines = append(lines, promptDebugBlockLines("thinking", msg.Thinking, width, chatHistoryTextStyle)...)
}
if msg.Role != "tool" && (strings.TrimSpace(msg.Content) != "" || (msg.Role != "assistant" && len(msg.ToolCalls) == 0 && len(msg.Images) == 0 && msg.Thinking == "")) {
lines = append(lines, promptDebugBlockLines("content", msg.Content, width, chatHistoryTextStyle)...)
}
if len(msg.ToolCalls) > 0 {
for i, call := range msg.ToolCalls {
lines = append(lines, promptDebugToolCallLines(i+1, call, width)...)
}
}
if msg.Role == "tool" {
if msg.ToolName != "" {
lines = append(lines, " "+chatHistoryLabelStyle.Render("tool_name:")+" "+chatHistoryTextStyle.Render(msg.ToolName))
}
if msg.ToolCallID != "" {
lines = append(lines, " "+chatHistoryLabelStyle.Render("tool_call_id:")+" "+chatHistoryTextStyle.Render(msg.ToolCallID))
}
lines = append(lines, promptDebugBlockLines("tool result", promptDebugToolResult(msg.Content), width, chatHistoryTextStyle)...)
}
if len(msg.Images) > 0 {
lines = append(lines, " "+chatHistoryLabelStyle.Render(fmt.Sprintf("%d image%s", len(msg.Images), pluralSuffix(len(msg.Images)))))
}
return lines
}
func promptDebugToolResult(content string) string {
runes := []rune(content)
if len(runes) <= maxPromptDebugToolResultRunes {
return content
}
return string(runes[:maxPromptDebugToolResultRunes-3]) + "..."
}
func promptDebugMapLines(label string, values map[string]any, width int) []string {
lines := []string{" " + chatHistoryLabelStyle.Render(label+":")}
if len(values) == 0 {
return append(lines, " "+chatMetaStyle.Render("_empty_"))
}
keys := make([]string, 0, len(values))
for key := range values {
keys = append(keys, key)
}
slices.Sort(keys)
for _, key := range keys {
lines = append(lines, promptDebugValueLine(4, key, values[key], width)...)
}
return lines
}
func promptDebugToolLines(index int, tool api.Tool, width int) []string {
name := strings.TrimSpace(tool.Function.Name)
if name == "" {
name = "_unnamed_"
}
lines := []string{historyRoleStyle("tool").Render(fmt.Sprintf("%d. %s", index, name))}
if strings.TrimSpace(tool.Function.Description) != "" {
lines = append(lines, promptDebugBlockLines("description", tool.Function.Description, width, chatHistoryTextStyle)...)
}
params := tool.Function.Parameters
if params.Type != "" || params.Properties != nil {
kind := params.Type
if kind == "" {
kind = "object"
}
lines = append(lines, " "+chatHistoryLabelStyle.Render("parameters:")+" "+chatHistoryTextStyle.Render(kind))
}
if params.Properties == nil || params.Properties.Len() == 0 {
return lines
}
lines = append(lines, " "+chatHistoryLabelStyle.Render("properties:"))
required := map[string]bool{}
for _, name := range params.Required {
required[name] = true
}
for name, property := range params.Properties.All() {
label := name
propertyType := property.ToTypeScriptType()
switch {
case propertyType != "" && required[name]:
label += " (" + propertyType + ", required)"
case propertyType != "":
label += " (" + propertyType + ")"
case required[name]:
label += " (required)"
}
value := strings.TrimSpace(property.Description)
if value == "" {
value = promptDebugPropertyDetails(property)
}
lines = append(lines, promptDebugTextLine(4, label, value, width)...)
}
return lines
}
func promptDebugToolCallLines(index int, call api.ToolCall, width int) []string {
name := strings.TrimSpace(call.Function.Name)
if name == "" {
name = "_unnamed_"
}
lines := []string{" " + chatHistoryLabelStyle.Render(fmt.Sprintf("tool call %d:", index)) + " " + chatHistoryTextStyle.Render(name)}
if strings.TrimSpace(call.ID) != "" {
lines = append(lines, promptDebugTextLine(4, "id", call.ID, width)...)
}
if call.Function.Arguments.Len() == 0 {
lines = append(lines, " "+chatHistoryLabelStyle.Render("arguments:")+" "+chatMetaStyle.Render("none"))
return lines
}
lines = append(lines, " "+chatHistoryLabelStyle.Render("arguments:"))
for key, value := range call.Function.Arguments.All() {
lines = append(lines, promptDebugValueLine(6, key, value, width)...)
}
return lines
}
func promptDebugPropertyDetails(property api.ToolProperty) string {
var parts []string
if len(property.Enum) > 0 {
values := make([]string, 0, len(property.Enum))
for _, value := range property.Enum {
values = append(values, promptDebugValueText(value))
}
parts = append(parts, "one of "+strings.Join(values, ", "))
}
if property.Properties != nil && property.Properties.Len() > 0 {
count := property.Properties.Len()
noun := "property"
if count != 1 {
noun = "properties"
}
parts = append(parts, fmt.Sprintf("%d nested %s", count, noun))
}
if property.Items != nil {
parts = append(parts, "array items: "+promptDebugValueText(property.Items))
}
if len(parts) == 0 {
return "_empty_"
}
return strings.Join(parts, "; ")
}
func promptDebugValueLine(indent int, label string, value any, width int) []string {
return promptDebugTextLine(indent, label, promptDebugValueText(value), width)
}
func promptDebugTextLine(indent int, label, value string, width int) []string {
prefix := strings.Repeat(" ", indent) + chatHistoryLabelStyle.Render(label+":")
value = strings.TrimSpace(value)
if value == "" {
value = "_empty_"
}
wrapWidth := max(20, width-indent-lipgloss.Width(label)-2)
wrapped := wrapChatText(value, wrapWidth)
if len(wrapped) == 0 {
return []string{prefix + " " + chatMetaStyle.Render("_empty_")}
}
lines := []string{prefix + " " + chatHistoryTextStyle.Render(wrapped[0])}
for _, line := range wrapped[1:] {
lines = append(lines, strings.Repeat(" ", indent+2)+chatHistoryTextStyle.Render(line))
}
return lines
}
func promptDebugValueText(value any) string {
switch v := value.(type) {
case nil:
return "null"
case string:
return v
case fmt.Stringer:
return v.String()
case []any:
parts := make([]string, 0, len(v))
for _, item := range v {
parts = append(parts, promptDebugValueText(item))
}
return strings.Join(parts, ", ")
case map[string]any:
keys := make([]string, 0, len(v))
for key := range v {
keys = append(keys, key)
}
slices.Sort(keys)
parts := make([]string, 0, len(keys))
for _, key := range keys {
parts = append(parts, key+": "+promptDebugValueText(v[key]))
}
return strings.Join(parts, ", ")
default:
return fmt.Sprint(value)
}
}
func promptDebugBlockLines(label, value string, width int, style lipgloss.Style) []string {
lines := []string{" " + chatHistoryLabelStyle.Render(label+":")}
if value == "" {
return append(lines, " "+chatMetaStyle.Render("_empty_"))
}
for _, raw := range strings.Split(strings.TrimRight(value, "\n"), "\n") {
if raw == "" {
lines = append(lines, "")
continue
}
for _, wrapped := range wrapChatText(raw, max(20, width-4)) {
lines = append(lines, " "+style.Render(wrapped))
}
}
return lines
}
func (m chatModel) promptTokenText(tokens int) string {
window := m.displayContextWindowTokens()
if window > 0 {
return fmt.Sprintf("%s / %s tokens", formatPromptTokenCount(max(tokens, 0)), formatPromptTokenCount(window))
}
return formatTokenCount(tokens)
}
func formatPromptTokenCount(count int) string {
sign := ""
if count < 0 {
sign = "-"
count = -count
}
if count < 100_000 {
return sign + fmt.Sprint(count)
}
if count >= 950_000 {
return fmt.Sprintf("%s%dM", sign, int(float64(count)/1_000_000+0.5))
}
return fmt.Sprintf("%s%dk", sign, int(float64(count)/1024+0.5))
}
func promptMessageLabel(msg api.Message) string {
if msg.Role == "tool" && msg.ToolName != "" {
return msg.Role + ":" + msg.ToolName
}
return msg.Role
}
func saveRequestFilename(args string) (string, error) {
args = strings.TrimSpace(args)
if args == "" {
return "", fmt.Errorf("usage: /save <filename>")
}
if strings.HasPrefix(args, ">") {
args = strings.TrimSpace(strings.TrimPrefix(args, ">"))
}
fields := strings.Fields(args)
if len(fields) != 1 {
return "", fmt.Errorf("usage: /save <filename>")
}
filename := strings.TrimSpace(fields[0])
if filename == "" || filename == "." || filename == ".." || strings.ContainsAny(filename, `/\`) || filepath.IsAbs(filename) {
return "", fmt.Errorf("save filename must be a file name, not a path")
}
if !strings.HasSuffix(strings.ToLower(filename), ".json") {
filename += ".json"
}
return filename, nil
}
func (m chatModel) debugWorkingDir() (string, error) {
dir := strings.TrimSpace(m.currentWorkingDir())
if dir != "" {
return dir, nil
}
return os.Getwd()
}
+373
View File
@@ -0,0 +1,373 @@
package chat
import (
"context"
"slices"
"strings"
"time"
tea "github.com/charmbracelet/bubbletea"
coreagent "github.com/ollama/ollama/agent"
"github.com/ollama/ollama/api"
)
type chatAgentMsg struct {
event coreagent.Event
}
type chatApprovalPromptMsg struct {
request coreagent.ApprovalRequest
reply chan<- coreagent.Approval
}
type chatRunDoneMsg struct {
result *coreagent.RunResult
err error
newMessagesPersisted bool
persistedMessages []api.Message
}
type chatCompactDoneMsg struct {
result coreagent.CompactionResult
err error
}
type chatCompactProgressMsg struct {
tokens int
}
// resetStreamingState clears the transient streaming flags that every
// non-streaming event resets before applying its own state.
func (m *chatModel) resetStreamingState() {
m.finishThinkingEntry()
m.awaitingModel = false
m.thinking = false
m.thinkingTokens = 0
}
// resetRunState clears all run-progress flags (streaming plus compaction
// progress) for terminal events that fully reset the run view.
func (m *chatModel) resetRunState() {
m.finishThinkingEntry()
m.awaitingModel = false
m.compacting = false
m.compactingTokens = 0
m.detectedToolCalls = nil
m.thinking = false
m.thinkingTokens = 0
}
type chatModelPreloadDoneMsg struct {
model string
contextWindowTokens int
err error
}
type chatEventsClosedMsg struct{}
type chatTickMsg struct{}
func (m *chatModel) applyAgentEvent(event coreagent.Event) {
contextChanged := false
switch event.Type {
case coreagent.EventThinkingDelta:
m.awaitingModel = false
if event.Thinking != "" {
m.thinking = true
if event.Tokens > 0 {
m.thinkingTokens = max(m.thinkingTokens, event.Tokens)
} else {
m.thinkingTokens += approximateTokenCount(event.Thinking)
}
idx := m.ensureLiveAssistantMessage()
m.liveMessages[idx].Thinking += event.Thinking
m.syncThinkingEntry()
contextChanged = true
}
case coreagent.EventMessageDelta:
m.resetStreamingState()
m.spinner = 0
m.detectedToolCalls = nil
idx := m.ensureAssistantEntry()
m.entries[idx].content += event.Content
m.markEntryDirty(idx)
msgIdx := m.ensureLiveAssistantMessage()
m.liveMessages[msgIdx].Content += event.Content
contextChanged = true
case coreagent.EventToolCallDetected:
m.finishThinkingEntry()
m.awaitingModel = m.running
m.thinking = false
m.thinkingTokens = 0
m.groupCompletedToolHistory()
m.detectedToolCalls = nil
m.addDetectedToolCalls(event.ToolCalls)
idx := m.ensureLiveAssistantMessage()
m.liveMessages[idx].ToolCalls = append(m.liveMessages[idx].ToolCalls, event.ToolCalls...)
contextChanged = true
case coreagent.EventToolStarted:
m.resetStreamingState()
m.refreshContextWindowTokens(m.opts.Model)
startedAt := time.Now()
idx := m.findActiveToolEntry(event.ToolCallID)
if idx < 0 {
m.groupCompletedToolHistory()
m.entries = append(m.entries, newChatEntry(chatEntry{role: "tool"}))
idx = len(m.entries) - 1
}
m.entries[idx].detail = event.ToolName
m.entries[idx].label = toolInvocationLabel(event.ToolName, event.Args)
m.entries[idx].status = "running"
m.entries[idx].toolID = event.ToolCallID
m.entries[idx].args = event.Args
m.entries[idx].startedAt = startedAt
m.applyToolOutputModeTo(idx)
m.markEntryDirty(idx)
case coreagent.EventToolFinished:
m.resetStreamingState()
m.refreshContextWindowTokens(m.opts.Model)
if event.WorkingDir != "" {
m.workingDir = event.WorkingDir
}
startedAt := m.toolStartedAt(event.ToolCallID)
status := toolFinishedStatus(event)
idx := m.findToolEntry(event.ToolCallID)
if idx < 0 {
m.entries = append(m.entries, newChatEntry(chatEntry{role: "tool"}))
idx = len(m.entries) - 1
}
m.entries[idx].content = event.Content
m.entries[idx].label = toolInvocationLabel(event.ToolName, event.Args)
m.entries[idx].detail = event.ToolName
m.entries[idx].status = status
if status != "denied" {
m.entries[idx].err = event.Error
}
m.entries[idx].toolID = event.ToolCallID
m.entries[idx].args = event.Args
m.entries[idx].startedAt = startedAt
m.entries[idx].finishedAt = time.Now()
m.applyToolOutputModeTo(idx)
m.markEntryDirty(idx)
m.liveMessages = append(m.liveMessages, api.Message{
Role: "tool",
Content: event.Content,
ToolName: event.ToolName,
ToolCallID: event.ToolCallID,
})
if m.running && status != "denied" && !m.hasPendingDetectedToolCalls() {
m.awaitingModel = true
}
contextChanged = true
case coreagent.EventCompacted:
m.resetRunState()
if len(event.Messages) > 0 {
m.liveMessages = slices.Clone(event.Messages)
m.messages = slices.Clone(event.Messages)
contextChanged = true
}
m.status = "compacted"
case coreagent.EventCompactionStarted:
m.awaitingModel = false
m.compacting = true
m.compactingTokens = 0
m.thinking = false
m.thinkingTokens = 0
m.status = "compacting"
case coreagent.EventCompactionProgress:
m.awaitingModel = false
m.compacting = true
m.thinking = false
m.thinkingTokens = 0
if event.Tokens > m.compactingTokens {
m.compactingTokens = event.Tokens
}
case coreagent.EventCompactionSkipped:
m.resetRunState()
message := event.Content
if strings.TrimSpace(message) == "" {
message = coreagent.CompactionSkippedMessage(event.Error)
}
m.entries = append(m.entries, newChatEntry(chatEntry{role: "system", content: message}))
m.status = "compact skipped"
case coreagent.EventError:
m.resetRunState()
m.eventErrorRendered = true
m.entries = append(m.entries, newChatEntry(chatEntry{role: "error", content: event.Error, err: event.Error}))
}
if contextChanged {
m.refreshLiveContextEstimate()
}
}
func (m *chatModel) addDetectedToolCalls(calls []api.ToolCall) {
if len(calls) == 0 {
return
}
seen := make(map[string]struct{}, len(m.detectedToolCalls)+len(calls))
for _, entry := range m.detectedToolCalls {
if entry.toolID != "" {
seen[entry.toolID] = struct{}{}
}
}
for _, call := range calls {
if call.ID != "" {
if _, ok := seen[call.ID]; ok {
continue
}
seen[call.ID] = struct{}{}
}
args := call.Function.Arguments.ToMap()
m.detectedToolCalls = append(m.detectedToolCalls, newChatEntry(chatEntry{
role: "tool",
label: toolInvocationLabel(call.Function.Name, args),
detail: call.Function.Name,
status: "queued",
toolID: call.ID,
args: args,
}))
}
}
func toolFinishedStatus(event coreagent.Event) string {
switch event.ToolStatus {
case coreagent.ToolStatusDenied:
return "denied"
case coreagent.ToolStatusDisabled:
return "disabled"
case coreagent.ToolStatusDone:
return "done"
}
// failed/skipped/unknown: derive from content and error fields.
if isDeniedToolResult(event.Content) || isDeniedToolResult(event.Error) {
return "denied"
}
if event.Error != "" {
return "error"
}
return "done"
}
func messagesEndWithCompactionResult(messages []api.Message) bool {
if len(messages) == 0 {
return false
}
return coreagent.IsCompactionToolResult(messages[len(messages)-1])
}
func (m chatModel) awaitingToolStart() bool {
for i := len(m.liveMessages) - 1; i >= 0; i-- {
msg := m.liveMessages[i]
if msg.Role != "assistant" {
continue
}
if len(msg.ToolCalls) == 0 {
return false
}
for _, call := range msg.ToolCalls {
if call.ID == "" || m.findToolEntry(call.ID) < 0 {
return true
}
}
return false
}
return false
}
func (m *chatModel) ensureLiveAssistantMessage() int {
if len(m.liveMessages) > 0 && m.liveMessages[len(m.liveMessages)-1].Role == "assistant" {
return len(m.liveMessages) - 1
}
m.liveMessages = append(m.liveMessages, api.Message{Role: "assistant"})
return len(m.liveMessages) - 1
}
func (m *chatModel) refreshLiveContextEstimate() {
messages := m.liveMessages
if len(messages) == 0 {
messages = m.messages
}
m.contextTokens = m.estimatePromptTokens(messages, "")
m.contextEstimate = true
}
//nolint:containedctx // event sinks need the session context to unblock sends on cancellation.
type chatEventSink struct {
ctx context.Context
ch chan<- tea.Msg
newMessagesPersisted *bool
}
func (s chatEventSink) Emit(event coreagent.Event) error {
if s.newMessagesPersisted != nil {
*s.newMessagesPersisted = true
}
select {
case s.ch <- chatAgentMsg{event: event}:
return nil
case <-s.ctx.Done():
return s.ctx.Err()
}
}
func waitForChatMsg(ch <-chan tea.Msg) tea.Cmd {
if ch == nil {
return nil
}
return func() tea.Msg {
msg, ok := <-ch
if !ok {
return chatEventsClosedMsg{}
}
return msg
}
}
func (m *chatModel) scheduleTick() tea.Cmd {
if m.tickActive {
return nil
}
m.tickActive = true
return chatTickCmd()
}
func chatTickCmd() tea.Cmd {
return tea.Tick(350*time.Millisecond, func(time.Time) tea.Msg {
return chatTickMsg{}
})
}
func preloadModelCmd(ctx context.Context, preload func(context.Context, string, *api.ThinkValue) (int, error), model string, think *api.ThinkValue) tea.Cmd {
if preload == nil || strings.TrimSpace(model) == "" {
return nil
}
if think != nil {
copied := *think
think = &copied
}
return func() tea.Msg {
if ctx == nil {
ctx = context.Background()
}
tokens, err := preload(ctx, model, think)
return chatModelPreloadDoneMsg{model: model, contextWindowTokens: tokens, err: err}
}
}
func isUnsupportedThinkingError(err error) bool {
if err == nil {
return false
}
text := strings.ToLower(err.Error())
return strings.Contains(text, "does not support thinking")
}
func thinkRequestsThinking(think *api.ThinkValue) bool {
if think == nil {
return false
}
return think.Bool()
}
+376
View File
@@ -0,0 +1,376 @@
package chat
import (
"strings"
"testing"
coreagent "github.com/ollama/ollama/agent"
"github.com/ollama/ollama/api"
)
func TestApplyAgentEventStreamsAssistantContent(t *testing.T) {
m := chatModel{running: true}
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventMessageDelta, Content: "hello"})
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventMessageDelta, Content: " world"})
if len(m.entries) != 1 || m.entries[0].role != "assistant" || m.entries[0].content != "hello world" {
t.Fatalf("entries = %#v", m.entries)
}
if len(m.liveMessages) != 1 || m.liveMessages[0].Content != "hello world" {
t.Fatalf("live messages = %#v", m.liveMessages)
}
}
func TestApplyAgentEventTracksToolLifecycle(t *testing.T) {
m := chatModel{running: true}
args := map[string]any{"command": "pwd"}
m.applyAgentEvent(coreagent.Event{
Type: coreagent.EventToolStarted,
ToolCallID: "call-1",
ToolName: "bash",
Args: args,
})
m.applyAgentEvent(coreagent.Event{
Type: coreagent.EventToolFinished,
ToolCallID: "call-1",
ToolName: "bash",
Args: args,
Content: "ok",
})
if len(m.entries) != 1 {
t.Fatalf("entries = %#v", m.entries)
}
entry := m.entries[0]
if entry.status != "done" || entry.content != "ok" || !strings.Contains(entry.label, "Bash") {
t.Fatalf("tool entry = %#v", entry)
}
if line := stripANSI(toolStatusLine(entry)); line != `Bash("pwd")` {
t.Fatalf("tool status line = %q, want command label", line)
}
if len(m.liveMessages) != 1 || m.liveMessages[0].Role != "tool" || m.liveMessages[0].Content != "ok" {
t.Fatalf("live messages = %#v", m.liveMessages)
}
}
func TestApplyAgentEventRendersDeniedCommandAsDenied(t *testing.T) {
m := chatModel{running: true}
args := map[string]any{"command": "pwd"}
m.applyAgentEvent(coreagent.Event{
Type: coreagent.EventToolFinished,
ToolStatus: coreagent.ToolStatusDenied,
ToolCallID: "call-1",
ToolName: "bash",
Args: args,
Content: "Tool execution denied.",
Error: "Tool execution denied.",
})
if len(m.entries) != 1 {
t.Fatalf("entries = %#v", m.entries)
}
entry := m.entries[0]
if entry.status != "denied" {
t.Fatalf("tool status = %q, want denied: %#v", entry.status, entry)
}
if line := stripANSI(toolStatusLine(entry)); line != `Bash("pwd") denied` {
t.Fatalf("tool status line = %q, want denied command label", line)
}
}
func TestApplyAgentEventShowsWorkingWhileAwaitingCloudToolStart(t *testing.T) {
args := api.NewToolCallFunctionArguments()
args.Set("command", "pwd")
m := chatModel{
running: true,
}
m.applyAgentEvent(coreagent.Event{
Type: coreagent.EventToolCallDetected,
ToolCalls: []api.ToolCall{{
ID: "call-1",
Function: api.ToolCallFunction{
Name: "bash",
Arguments: args,
},
}},
})
if line := stripANSI(m.activityLine()); !strings.Contains(line, "Working") {
t.Fatalf("activityLine = %q, want Working while tool call is pending", line)
}
}
func TestActivityLineShowsWorkingWhileAwaitingModelBeforeFirstEvent(t *testing.T) {
m := chatModel{
running: true,
awaitingModel: true,
spinner: 0,
}
if line := stripANSI(m.activityLine()); !strings.Contains(line, "Working") {
t.Fatalf("activityLine = %q, want Working while stream is open before first event", line)
}
}
func TestActivityLineShowsWorkingAfterAssistantContentGoesIdle(t *testing.T) {
m := chatModel{
running: true,
spinner: idleWorkingDelayTicks,
}
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventMessageDelta, Content: "I will inspect that next."})
if line := strings.TrimSpace(stripANSI(m.activityLine())); line != "" {
t.Fatalf("activityLine immediately after content = %q, want quiet until the idle delay", line)
}
m.spinner = idleWorkingDelayTicks
if line := stripANSI(m.activityLine()); !strings.Contains(line, "Working") {
t.Fatalf("activityLine after idle content stream = %q, want Working while stream remains open", line)
}
}
func TestApplyAgentEventKeepsDetectedBatchStableUntilComplete(t *testing.T) {
firstArgs := api.NewToolCallFunctionArguments()
firstArgs.Set("command", "pwd")
secondArgs := api.NewToolCallFunctionArguments()
secondArgs.Set("command", "ls")
m := chatModel{
running: true,
spinner: idleWorkingDelayTicks,
}
m.applyAgentEvent(coreagent.Event{
Type: coreagent.EventToolCallDetected,
ToolCalls: []api.ToolCall{
{ID: "call-1", Function: api.ToolCallFunction{Name: "bash", Arguments: firstArgs}},
{ID: "call-2", Function: api.ToolCallFunction{Name: "bash", Arguments: secondArgs}},
},
})
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventToolStarted, ToolCallID: "call-1", ToolName: "bash", Args: firstArgs.ToMap()})
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventToolFinished, ToolCallID: "call-1", ToolName: "bash", Args: firstArgs.ToMap(), Content: "one"})
if len(m.entries) != 1 {
t.Fatalf("entries = %d, want first completed command row: %#v", len(m.entries), m.entries)
}
if m.entries[0].role != "tool" || m.entries[0].status != "done" {
t.Fatalf("first command should remain stable while second is pending: %#v", m.entries[0])
}
if line := stripANSI(toolStatusLine(m.entries[0])); line != `Bash("pwd")` {
t.Fatalf("completed command line = %q", line)
}
if line := stripANSI(m.activityLine()); !strings.Contains(line, "Working") {
t.Fatalf("activityLine = %q, want Working while second command is pending", line)
}
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventToolStarted, ToolCallID: "call-2", ToolName: "bash", Args: secondArgs.ToMap()})
if len(m.entries) != 2 {
t.Fatalf("entries after second start = %d, want finished command plus running command: %#v", len(m.entries), m.entries)
}
if line := stripANSI(toolStatusLine(m.entries[0])); line != `Bash("pwd")` {
t.Fatalf("finished command line after second start = %q", line)
}
if line := stripANSI(toolStatusLine(m.entries[1])); line != `Bash("ls")` {
t.Fatalf("running command line = %q", line)
}
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventToolFinished, ToolCallID: "call-2", ToolName: "bash", Args: secondArgs.ToMap(), Content: "two"})
if line := stripANSI(m.activityLine()); !strings.Contains(line, "Working") {
t.Fatalf("activityLine after completed batch = %q, want Working while waiting for next model response", line)
}
if len(m.entries) != 2 {
t.Fatalf("entries after batch completion = %d, want stable command rows until the next tool boundary: %#v", len(m.entries), m.entries)
}
for i, want := range []string{`Bash("pwd")`, `Bash("ls")`} {
if line := stripANSI(toolStatusLine(m.entries[i])); line != want {
t.Fatalf("completed command row %d = %q, want %q", i, line, want)
}
}
thirdArgs := api.NewToolCallFunctionArguments()
thirdArgs.Set("command", "date")
m.applyAgentEvent(coreagent.Event{
Type: coreagent.EventToolCallDetected,
ToolCalls: []api.ToolCall{
{ID: "call-3", Function: api.ToolCallFunction{Name: "bash", Arguments: thirdArgs}},
},
})
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventToolStarted, ToolCallID: "call-3", ToolName: "bash", Args: thirdArgs.ToMap()})
if len(m.entries) != 2 {
t.Fatalf("entries after next tool boundary = %d, want grouped history plus running command: %#v", len(m.entries), m.entries)
}
if m.entries[0].role != "tool_group" || len(m.entries[0].tools) != 2 {
t.Fatalf("completed detected batch should collapse at the next tool boundary: %#v", m.entries[0])
}
if line := stripANSI(toolGroupStatusLine(m.entries[0])); line != "Ran 2 commands" {
t.Fatalf("grouped command line = %q", line)
}
if line := stripANSI(toolStatusLine(m.entries[1])); line != `Bash("date")` {
t.Fatalf("running command line = %q", line)
}
}
func TestApplyAgentEventDoesNotCollapsePartialDetectedBatch(t *testing.T) {
firstArgs := api.NewToolCallFunctionArguments()
firstArgs.Set("command", "pwd")
secondArgs := api.NewToolCallFunctionArguments()
secondArgs.Set("command", "ls")
thirdArgs := api.NewToolCallFunctionArguments()
thirdArgs.Set("command", "date")
m := chatModel{running: true}
m.applyAgentEvent(coreagent.Event{
Type: coreagent.EventToolCallDetected,
ToolCalls: []api.ToolCall{
{ID: "call-1", Function: api.ToolCallFunction{Name: "bash", Arguments: firstArgs}},
{ID: "call-2", Function: api.ToolCallFunction{Name: "bash", Arguments: secondArgs}},
{ID: "call-3", Function: api.ToolCallFunction{Name: "bash", Arguments: thirdArgs}},
},
})
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventToolStarted, ToolCallID: "call-1", ToolName: "bash", Args: firstArgs.ToMap()})
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventToolFinished, ToolCallID: "call-1", ToolName: "bash", Args: firstArgs.ToMap(), Content: "one"})
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventToolStarted, ToolCallID: "call-2", ToolName: "bash", Args: secondArgs.ToMap()})
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventToolFinished, ToolCallID: "call-2", ToolName: "bash", Args: secondArgs.ToMap(), Content: "two"})
if line := stripANSI(m.activityLine()); !strings.Contains(line, "Working") {
t.Fatalf("activityLine before final detected call = %q, want Working while final tool is pending", line)
}
if len(m.entries) != 2 {
t.Fatalf("entries before final detected call = %d, want two stable rows: %#v", len(m.entries), m.entries)
}
for i, want := range []string{`Bash("pwd")`, `Bash("ls")`} {
if line := stripANSI(toolStatusLine(m.entries[i])); line != want {
t.Fatalf("tool row %d = %q, want %q", i, line, want)
}
}
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventToolStarted, ToolCallID: "call-3", ToolName: "bash", Args: thirdArgs.ToMap()})
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventToolFinished, ToolCallID: "call-3", ToolName: "bash", Args: thirdArgs.ToMap(), Content: "three"})
if len(m.entries) != 3 {
t.Fatalf("entries after full detected batch = %#v, want stable tool rows until the next tool boundary", m.entries)
}
for i, want := range []string{`Bash("pwd")`, `Bash("ls")`, `Bash("date")`} {
if line := stripANSI(toolStatusLine(m.entries[i])); line != want {
t.Fatalf("tool row %d = %q, want %q", i, line, want)
}
}
fourthArgs := api.NewToolCallFunctionArguments()
fourthArgs.Set("command", "whoami")
m.applyAgentEvent(coreagent.Event{
Type: coreagent.EventToolCallDetected,
ToolCalls: []api.ToolCall{
{ID: "call-4", Function: api.ToolCallFunction{Name: "bash", Arguments: fourthArgs}},
},
})
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventToolStarted, ToolCallID: "call-4", ToolName: "bash", Args: fourthArgs.ToMap()})
if len(m.entries) != 2 || m.entries[0].role != "tool_group" || len(m.entries[0].tools) != 3 {
t.Fatalf("entries after next detected batch starts = %#v, want one grouped history entry plus active tool", m.entries)
}
if line := stripANSI(toolGroupStatusLine(m.entries[0])); line != "Ran 3 commands" {
t.Fatalf("grouped command line = %q", line)
}
}
func TestApplyAgentEventGroupsCompletedCommandsAtNextToolBoundary(t *testing.T) {
m := chatModel{running: true}
firstArgs := map[string]any{"command": "pwd"}
secondArgs := map[string]any{"command": "ls"}
thirdArgs := map[string]any{"command": "date"}
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventToolStarted, ToolCallID: "call-1", ToolName: "bash", Args: firstArgs})
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventToolFinished, ToolCallID: "call-1", ToolName: "bash", Args: firstArgs, Content: "one"})
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventToolStarted, ToolCallID: "call-2", ToolName: "bash", Args: secondArgs})
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventToolFinished, ToolCallID: "call-2", ToolName: "bash", Args: secondArgs, Content: "two"})
if len(m.entries) != 2 {
t.Fatalf("entries after second finish = %d, want two stable command rows: %#v", len(m.entries), m.entries)
}
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventToolStarted, ToolCallID: "call-3", ToolName: "bash", Args: thirdArgs})
if len(m.entries) != 2 {
t.Fatalf("entries = %d, want grouped command history plus active command: %#v", len(m.entries), m.entries)
}
if m.entries[0].role != "tool_group" || len(m.entries[0].tools) != 2 {
t.Fatalf("completed commands should be grouped when the next command starts: %#v", m.entries[0])
}
if line := stripANSI(toolGroupStatusLine(m.entries[0])); line != "Ran 2 commands" {
t.Fatalf("grouped command line = %q", line)
}
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventToolFinished, ToolCallID: "call-3", ToolName: "bash", Args: thirdArgs, Content: "three"})
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventMessageDelta, Content: "done"})
if len(m.entries) != 3 {
t.Fatalf("entries after assistant content = %d, want grouped history, last command, assistant: %#v", len(m.entries), m.entries)
}
transcript := stripANSI(m.renderTranscript(100))
if !strings.Contains(transcript, "• Ran 2 commands\n\n• Bash(\"date\")\n\n done") {
t.Fatalf("tool history should stay visually separated from assistant content:\n%s", transcript)
}
}
func TestApplyAgentEventDoesNotGroupCompletedCommandsOnMessageDelta(t *testing.T) {
m := chatModel{running: true}
firstArgs := map[string]any{"command": "pwd"}
secondArgs := map[string]any{"command": "ls"}
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventToolStarted, ToolCallID: "call-1", ToolName: "bash", Args: firstArgs})
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventToolFinished, ToolCallID: "call-1", ToolName: "bash", Args: firstArgs, Content: "one"})
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventToolStarted, ToolCallID: "call-2", ToolName: "bash", Args: secondArgs})
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventToolFinished, ToolCallID: "call-2", ToolName: "bash", Args: secondArgs, Content: "two"})
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventMessageDelta, Content: "done"})
if len(m.entries) != 3 {
t.Fatalf("entries = %d, want two command rows plus assistant content: %#v", len(m.entries), m.entries)
}
if m.entries[0].role != "tool" || m.entries[1].role != "tool" || m.entries[2].role != "assistant" {
t.Fatalf("completed commands should not collapse on assistant content: %#v", m.entries)
}
}
func TestApplyAgentEventGroupsPreviouslyDeniedCommandsAtNextToolBoundary(t *testing.T) {
m := chatModel{running: true}
firstArgs := map[string]any{"command": "pwd"}
secondArgs := map[string]any{"command": "ls"}
thirdArgs := map[string]any{"command": "date"}
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventToolFinished, ToolStatus: coreagent.ToolStatusDenied, ToolCallID: "call-1", ToolName: "bash", Args: firstArgs, Content: "Tool execution denied.", Error: "Tool execution denied."})
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventToolFinished, ToolStatus: coreagent.ToolStatusDenied, ToolCallID: "call-2", ToolName: "bash", Args: secondArgs, Content: "Tool execution denied.", Error: "Tool execution denied."})
if len(m.entries) != 2 {
t.Fatalf("entries = %d, want two stable denied command rows: %#v", len(m.entries), m.entries)
}
m.applyAgentEvent(coreagent.Event{Type: coreagent.EventToolStarted, ToolCallID: "call-3", ToolName: "bash", Args: thirdArgs})
if len(m.entries) != 2 {
t.Fatalf("entries = %d, want grouped denied command entry plus active command: %#v", len(m.entries), m.entries)
}
if m.entries[0].role != "tool_group" || len(m.entries[0].tools) != 2 {
t.Fatalf("denied commands should be grouped at the next tool boundary: %#v", m.entries[0])
}
if line := stripANSI(toolGroupStatusLine(m.entries[0])); line != "Denied 2 commands" {
t.Fatalf("grouped command line = %q", line)
}
if line := stripANSI(toolStatusLine(m.entries[1])); line != `Bash("date")` {
t.Fatalf("running command line = %q", line)
}
}
func TestMessagesEndWithCompactionResult(t *testing.T) {
messages := []api.Message{{
Role: "tool",
ToolName: coreagent.CompactionToolName,
ToolCallID: coreagent.CompactionToolCallID,
Content: coreagent.CompactionSummaryMessagePrefix + "summary",
}}
if !messagesEndWithCompactionResult(messages) {
t.Fatal("expected compaction result")
}
}
File diff suppressed because it is too large. Load diff
File diff suppressed because it is too large. Load diff
+294
View File
@@ -0,0 +1,294 @@
package chat
import (
"strings"
"github.com/charmbracelet/lipgloss"
)
func renderMarkdownForView(markdown string, width int) string {
if width < 20 {
width = 20
}
source := strings.Split(strings.TrimRight(markdown, "\n"), "\n")
var rendered []string
inCodeBlock := false
for i := 0; i < len(source); i++ {
line := strings.TrimRight(source[i], "\r")
trimmed := strings.TrimSpace(line)
if strings.HasPrefix(trimmed, "```") {
inCodeBlock = !inCodeBlock
continue
}
if inCodeBlock {
rendered = append(rendered, renderMarkdownCodeLine(line, width)...)
continue
}
if table, consumed := renderMarkdownTable(source[i:], width); consumed > 0 {
rendered = append(rendered, table...)
i += consumed - 1
continue
}
if heading, ok := markdownHeading(trimmed); ok {
rendered = append(rendered, chatHeaderStyle.Render(heading))
continue
}
if trimmed == "" {
rendered = append(rendered, "")
continue
}
for _, wrapped := range wrapChatText(line, width) {
rendered = append(rendered, renderMarkdownInline(wrapped))
}
}
return strings.Join(rendered, "\n")
}
func splitRenderedBody(body string) []string {
body = strings.TrimRight(body, "\n")
if body == "" {
return []string{""}
}
return strings.Split(body, "\n")
}
func markdownHeading(line string) (string, bool) {
if !strings.HasPrefix(line, "#") {
return "", false
}
level := 0
for level < len(line) && line[level] == '#' {
level++
}
if level == 0 || level > 6 || level >= len(line) || line[level] != ' ' {
return "", false
}
return strings.TrimSpace(line[level:]), true
}
func renderMarkdownInline(line string) string {
var b strings.Builder
for {
before, rest, ok := strings.Cut(line, "`")
b.WriteString(before)
if !ok {
break
}
code, after, ok := strings.Cut(rest, "`")
if !ok {
b.WriteString("`")
b.WriteString(rest)
break
}
b.WriteString(chatInlineCodeStyle.Render(code))
line = after
}
return b.String()
}
func renderMarkdownCodeLine(line string, width int) []string {
codeWidth := max(1, width-2)
lines := wrapChatText(line, codeWidth)
for i, wrapped := range lines {
lines[i] = " " + chatCodeBlockStyle.Render(wrapped)
}
return lines
}
func renderMarkdownTable(lines []string, width int) ([]string, int) {
if len(lines) < 2 || !looksLikeMarkdownTableRow(lines[0]) || !isMarkdownTableSeparator(lines[1]) {
return nil, 0
}
var rows [][]string
consumed := 0
for consumed < len(lines) && looksLikeMarkdownTableRow(lines[consumed]) {
if consumed == 1 && isMarkdownTableSeparator(lines[consumed]) {
consumed++
continue
}
rows = append(rows, parseMarkdownTableRow(lines[consumed]))
consumed++
}
if len(rows) == 0 {
return nil, 0
}
columnCount := 0
for _, row := range rows {
columnCount = max(columnCount, len(row))
}
naturalWidths := make([]int, columnCount)
for _, row := range rows {
for i := range columnCount {
cell := ""
if i < len(row) {
cell = row[i]
}
naturalWidths[i] = max(naturalWidths[i], lipglossWidth(cell))
}
}
widths := markdownTableColumnWidths(naturalWidths, width)
var rendered []string
for rowIndex, row := range rows {
wrappedCells := make([][]string, columnCount)
rowHeight := 1
for i := range columnCount {
cell := ""
if i < len(row) {
cell = row[i]
}
wrappedCells[i] = wrapMarkdownTableCell(cell, widths[i])
rowHeight = max(rowHeight, len(wrappedCells[i]))
}
for lineIndex := range rowHeight {
cells := make([]string, columnCount)
for i := range columnCount {
cellLine := ""
if lineIndex < len(wrappedCells[i]) {
cellLine = wrappedCells[i][lineIndex]
}
cells[i] = padPlainLine(cellLine, widths[i])
}
line := strings.Join(cells, chatTableBorderStyle.Render(" | "))
if rowIndex == 0 {
line = chatHeaderStyle.Render(stripANSIForWidth(line))
}
rendered = append(rendered, line)
}
}
return rendered, consumed
}
func markdownTableColumnWidths(naturalWidths []int, width int) []int {
if len(naturalWidths) == 0 {
return nil
}
separatorWidth := max(0, len(naturalWidths)-1) * lipglossWidth(" | ")
available := max(1, width-separatorWidth)
widths := make([]int, len(naturalWidths))
minWidths := make([]int, len(naturalWidths))
for i, natural := range naturalWidths {
widths[i] = max(1, natural)
minWidth := min(widths[i], 12)
if i == 0 {
minWidth = min(widths[i], 4)
}
minWidths[i] = max(1, minWidth)
}
for sumInts(widths) > available {
index := widestShrinkableColumn(widths, minWidths)
if index < 0 {
break
}
widths[index]--
}
for sumInts(widths) > available {
index := widestColumn(widths)
if index < 0 || widths[index] <= 1 {
break
}
widths[index]--
}
return widths
}
func widestShrinkableColumn(widths, minWidths []int) int {
index := -1
for i, width := range widths {
if width <= minWidths[i] {
continue
}
if index < 0 || width > widths[index] {
index = i
}
}
return index
}
func widestColumn(widths []int) int {
index := -1
for i, width := range widths {
if index < 0 || width > widths[index] {
index = i
}
}
return index
}
func sumInts(values []int) int {
sum := 0
for _, value := range values {
sum += value
}
return sum
}
func wrapMarkdownTableCell(cell string, width int) []string {
width = max(1, width)
var out []string
line := strings.TrimSpace(cell)
for lipglossWidth(line) > width {
cut := chatDisplayWidthCut(line, width)
out = append(out, strings.TrimSpace(line[:cut]))
line = strings.TrimSpace(line[cut:])
}
out = append(out, line)
if len(out) == 0 {
return []string{""}
}
return out
}
func looksLikeMarkdownTableRow(line string) bool {
line = strings.TrimSpace(line)
return strings.Contains(line, "|") && strings.Count(line, "|") >= 1
}
func isMarkdownTableSeparator(line string) bool {
cells := parseMarkdownTableRow(line)
if len(cells) == 0 {
return false
}
for _, cell := range cells {
cell = strings.Trim(cell, " :-")
if cell != "" {
return false
}
}
return true
}
func parseMarkdownTableRow(line string) []string {
line = strings.TrimSpace(line)
line = strings.TrimPrefix(line, "|")
line = strings.TrimSuffix(line, "|")
raw := strings.Split(line, "|")
cells := make([]string, 0, len(raw))
for _, cell := range raw {
cells = append(cells, strings.TrimSpace(cell))
}
return cells
}
func padPlainLine(line string, width int) string {
if extra := width - lipglossWidth(line); extra > 0 {
return line + strings.Repeat(" ", extra)
}
return line
}
func stripANSIForWidth(line string) string {
return stripChatANSI(line)
}
func lipglossWidth(line string) int {
return lipgloss.Width(line)
}
+263
View File
@@ -0,0 +1,263 @@
package chat
import (
"context"
"fmt"
"slices"
"strings"
tea "github.com/charmbracelet/bubbletea"
apptui "github.com/ollama/ollama/cmd/tui"
)
type chatModelPicker = apptui.SelectorModel
func (m *chatModel) openModelPicker(filter string) (tea.Model, tea.Cmd) {
if m.opts.ModelOptions == nil {
m.entries = append(m.entries, newChatEntry(chatEntry{role: "error", content: "Model picker is unavailable.", err: "Model picker is unavailable."}))
m.status = "error"
return *m, nil
}
ctx := m.ctx
if ctx == nil {
ctx = context.Background()
}
models, err := m.opts.ModelOptions(ctx)
if err != nil {
m.entries = append(m.entries, newChatEntry(chatEntry{role: "error", content: fmt.Sprintf("Could not list models: %v", err), err: err.Error()}))
m.status = "error"
return *m, nil
}
models = normalizeModelOptions(models)
if len(models) == 0 {
m.entries = append(m.entries, newChatEntry(chatEntry{role: "system", content: "No models available."}))
m.status = "ready"
return *m, nil
}
items := modelSelectorItems(models, m.opts.Model)
current := m.opts.Model
if !m.openModelOnInit {
items = compactModelSelectorItems(models, m.opts.Model)
current = ""
}
picker := apptui.NewModelSelectorModel("Select model", items, current, filter)
picker.SetHelpText("↑/↓ navigate • enter select • type search • esc cancel")
m.modelPicker = &picker
m.modelPickerModels = models
m.status = "model"
return *m, nil
}
func normalizeModelOptions(models []ModelOption) []ModelOption {
seen := make(map[string]struct{}, len(models))
out := make([]ModelOption, 0, len(models))
for _, model := range models {
model.Name = strings.TrimSpace(model.Name)
model.Description = strings.TrimSpace(model.Description)
if model.Name == "" {
continue
}
key := strings.ToLower(model.Name)
if _, ok := seen[key]; ok {
continue
}
seen[key] = struct{}{}
out = append(out, model)
}
slices.SortStableFunc(out, func(a, b ModelOption) int {
if a.Recommended == b.Recommended {
return 0
}
if a.Recommended {
return -1
}
return 1
})
return out
}
func modelSelectorItems(models []ModelOption, current string) []apptui.SelectItem {
return modelSelectorItemsWithCurrentPriority(models, current, true)
}
func compactModelSelectorItems(models []ModelOption, current string) []apptui.SelectItem {
return modelSelectorItemsWithCurrentPriority(models, current, false)
}
func modelSelectorItemsWithCurrentPriority(models []ModelOption, current string, pinCurrent bool) []apptui.SelectItem {
ordered := slices.Clone(models)
slices.SortStableFunc(ordered, func(a, b ModelOption) int {
if cmp := compareModelPickerGroup(modelPickerGroup(a, current, pinCurrent), modelPickerGroup(b, current, pinCurrent)); cmp != 0 {
return cmp
}
return 0
})
items := make([]apptui.SelectItem, 0, len(ordered))
for _, model := range ordered {
items = append(items, apptui.SelectItem{
Name: model.Name,
Description: modelOptionMeta(model),
Recommended: model.Name == current || !model.Cloud || model.Recommended,
AvailabilityBadge: model.AvailabilityBadge,
})
}
return items
}
func modelPickerGroup(model ModelOption, current string, pinCurrent bool) int {
if pinCurrent && model.Name == current {
return 0
}
if model.Recommended {
return 1
}
if model.Name == current {
return 2
}
if !model.Cloud {
return 3
}
return 4
}
func compareModelPickerGroup(a, b int) int {
switch {
case a < b:
return -1
case a > b:
return 1
default:
return 0
}
}
func (m chatModel) updateModelPicker(msg tea.KeyMsg) (tea.Model, tea.Cmd) {
if m.modelPicker == nil {
return m, nil
}
switch msg.Type {
case tea.KeyCtrlC, tea.KeyEsc:
m.modelPicker = nil
m.modelPickerModels = nil
m.openModelOnInit = false
m.status = "ready"
return m, nil
case tea.KeyEnter:
return m.selectModel()
default:
m.modelPicker.UpdateNavigation(msg)
}
return m, nil
}
func (m chatModel) selectModel() (tea.Model, tea.Cmd) {
if m.modelPicker == nil {
return m, nil
}
selectedItem, ok := m.modelPicker.SelectedItem()
if !ok {
return m, nil
}
selected, ok := m.modelOptionForSelection(selectedItem.Name)
if !ok {
return m, nil
}
// Cloud models need auth + plan check before switching. If we already
// know the badge state from the model list, go directly to the right
// prompt — no "checking" spinner.
if selected.Cloud && m.opts.CheckCloudModel != nil {
switch selected.AvailabilityBadge {
case "Sign in required":
return m.startCloudAuthSignIn(selected.Name, selected.RequiredPlan, selected.SignInURL)
case "Upgrade required":
return m.startCloudAuthUpgrade(selected.Name, selected.RequiredPlan)
}
// Badge is empty — auth is satisfied (confirmed via Whoami when the
// list was built). Apply directly.
}
m.modelPicker = nil
m.modelPickerModels = nil
m.openModelOnInit = false
if err := m.applyModelSelection(selected.Name, true); err != nil {
m.entries = append(m.entries, newChatEntry(chatEntry{role: "error", content: fmt.Sprintf("Could not switch model: %v", err), err: err.Error()}))
m.status = "error"
return m, nil
}
m.status = "ready"
return m, tea.Batch(m.startModelPreload(selected.Name), cloudModelPreflightCmd(m.ctx, m.opts, selected.Name, selected.RequiredPlan))
}
func (m chatModel) modelOptionForSelection(name string) (ModelOption, bool) {
for _, model := range m.modelPickerModels {
if model.Name == name {
return model, true
}
}
return ModelOption{}, false
}
func (m *chatModel) applyModelSelection(modelName string, persist bool) error {
modelName = strings.TrimSpace(modelName)
if modelName == "" {
return nil
}
m.opts.Model = modelName
m.opts.ContextWindowTokens = 0
if m.opts.ToolRegistryForModel != nil {
m.opts.Tools = m.opts.ToolRegistryForModel(m.ctx, modelName)
}
if m.opts.SystemPromptForModel != nil {
m.opts.SystemPrompt = m.opts.SystemPromptForModel(m.ctx, modelName, m.opts.Tools, m.opts.ToolsDisabled)
}
if m.opts.MultiModalForModel != nil {
ctx := m.ctx
if ctx == nil {
ctx = context.Background()
}
m.opts.MultiModal = m.opts.MultiModalForModel(ctx, modelName)
}
m.refreshContextWindowTokens(modelName)
m.contextTokens = m.estimatePromptTokens(m.messages, "")
m.contextEstimate = true
if persist && m.opts.OnModelSelected != nil {
ctx := m.ctx
if ctx == nil {
ctx = context.Background()
}
return m.opts.OnModelSelected(ctx, modelName)
}
return nil
}
func (m *chatModel) startModelPreload(modelName string) tea.Cmd {
modelName = strings.TrimSpace(modelName)
if m == nil || modelName == "" || m.opts.PreloadModel == nil {
return nil
}
m.preloadingModel = modelName
m.spinner = 0
return tea.Batch(preloadModelCmd(m.ctx, m.opts.PreloadModel, modelName, m.opts.Think), m.scheduleTick())
}
func (m chatModel) renderModelPicker(width int) string {
return m.modelPicker.RenderContent()
}
func (m chatModel) renderInlineModelPicker(width int) []string {
rendered := m.modelPicker.RenderCompactContent(maxInlineModelPickerItems)
lines := strings.Split(strings.TrimRight(rendered, "\n"), "\n")
for i := range lines {
lines[i] = truncateRenderedLine(lines[i], width)
}
return lines
}
func modelOptionMeta(model ModelOption) string {
return strings.TrimSpace(model.Description)
}
+441
View File
@@ -0,0 +1,441 @@
package chat
import (
"context"
"slices"
"strings"
"testing"
tea "github.com/charmbracelet/bubbletea"
coreagent "github.com/ollama/ollama/agent"
"github.com/ollama/ollama/api"
apptui "github.com/ollama/ollama/cmd/tui"
)
func TestChatModelCommandOpensPicker(t *testing.T) {
m := chatModel{
ctx: context.Background(),
input: []rune("/model"),
width: 100,
height: 20,
opts: Options{
Model: "llama3.2",
ContextWindowTokens: 131072,
ModelOptions: func(context.Context) ([]ModelOption, error) {
return []ModelOption{
{Name: "kimi-k2.6:cloud", Description: "cloud coding"},
{Name: "llama3.2", Description: "local"},
}, nil
},
},
}
updated, cmd := m.handleSubmit()
if cmd != nil {
t.Fatal("model command should not return a command")
}
m = updated.(chatModel)
if m.modelPicker == nil {
t.Fatal("model picker was not opened")
}
view := stripANSI(m.View())
if !strings.Contains(view, "Select model") ||
!strings.Contains(view, "Type to filter") ||
!strings.Contains(view, "kimi-k2.6:cloud") ||
!strings.Contains(view, "llama3.2") {
t.Fatalf("model picker view missing content: %q", view)
}
if strings.Contains(view, "Search...") {
t.Fatalf("model picker should render inline without full search box: %q", view)
}
if strings.Contains(view, "local") || strings.Contains(view, "cloud coding") {
t.Fatalf("inline model picker should stay compact without descriptions: %q", view)
}
if !strings.Contains(view, "│ █") {
t.Fatalf("inline model picker should keep input box visible: %q", view)
}
}
func TestChatModelCommandShowsRecommendedFirstWithoutSections(t *testing.T) {
m := chatModel{
ctx: context.Background(),
input: []rune("/model"),
width: 100,
height: 20,
opts: Options{
Model: "llama3.2",
ModelOptions: func(context.Context) ([]ModelOption, error) {
return []ModelOption{
{Name: "llama3.2", Description: "selected local"},
{Name: "gemma4", Description: "local"},
{Name: "glm-5.2:cloud", Description: "recommended cloud", Recommended: true, Cloud: true},
{Name: "kimi-k2.7-code:cloud", Description: "another recommended cloud", Recommended: true, Cloud: true},
}, nil
},
},
}
updated, cmd := m.handleSubmit()
if cmd != nil {
t.Fatal("model command should not return a command")
}
view := stripANSI(updated.(chatModel).View())
for _, unwanted := range []string{"Recommended", "More", "recommended cloud", "selected local"} {
if strings.Contains(view, unwanted) {
t.Fatalf("compact model picker should be flat and description-free; found %q in %q", unwanted, view)
}
}
firstRecommended := strings.Index(view, "glm-5.2:cloud")
secondRecommended := strings.Index(view, "kimi-k2.7-code:cloud")
current := strings.Index(view, "llama3.2")
local := strings.Index(view, "gemma4")
if firstRecommended < 0 || secondRecommended < 0 || current < 0 || local < 0 {
t.Fatalf("compact model picker missing expected models: %q", view)
}
if !(firstRecommended < current && secondRecommended < current && current < local) {
t.Fatalf("compact model picker order should be recommended, current, local: %q", view)
}
}
func TestChatModelCommandOpensSmallPicker(t *testing.T) {
m := chatModel{
ctx: context.Background(),
input: []rune("/model"),
width: 100,
height: 24,
opts: Options{
Model: "model-1",
ModelOptions: func(context.Context) ([]ModelOption, error) {
return []ModelOption{
{Name: "model-1"},
{Name: "model-2"},
{Name: "model-3"},
{Name: "model-4"},
{Name: "model-5"},
{Name: "model-6"},
{Name: "model-7"},
}, nil
},
},
}
updated, cmd := m.handleSubmit()
if cmd != nil {
t.Fatal("model command should not return a command")
}
m = updated.(chatModel)
view := stripANSI(m.View())
for _, want := range []string{"model-1", "model-5", "... and 2 more"} {
if !strings.Contains(view, want) {
t.Fatalf("small model picker missing %q: %q", want, view)
}
}
if strings.Contains(view, "model-6") || strings.Contains(view, "model-7") {
t.Fatalf("small model picker rendered too many items: %q", view)
}
}
func TestChatModelPickerStaysInlineWhenSmall(t *testing.T) {
m := chatModel{
ctx: context.Background(),
input: []rune("/model"),
width: 44,
height: 10,
opts: Options{
Model: "llama3.2",
ModelOptions: func(context.Context) ([]ModelOption, error) {
return []ModelOption{
{Name: "kimi-k2.6:cloud", Description: "cloud coding"},
{Name: "llama3.2", Description: "local"},
}, nil
},
},
}
updated, cmd := m.handleSubmit()
if cmd != nil {
t.Fatal("model command should not return a command")
}
m = updated.(chatModel)
view := stripANSI(m.View())
if !strings.Contains(view, "Select model") || !strings.Contains(view, "Type to filter") {
t.Fatalf("small model picker should stay inline: %q", view)
}
if strings.Contains(view, "Search...") {
t.Fatalf("small model picker should not use bespoke full-frame search: %q", view)
}
}
func TestChatModelPickerShowsRecommendedModelsFirst(t *testing.T) {
models := normalizeModelOptions([]ModelOption{
{Name: "llama3.2", Description: "local"},
{Name: "kimi-k2.6:cloud", Description: "cloud coding", Recommended: true},
{Name: "qwen3.5:cloud", Description: "cloud reasoning", Recommended: true},
{Name: "gemma4", Description: "local"},
})
got := make([]string, 0, len(models))
for _, model := range models {
got = append(got, model.Name)
}
want := []string{"kimi-k2.6:cloud", "qwen3.5:cloud", "llama3.2", "gemma4"}
if !slices.Equal(got, want) {
t.Fatalf("model order = %#v, want %#v", got, want)
}
}
func TestChatModelPickerPinsCurrentThenRecommendedModels(t *testing.T) {
models := normalizeModelOptions([]ModelOption{
{Name: "llama3.2", Description: "local"},
{Name: "glm-5.2:cloud", Description: "cloud selected", Recommended: true, Cloud: true},
{Name: "kimi-k2.7-code:cloud", Description: "cloud coding", Recommended: true, Cloud: true},
{Name: "gemma4", Description: "local"},
})
items := modelSelectorItems(models, "glm-5.2:cloud")
got := make([]string, 0, len(items))
for _, item := range items {
got = append(got, item.Name)
}
want := []string{"glm-5.2:cloud", "kimi-k2.7-code:cloud", "llama3.2", "gemma4"}
if !slices.Equal(got, want) {
t.Fatalf("selector item order = %#v, want %#v", got, want)
}
for _, item := range items[:3] {
if !item.Recommended {
t.Fatalf("%q should be pinned in the first picker section", item.Name)
}
}
if items[0].Description != "cloud selected" {
t.Fatalf("current model description = %q, want plain model description", items[0].Description)
}
}
func TestInitialModelPickerRendersBeforeChatShell(t *testing.T) {
models := normalizeModelOptions([]ModelOption{
{Name: "glm-5.2:cloud", Description: "cloud selected", Recommended: true, Cloud: true},
{Name: "llama3.2", Description: "local"},
})
picker := apptui.NewModelSelectorModel("Select model", modelSelectorItems(models, "glm-5.2:cloud"), "glm-5.2:cloud", "")
m := chatModel{
width: 100,
height: 20,
openModelOnInit: true,
modelPicker: &picker,
entries: []chatEntry{{role: "assistant", content: "old chat content"}},
}
view := stripANSI(m.View())
if !strings.Contains(view, "Select model") || !strings.Contains(view, "llama3.2") {
t.Fatalf("initial picker view missing model content: %q", view)
}
if strings.Contains(view, "old chat content") || strings.Contains(view, "│ █") {
t.Fatalf("initial picker should render before chat shell: %q", view)
}
}
func TestChatModelPickerRanksClosestFilteredModelFirst(t *testing.T) {
models := normalizeModelOptions([]ModelOption{
{Name: "gemma3:27b", Description: "recommended but longer", Recommended: true},
{Name: "llama3.2", Description: "mentions gemm in description"},
{Name: "gemma4:27b", Description: "longer local"},
{Name: "gemma4", Description: "short local"},
})
picker := apptui.NewModelSelectorModel("Select model", modelSelectorItems(models, ""), "", "gemm")
filtered := picker.FilteredItems()
got := make([]string, 0, len(filtered))
for _, model := range filtered {
got = append(got, model.Name)
}
want := []string{"gemma4", "gemma3:27b", "gemma4:27b", "llama3.2"}
if !slices.Equal(got, want) {
t.Fatalf("filtered model order = %#v, want %#v", got, want)
}
}
func TestChatModelPickerFiltersAndSwitchesModel(t *testing.T) {
var savedModel string
originalMessages := []api.Message{{Role: "user", Content: "keep me"}}
m := chatModel{
ctx: context.Background(),
chatID: "chat-1",
input: []rune("/model qwen"),
width: 100,
height: 20,
messages: slices.Clone(originalMessages),
opts: Options{
Model: "llama3.2",
ModelOptions: func(context.Context) ([]ModelOption, error) {
return []ModelOption{
{Name: "llama3.2", Description: "local"},
{Name: "qwen3.5:cloud", Description: "cloud reasoning"},
}, nil
},
ToolRegistryForModel: func(ctx context.Context, model string) *coreagent.Registry {
if model != "qwen3.5:cloud" {
t.Fatalf("tool registry model = %q, want qwen3.5:cloud", model)
}
registry := &coreagent.Registry{}
registry.Register(chatTestTool{})
return registry
},
ContextWindowTokensForModel: func(ctx context.Context, model string, fallback int) int {
if model != "qwen3.5:cloud" {
t.Fatalf("context model = %q, want qwen3.5:cloud", model)
}
if fallback != 0 {
t.Fatalf("context fallback = %d, want 0 after model switch", fallback)
}
return 262144
},
SystemPromptForModel: func(ctx context.Context, model string, registry *coreagent.Registry, toolsDisabled bool) string {
if model != "qwen3.5:cloud" {
t.Fatalf("system prompt model = %q, want qwen3.5:cloud", model)
}
if registry == nil {
t.Fatalf("system prompt registry missing fake tool: %#v", registry)
}
if _, ok := registry.Get("fake_tool"); !ok {
t.Fatalf("system prompt registry missing fake tool: %#v", registry)
}
return "system for " + model
},
OnModelSelected: func(ctx context.Context, model string) error {
savedModel = model
return nil
},
},
}
updated, cmd := m.handleSubmit()
if cmd != nil {
t.Fatal("model command should not return a command")
}
m = updated.(chatModel)
if m.modelPicker == nil || m.modelPicker.Filter() != "qwen" {
t.Fatalf("model picker = %#v, want qwen filter", m.modelPicker)
}
if view := stripANSI(m.View()); !strings.Contains(view, "qwen3.5:cloud") || strings.Contains(view, "llama3.2") {
t.Fatalf("filtered model picker view = %q", view)
}
updated, cmd = m.Update(tea.KeyMsg{Type: tea.KeyEnter})
m = updated.(chatModel)
if cmd != nil {
t.Fatal("switching models should not start a command")
}
if m.modelPicker != nil {
t.Fatal("model picker should close after selection")
}
if m.status != "ready" || m.notificationLine() != "" {
t.Fatalf("model switch should not show action status, status=%q notification=%q", m.status, m.notificationLine())
}
if m.opts.Model != "qwen3.5:cloud" {
t.Fatalf("model = %q, want qwen3.5:cloud", m.opts.Model)
}
if m.chatID != "chat-1" {
t.Fatalf("chatID = %q, want chat-1", m.chatID)
}
if len(m.messages) != len(originalMessages) || m.messages[0].Content != originalMessages[0].Content {
t.Fatalf("messages changed on model switch: %#v", m.messages)
}
if len(m.entries) != 0 {
t.Fatalf("model switch should not append transcript entries: %#v", m.entries)
}
if savedModel != "qwen3.5:cloud" {
t.Fatalf("saved model = %q, want qwen3.5:cloud", savedModel)
}
if m.opts.Tools == nil {
t.Fatalf("tools registry was not rebuilt for model: %#v", m.opts.Tools)
}
if _, ok := m.opts.Tools.Get("fake_tool"); !ok {
t.Fatalf("tools registry was not rebuilt for model: %#v", m.opts.Tools)
}
if m.opts.ContextWindowTokens != 262144 {
t.Fatalf("context window = %d, want 262144", m.opts.ContextWindowTokens)
}
if m.opts.SystemPrompt != "system for qwen3.5:cloud" {
t.Fatalf("system prompt = %q", m.opts.SystemPrompt)
}
}
func TestChatModelSelectionStartsBackgroundPreload(t *testing.T) {
m := chatModel{
ctx: context.Background(),
opts: Options{
Model: "llama3.2",
PreloadModel: func(context.Context, string, *api.ThinkValue) (int, error) {
return 0, nil
},
},
}
if err := m.applyModelSelection("qwen3", false); err != nil {
t.Fatal(err)
}
cmd := m.startModelPreload("qwen3")
if cmd == nil {
t.Fatal("model switch should start background preload when configured")
}
if m.preloadingModel != "qwen3" {
t.Fatalf("preloadingModel = %q, want qwen3", m.preloadingModel)
}
}
func TestChatModelSwitchNextRunKeepsHistory(t *testing.T) {
client := &chatCaptureClient{}
history := []api.Message{
{Role: "user", Content: "old question"},
{Role: "assistant", Content: "old answer"},
}
m := chatModel{
ctx: context.Background(),
chatID: "chat-1",
messages: slices.Clone(history),
input: []rune("continue"),
opts: Options{
Model: "llama3.2",
Client: client,
SystemPromptForModel: func(_ context.Context, model string, _ *coreagent.Registry, _ bool) string {
return "system for " + model
},
},
}
if err := m.applyModelSelection("qwen3", true); err != nil {
t.Fatal(err)
}
updated, cmd := m.handleSubmit()
m = updated.(chatModel)
if cmd == nil {
t.Fatal("next prompt should start a model run")
}
done := waitForRunDone(t, m.events)
if done.err != nil {
t.Fatal(done.err)
}
if len(client.requests) != 1 {
t.Fatalf("requests = %d, want 1", len(client.requests))
}
req := client.requests[0]
if req.Model != "qwen3" {
t.Fatalf("request model = %q, want qwen3", req.Model)
}
if len(req.Messages) != 4 {
t.Fatalf("request messages = %#v, want system + 2 history + new user", req.Messages)
}
if req.Messages[0].Role != "system" || req.Messages[0].Content != "system for qwen3" {
t.Fatalf("system message = %#v", req.Messages[0])
}
for i, want := range history {
got := req.Messages[i+1]
if got.Role != want.Role || got.Content != want.Content {
t.Fatalf("history message %d = %#v, want %#v", i, got, want)
}
}
if req.Messages[3].Role != "user" || req.Messages[3].Content != "continue" {
t.Fatalf("new user message = %#v", req.Messages[3])
}
}
+307
View File
@@ -0,0 +1,307 @@
package chat
import (
"context"
"net/url"
"os"
"path/filepath"
"strings"
"testing"
tea "github.com/charmbracelet/bubbletea"
)
func TestChatStartRunAttachesDroppedImagePath(t *testing.T) {
fp := writeTestPNG(t)
m := chatModel{
ctx: context.Background(),
opts: Options{
Model: "test",
Client: chatTestClient{},
MultiModal: true,
},
}
updated, cmd := m.startRun("describe " + fp)
m = updated.(chatModel)
if cmd == nil {
t.Fatal("startRun should return a command")
}
if len(m.liveMessages) != 1 {
t.Fatalf("liveMessages = %d, want 1", len(m.liveMessages))
}
if got := m.liveMessages[0].Content; got != "describe" {
t.Fatalf("content = %q, want describe", got)
}
if got := len(m.liveMessages[0].Images); got != 1 {
t.Fatalf("images = %d, want 1", got)
}
if len(m.entries) == 0 {
t.Fatal("missing user transcript entry")
}
entry := m.entries[0].content
if strings.Contains(entry, fp) {
t.Fatalf("transcript entry should hide local file path: %q", entry)
}
if !strings.Contains(entry, "describe") || !strings.Contains(entry, "[attached 1 file]") {
t.Fatalf("transcript entry = %q, want prompt plus attachment note", entry)
}
}
func TestChatStartRunAttachesDroppedFileURL(t *testing.T) {
fp := writeTestPNG(t)
fileURL := (&url.URL{Scheme: "file", Path: fp}).String()
m := chatModel{
ctx: context.Background(),
opts: Options{
Model: "test",
Client: chatTestClient{},
MultiModal: true,
},
}
updated, _ := m.startRun(fileURL)
m = updated.(chatModel)
if got := m.liveMessages[0].Content; got != "" {
t.Fatalf("content = %q, want empty prompt after extracting file URL", got)
}
if got := len(m.liveMessages[0].Images); got != 1 {
t.Fatalf("images = %d, want 1", got)
}
if got := m.entries[0].content; got != "[attached 1 file]" {
t.Fatalf("transcript entry = %q, want attachment-only note", got)
}
}
func TestChatPasteImagePathAttachesOnSubmit(t *testing.T) {
fp := writeTestPNG(t)
m := chatModel{
ctx: context.Background(),
opts: Options{
Model: "test",
Client: chatTestClient{},
MultiModal: true,
},
}
updated, _ := m.Update(tea.KeyMsg{Type: tea.KeyRunes, Runes: []rune("describe " + fp), Paste: true})
m = updated.(chatModel)
if got := string(m.input); got != "describe [Image #0]" {
t.Fatalf("pasted path input = %q, want placeholder", got)
}
if got := m.notificationLine(); got != "" {
t.Fatalf("notification = %q, want no attachment notification", got)
}
if got := string(m.input); strings.Contains(got, fp) {
t.Fatalf("pasted path should be hidden behind placeholder, input = %q", got)
}
if completions := m.slashCompletions(); len(completions) != 0 {
t.Fatalf("placeholder input should not show slash completions: %#v", completions)
}
updated, cmd := m.Update(tea.KeyMsg{Type: tea.KeyEnter})
m = updated.(chatModel)
if cmd == nil {
t.Fatal("submit should start a run")
}
if got := m.liveMessages[0].Content; got != "describe [Image #0]" {
t.Fatalf("content = %q, want prompt with placeholder", got)
}
if got := len(m.liveMessages[0].Images); got != 1 {
t.Fatalf("images = %d, want 1", got)
}
if strings.Contains(m.entries[0].content, fp) {
t.Fatalf("transcript entry should hide pasted file path: %q", m.entries[0].content)
}
if !strings.Contains(m.entries[0].content, "[Image #0]") {
t.Fatalf("transcript entry should show placeholder: %q", m.entries[0].content)
}
}
func TestChatPasteImagePathAfterSwitchingToMultimodalModel(t *testing.T) {
fp := writeTestPNG(t)
m := chatModel{
ctx: context.Background(),
input: []rune("/model vision"),
opts: Options{
Model: "text",
ModelOptions: func(context.Context) ([]ModelOption, error) {
return []ModelOption{
{Name: "text", Description: "local"},
{Name: "vision", Description: "local vision"},
}, nil
},
MultiModalForModel: func(_ context.Context, model string) bool {
return model == "vision"
},
},
}
updated, cmd := m.handleSubmit()
if cmd != nil {
t.Fatal("model picker should not return a command")
}
m = updated.(chatModel)
updated, cmd = m.Update(tea.KeyMsg{Type: tea.KeyEnter})
if cmd != nil {
t.Fatal("model switch should not preload without a preload hook")
}
m = updated.(chatModel)
if m.opts.Model != "vision" {
t.Fatalf("model = %q, want vision", m.opts.Model)
}
if !m.opts.MultiModal {
t.Fatal("switching to a multimodal model should enable image paste handling")
}
updated, _ = m.Update(tea.KeyMsg{Type: tea.KeyRunes, Runes: []rune("describe " + fp), Paste: true})
m = updated.(chatModel)
if got := string(m.input); got != "describe [Image #0]" {
t.Fatalf("pasted path input = %q, want placeholder", got)
}
if strings.Contains(string(m.input), fp) {
t.Fatalf("pasted path should be hidden behind placeholder, input = %q", string(m.input))
}
}
func TestChatImagePlaceholdersUseSessionNumbers(t *testing.T) {
first := writeTestPNG(t)
second := writeTestPNG(t)
m := chatModel{
ctx: context.Background(),
opts: Options{
Model: "test",
Client: chatTestClient{},
MultiModal: true,
},
}
updated, _ := m.Update(tea.KeyMsg{Type: tea.KeyRunes, Runes: []rune(first), Paste: true})
m = updated.(chatModel)
if got := string(m.input); got != "[Image #0]" {
t.Fatalf("first placeholder = %q, want [Image #0]", got)
}
updated, cmd := m.Update(tea.KeyMsg{Type: tea.KeyEnter})
m = updated.(chatModel)
if cmd == nil {
t.Fatal("submit should start a run")
}
updated, _ = m.Update(tea.KeyMsg{Type: tea.KeyRunes, Runes: []rune(second), Paste: true})
m = updated.(chatModel)
if got := string(m.input); got != "[Image #1]" {
t.Fatalf("second placeholder = %q, want [Image #1]", got)
}
}
func TestChatAbsoluteImagePathBypassesSlashCommandParsing(t *testing.T) {
fp := writeTestPNG(t)
m := chatModel{
ctx: context.Background(),
opts: Options{
Model: "test",
Client: chatTestClient{},
MultiModal: true,
},
}
m.input = []rune(fp)
m.inputCursor = len(m.input)
m.inputCursorSet = true
updated, cmd := m.handleSubmit()
m = updated.(chatModel)
if cmd == nil {
t.Fatal("absolute image path should start a run instead of being parsed as a slash command")
}
if got := len(m.liveMessages[0].Images); got != 1 {
t.Fatalf("images = %d, want 1", got)
}
if got := m.entries[0].role; got != "user" {
t.Fatalf("entry role = %q, want user", got)
}
}
func TestChatDeletingImagePlaceholderRemovesAttachment(t *testing.T) {
fp := writeTestPNG(t)
m := chatModel{
ctx: context.Background(),
opts: Options{
Model: "test",
Client: chatTestClient{},
MultiModal: true,
},
}
updated, _ := m.Update(tea.KeyMsg{Type: tea.KeyRunes, Runes: []rune("describe " + fp), Paste: true})
m = updated.(chatModel)
if got := len(m.inputAttachments); got != 1 {
t.Fatalf("input attachments = %d, want 1", got)
}
updated, _ = m.Update(tea.KeyMsg{Type: tea.KeyBackspace})
m = updated.(chatModel)
if got := string(m.input); got != "describe " {
t.Fatalf("input after backspace = %q, want image placeholder removed", got)
}
if got := len(m.inputAttachments); got != 0 {
t.Fatalf("input attachments after editing placeholder = %d, want 0", got)
}
m.input = []rune("describe")
m.inputCursor = len(m.input)
m.inputCursorSet = true
updated, cmd := m.Update(tea.KeyMsg{Type: tea.KeyEnter})
m = updated.(chatModel)
if cmd == nil {
t.Fatal("submit should start a run")
}
if got := len(m.liveMessages[0].Images); got != 0 {
t.Fatalf("images = %d, want 0 after deleting placeholder", got)
}
if got := m.liveMessages[0].Content; got != "describe" {
t.Fatalf("content = %q, want describe", got)
}
}
func TestChatWordDeletingImagePlaceholderRemovesAttachment(t *testing.T) {
fp := writeTestPNG(t)
m := chatModel{
ctx: context.Background(),
opts: Options{
Model: "test",
Client: chatTestClient{},
MultiModal: true,
},
}
updated, _ := m.Update(tea.KeyMsg{Type: tea.KeyRunes, Runes: []rune("describe " + fp), Paste: true})
m = updated.(chatModel)
updated, _ = m.Update(tea.KeyMsg{Type: tea.KeySpace})
m = updated.(chatModel)
updated, _ = m.Update(tea.KeyMsg{Type: tea.KeyBackspace, Alt: true})
m = updated.(chatModel)
if got := string(m.input); got != "describe " {
t.Fatalf("input after word backspace = %q, want image placeholder removed", got)
}
if got := len(m.inputAttachments); got != 0 {
t.Fatalf("input attachments after word backspace = %d, want 0", got)
}
}
func writeTestPNG(t *testing.T) string {
t.Helper()
dir := t.TempDir()
fp := filepath.Join(dir, "dragged image.png")
data := make([]byte, 600)
copy(data, []byte{0x89, 'P', 'N', 'G', '\r', '\n', 0x1a, '\n'})
if err := os.WriteFile(fp, data, 0o600); err != nil {
t.Fatalf("failed to write test image: %v", err)
}
return fp
}
File diff suppressed because it is too large. Load diff
File diff suppressed because it is too large. Load diff
+113
View File
@@ -0,0 +1,113 @@
package chat
import (
"context"
"fmt"
"regexp"
"testing"
"time"
tea "github.com/charmbracelet/bubbletea"
coreagent "github.com/ollama/ollama/agent"
"github.com/ollama/ollama/api"
)
type chatTestTool struct{}
type chatTestClient struct{}
type chatCaptureClient struct {
requests []*api.ChatRequest
}
type chatToolLoopClient struct {
calls int
toolRounds int
}
func (chatTestClient) Chat(ctx context.Context, req *api.ChatRequest, fn api.ChatResponseFunc) error {
if err := ctx.Err(); err != nil {
return err
}
return fn(api.ChatResponse{
Message: api.Message{Role: "assistant", Content: "ok"},
Done: true,
})
}
func (c *chatCaptureClient) Chat(ctx context.Context, req *api.ChatRequest, fn api.ChatResponseFunc) error {
if err := ctx.Err(); err != nil {
return err
}
c.requests = append(c.requests, req)
return fn(api.ChatResponse{
Message: api.Message{Role: "assistant", Content: "ok"},
Done: true,
})
}
func (c *chatToolLoopClient) Chat(ctx context.Context, req *api.ChatRequest, fn api.ChatResponseFunc) error {
if err := ctx.Err(); err != nil {
return err
}
c.calls++
if c.calls > c.toolRounds {
return fn(api.ChatResponse{Message: api.Message{Role: "assistant", Content: "done"}, Done: true})
}
args := api.NewToolCallFunctionArguments()
args.Set("value", "keep going")
return fn(api.ChatResponse{Message: api.Message{Role: "assistant", ToolCalls: []api.ToolCall{{
ID: fmt.Sprintf("call-%d", c.calls),
Function: api.ToolCallFunction{
Name: "fake_tool",
Arguments: args,
},
}}}})
}
func (chatTestTool) Name() string {
return "fake_tool"
}
func (chatTestTool) Description() string {
return "does test work"
}
func (chatTestTool) Schema() api.ToolFunction {
return api.ToolFunction{
Name: "fake_tool",
Description: "does test work",
Parameters: api.ToolFunctionParameters{
Type: "object",
},
}
}
func (chatTestTool) Execute(context.Context, coreagent.ToolContext, map[string]any) (coreagent.ToolResult, error) {
return coreagent.ToolResult{Content: "ok"}, nil
}
func waitForRunDone(t *testing.T, events <-chan tea.Msg) chatRunDoneMsg {
t.Helper()
timeout := time.After(2 * time.Second)
for {
select {
case msg, ok := <-events:
if !ok {
t.Fatal("events closed before run done")
}
if done, ok := msg.(chatRunDoneMsg); ok {
return done
}
case <-timeout:
t.Fatal("timed out waiting for run done")
}
}
}
func stripANSI(s string) string {
re := regexp.MustCompile(`\x1b\[[0-9;:]*[A-Za-z]`)
return re.ReplaceAllString(s, "")
}
+126
View File
@@ -0,0 +1,126 @@
package chat
import "github.com/charmbracelet/lipgloss"
const (
chatAnsiRed = "1"
chatAnsiGreen = "2"
chatAnsiYellow = "3"
chatAnsiBlue = "4"
chatAnsiCyan = "6"
chatAnsiBrightBlack = "8"
)
var (
chatHeaderStyle = lipgloss.NewStyle().
Bold(true)
chatMetaStyle = lipgloss.NewStyle().
Faint(true)
chatFooterStyle = lipgloss.NewStyle().
Faint(true)
chatInputBorderStyle = lipgloss.NewStyle().
Faint(true)
chatInputPlaceholderStyle = lipgloss.NewStyle().
Foreground(lipgloss.Color("8"))
chatCursorStyle = lipgloss.NewStyle().
Reverse(true)
chatBlankCursorStyle = lipgloss.NewStyle().
Faint(true)
chatNotificationStyle = chatMetaStyle
chatUserStyle = lipgloss.NewStyle()
chatUserBlockStyle = lipgloss.NewStyle().
Foreground(lipgloss.AdaptiveColor{Light: "#777777", Dark: "#8a8a8a"})
chatToolStyle = lipgloss.NewStyle()
chatInlineCodeStyle = lipgloss.NewStyle().
Bold(true)
chatCodeBlockStyle = lipgloss.NewStyle()
chatTableBorderStyle = lipgloss.NewStyle().
Faint(true)
chatToolRunningStyle = lipgloss.NewStyle().
Foreground(lipgloss.Color(chatAnsiYellow))
chatToolDoneStyle = lipgloss.NewStyle().
Foreground(lipgloss.Color(chatAnsiGreen))
// chatToolMixedStyle marks a tool group with both succeeded and failed
// calls (partial success). Amber/orange is distinct from green (success),
// red (failure), and yellow (running).
chatToolMixedStyle = lipgloss.NewStyle().
Foreground(lipgloss.Color("208"))
chatToolOutputStyle = lipgloss.NewStyle().
Foreground(lipgloss.AdaptiveColor{Light: "#666666", Dark: "#a0a0a0"})
chatDiffMetaStyle = lipgloss.NewStyle().
Faint(true)
chatDiffFileStyle = lipgloss.NewStyle().
Bold(true).
Foreground(lipgloss.Color(chatAnsiCyan))
chatDiffHunkStyle = lipgloss.NewStyle().
Foreground(lipgloss.Color(chatAnsiBlue))
chatDiffAddStyle = lipgloss.NewStyle().
Foreground(lipgloss.Color(chatAnsiGreen))
chatDiffDeleteStyle = lipgloss.NewStyle().
Foreground(lipgloss.Color(chatAnsiRed))
chatErrorStyle = lipgloss.NewStyle().
Foreground(lipgloss.Color(chatAnsiRed))
chatFullAccessStyle = lipgloss.NewStyle().
Foreground(lipgloss.AdaptiveColor{Light: "#9f5f5f", Dark: "#b87373"})
chatCommandNameStyle = lipgloss.NewStyle()
chatPickerTextStyle = lipgloss.NewStyle()
chatPickerTitleStyle = lipgloss.NewStyle().
Bold(true)
chatPickerSelectedStyle = lipgloss.NewStyle().
Bold(true)
chatPickerMetaStyle = lipgloss.NewStyle().
Faint(true)
chatHistoryTitleStyle = lipgloss.NewStyle().
Bold(true)
chatHistorySystemRoleStyle = lipgloss.NewStyle().
Bold(true).
Faint(true)
chatHistoryUserRoleStyle = lipgloss.NewStyle().
Bold(true).
Foreground(lipgloss.Color(chatAnsiBlue))
chatHistoryAssistantRoleStyle = lipgloss.NewStyle().
Bold(true).
Foreground(lipgloss.Color(chatAnsiYellow))
chatHistoryToolRoleStyle = lipgloss.NewStyle().
Bold(true).
Foreground(lipgloss.Color(chatAnsiGreen))
chatHistoryLabelStyle = lipgloss.NewStyle().
Faint(true)
chatHistoryTextStyle = lipgloss.NewStyle()
)
+166
View File
@@ -0,0 +1,166 @@
package chat
import (
"fmt"
"strings"
tea "github.com/charmbracelet/bubbletea"
"github.com/ollama/ollama/api"
)
type chatThinkOption struct {
value string
label string
description string
}
type chatThinkPicker struct {
options []chatThinkOption
cursor int
}
var chatThinkOptions = []chatThinkOption{
{value: "auto", label: "auto", description: "use the model default"},
{value: "on", label: "on", description: "enable thinking"},
{value: "off", label: "off", description: "disable thinking"},
{value: "low", label: "low", description: "use low thinking effort"},
{value: "medium", label: "medium", description: "use medium thinking effort"},
{value: "high", label: "high", description: "use high thinking effort"},
{value: "max", label: "max", description: "use maximum thinking effort"},
}
func (m *chatModel) openThinkPicker() (tea.Model, tea.Cmd) {
m.thinkPicker = newChatThinkPicker(m.opts.Think)
m.status = "think"
return *m, nil
}
func newChatThinkPicker(current *api.ThinkValue) *chatThinkPicker {
picker := &chatThinkPicker{options: append([]chatThinkOption(nil), chatThinkOptions...)}
currentValue := thinkValueLabel(current)
for i, option := range picker.options {
if option.value == currentValue {
picker.cursor = i
break
}
}
return picker
}
func (m chatModel) updateThinkPicker(msg tea.KeyMsg) (tea.Model, tea.Cmd) {
switch msg.Type {
case tea.KeyCtrlC, tea.KeyEsc:
m.thinkPicker = nil
m.status = "ready"
case tea.KeyEnter:
return m.selectThinkOption()
case tea.KeyUp:
m.thinkPicker.move(-1)
case tea.KeyDown:
m.thinkPicker.move(1)
}
return m, nil
}
func (p *chatThinkPicker) move(delta int) {
if p == nil || len(p.options) == 0 || delta == 0 {
return
}
p.cursor = clamp(p.cursor+delta, 0, len(p.options)-1)
}
func (p *chatThinkPicker) selected() (chatThinkOption, bool) {
if p == nil || len(p.options) == 0 {
return chatThinkOption{}, false
}
return p.options[clamp(p.cursor, 0, len(p.options)-1)], true
}
func (m chatModel) selectThinkOption() (tea.Model, tea.Cmd) {
option, ok := m.thinkPicker.selected()
if !ok {
return m, nil
}
m.thinkPicker = nil
return m.applyThinkValue(option.value)
}
func (m *chatModel) handleThinkCommand(value string) (tea.Model, tea.Cmd) {
return m.applyThinkValue(value)
}
func (m *chatModel) applyThinkValue(value string) (tea.Model, tea.Cmd) {
think, label, err := parseThinkValue(value)
if err != nil {
m.entries = append(m.entries, newChatEntry(chatEntry{role: "error", content: err.Error(), err: err.Error()}))
m.status = "error"
return *m, nil
}
m.opts.Think = think
m.status = "think " + label
return *m, nil
}
func parseThinkValue(value string) (*api.ThinkValue, string, error) {
switch strings.ToLower(strings.TrimSpace(value)) {
case "", "auto", "default", "unset":
return nil, "auto", nil
case "on", "true", "think", "thinking":
return &api.ThinkValue{Value: true}, "on", nil
case "off", "false", "nothink", "no-think":
return &api.ThinkValue{Value: false}, "off", nil
case "low", "medium", "high", "max":
value = strings.ToLower(strings.TrimSpace(value))
return &api.ThinkValue{Value: value}, value, nil
default:
return nil, "", fmt.Errorf("Usage: /think [auto|on|off|low|medium|high|max]")
}
}
func thinkValueLabel(value *api.ThinkValue) string {
if value == nil || value.Value == nil {
return "auto"
}
switch v := value.Value.(type) {
case bool:
if v {
return "on"
}
return "off"
case string:
return strings.ToLower(v)
default:
return "auto"
}
}
func (m chatModel) renderThinkPicker(width int) string {
picker := m.thinkPicker
if picker == nil {
return ""
}
var b strings.Builder
b.WriteString(chatPickerTitleStyle.Render("Thinking mode"))
b.WriteString("\n\n")
for i, option := range picker.options {
selected := i == picker.cursor
if selected {
b.WriteString(chatPickerSelectedStyle.Render(" " + option.label))
} else {
b.WriteString(" ")
b.WriteString(chatPickerTextStyle.Render(option.label))
}
b.WriteByte('\n')
b.WriteString(chatPickerMetaStyle.Render(" " + option.description))
b.WriteByte('\n')
if i < len(picker.options)-1 {
b.WriteByte('\n')
}
}
b.WriteString("\n")
b.WriteString(chatPickerMetaStyle.Render("↑/↓ navigate • enter select • esc cancel"))
return b.String()
}
+2 -2
View File
@@ -39,7 +39,7 @@ func (m confirmModel) Update(msg tea.Msg) (tea.Model, tea.Cmd) {
wasSet := m.width > 0
m.width = msg.Width
if wasSet {
return m, tea.EnterAltScreen
return m, tea.ClearScreen
}
return m, nil
@@ -115,7 +115,7 @@ func RunConfirmWithOptions(prompt string, options ConfirmOptions) (bool, error)
prompt: prompt,
yesLabel: yesLabel,
noLabel: noLabel,
yes: true, // default to yes
yes: options.Default != launch.ConfirmDefaultNo,
}
p := tea.NewProgram(m)
+245 -10
View File
@@ -2,6 +2,7 @@ package tui
import (
"fmt"
"sort"
"strings"
tea "github.com/charmbracelet/bubbletea"
@@ -63,6 +64,8 @@ type SelectItem struct {
AvailabilityBadge string
}
type SelectorModel = selectorModel
type selectorItemsUpdatedMsg struct {
items []SelectItem
}
@@ -121,6 +124,7 @@ type selectorModel struct {
cancelled bool
helpText string
width int
rankFiltered bool
}
func selectorModelWithCurrent(title string, items []SelectItem, current string) selectorModel {
@@ -133,6 +137,21 @@ func selectorModelWithCurrent(title string, items []SelectItem, current string)
return m
}
func NewSelectorModel(title string, items []SelectItem, current string) SelectorModel {
return selectorModelWithCurrent(title, items, current)
}
func NewModelSelectorModel(title string, items []SelectItem, current, filter string) SelectorModel {
m := selectorModelWithCurrent(title, items, current)
m.filter = strings.TrimSpace(filter)
m.rankFiltered = true
if m.filter != "" {
m.cursor = 0
m.scrollOffset = 0
}
return m
}
func currentItemName(items []SelectItem, cursor int) string {
if cursor < 0 || cursor >= len(items) {
return ""
@@ -140,15 +159,22 @@ func currentItemName(items []SelectItem, cursor int) string {
return items[cursor].Name
}
func indexOfItemName(items []SelectItem, name string) int {
for i, item := range items {
if item.Name == name {
return i
}
}
return -1
}
func cursorForItemName(items []SelectItem, name string, fallback int) int {
if len(items) == 0 {
return 0
}
if name != "" {
for i, item := range items {
if item.Name == name {
return i
}
if i := indexOfItemName(items, name); i >= 0 {
return i
}
}
if fallback < 0 {
@@ -167,10 +193,19 @@ func (m selectorModel) filteredItems() []SelectItem {
filterLower := strings.ToLower(m.filter)
var result []SelectItem
for _, item := range m.items {
if m.rankFiltered {
if selectItemMatchScore(item, filterLower).ok {
result = append(result, item)
}
continue
}
if strings.Contains(strings.ToLower(item.Name), filterLower) {
result = append(result, item)
}
}
if m.rankFiltered {
sortSelectItemsForFilter(result, filterLower)
}
return result
}
@@ -241,6 +276,54 @@ func (m *selectorModel) updateNavigation(msg tea.KeyMsg) {
}
}
func (m *selectorModel) UpdateNavigation(msg tea.KeyMsg) {
m.updateNavigation(msg)
}
func (m *selectorModel) Move(delta int) {
if delta == 0 {
return
}
filtered := m.filteredItems()
if len(filtered) == 0 {
m.cursor = 0
m.scrollOffset = 0
return
}
m.cursor += delta
if m.cursor < 0 {
m.cursor = 0
}
if m.cursor >= len(filtered) {
m.cursor = len(filtered) - 1
}
m.updateScroll(m.otherStart())
}
func (m *selectorModel) SetHelpText(help string) {
m.helpText = help
}
func (m selectorModel) Filter() string {
return m.filter
}
func (m selectorModel) FilteredItems() []SelectItem {
return append([]SelectItem(nil), m.filteredItems()...)
}
func (m selectorModel) SelectedItem() (SelectItem, bool) {
filtered := m.filteredItems()
if len(filtered) == 0 || m.cursor < 0 || m.cursor >= len(filtered) {
return SelectItem{}, false
}
return filtered[m.cursor], true
}
func (m selectorModel) RenderContent() string {
return m.renderContent()
}
// updateScroll adjusts scrollOffset based on cursor position.
// When not filtering, scrollOffset is relative to the "More" (non-recommended) section.
// When filtering, it's relative to the full filtered list.
@@ -281,7 +364,7 @@ func (m selectorModel) Update(msg tea.Msg) (tea.Model, tea.Cmd) {
wasSet := m.width > 0
m.width = msg.Width
if wasSet {
return m, tea.EnterAltScreen
return m, tea.ClearScreen
}
return m, nil
@@ -338,6 +421,16 @@ func (m selectorModel) renderItem(s *strings.Builder, item SelectItem, idx int)
}
}
func (m selectorModel) renderCompactItem(s *strings.Builder, item SelectItem, idx int) {
if idx == m.cursor {
s.WriteString(selectorSelectedItemStyle.Render("▸ " + item.Name))
s.WriteString(cursorItemSuffix(item))
} else {
s.WriteString(selectorItemStyle.Render(item.Name))
}
s.WriteString("\n")
}
// renderContent renders the selector content (title, items, help text) without
// checking the cancelled/selected state. This is used by both View() (standalone mode)
// and by the TUI modal which embeds a selectorModel.
@@ -431,6 +524,57 @@ func (m selectorModel) renderContent() string {
return s.String()
}
func (m selectorModel) RenderCompactContent(maxItems int) string {
var s strings.Builder
s.WriteString(selectorTitleStyle.Render(m.title))
s.WriteString(" ")
if m.filter == "" {
s.WriteString(selectorFilterStyle.Render("Type to filter..."))
} else {
s.WriteString(selectorInputStyle.Render(m.filter))
}
s.WriteString("\n")
filtered := m.filteredItems()
if len(filtered) == 0 {
s.WriteString(selectorItemStyle.Render(selectorDescStyle.Render("(no matches)")))
s.WriteString("\n")
} else {
maxItems = max(1, maxItems)
start := 0
if len(filtered) > maxItems {
start = m.cursor - maxItems/2
if start < 0 {
start = 0
}
if maxStart := len(filtered) - maxItems; start > maxStart {
start = maxStart
}
}
end := min(len(filtered), start+maxItems)
if start > 0 {
s.WriteString(selectorMoreStyle.Render(fmt.Sprintf("... %d more above", start)))
s.WriteString("\n")
}
for idx := start; idx < end; idx++ {
m.renderCompactItem(&s, filtered[idx], idx)
}
if remaining := len(filtered) - end; remaining > 0 {
s.WriteString(selectorMoreStyle.Render(fmt.Sprintf("... and %d more", remaining)))
s.WriteString("\n")
}
}
help := "↑/↓ navigate • enter select • esc cancel"
if m.helpText != "" {
help = m.helpText
}
s.WriteString(selectorHelpStyle.Render(help))
return s.String()
}
func (m selectorModel) View() string {
if m.cancelled || m.selected != "" {
return ""
@@ -443,6 +587,99 @@ func (m selectorModel) View() string {
return s
}
type selectItemScore struct {
ok bool
rank int
index int
lengthDelta int
recommended int
name string
}
func sortSelectItemsForFilter(items []SelectItem, filter string) {
filter = strings.ToLower(strings.TrimSpace(filter))
sort.SliceStable(items, func(i, j int) bool {
return compareSelectItemsForFilter(items[i], items[j], filter) < 0
})
}
func compareSelectItemsForFilter(a, b SelectItem, filter string) int {
aScore := selectItemMatchScore(a, filter)
bScore := selectItemMatchScore(b, filter)
for _, cmp := range []int{
compareSelectorInt(aScore.rank, bScore.rank),
compareSelectorInt(aScore.index, bScore.index),
compareSelectorInt(aScore.lengthDelta, bScore.lengthDelta),
compareSelectorInt(aScore.recommended, bScore.recommended),
strings.Compare(aScore.name, bScore.name),
} {
if cmp != 0 {
return cmp
}
}
return 0
}
func selectItemMatchScore(item SelectItem, filter string) selectItemScore {
filter = strings.ToLower(strings.TrimSpace(filter))
name := strings.ToLower(strings.TrimSpace(item.Name))
description := strings.ToLower(strings.TrimSpace(item.Description))
score := selectItemScore{
rank: 4,
index: 1 << 20,
lengthDelta: 1 << 20,
name: name,
}
if item.Recommended {
score.recommended = -1
}
if filter == "" {
score.ok = true
return score
}
nameRunes := len([]rune(name))
filterRunes := len([]rune(filter))
if name == filter {
score.ok = true
score.rank = 0
score.index = 0
score.lengthDelta = 0
return score
}
if strings.HasPrefix(name, filter) {
score.ok = true
score.rank = 1
score.index = 0
score.lengthDelta = max(0, nameRunes-filterRunes)
return score
}
if index := strings.Index(name, filter); index >= 0 {
score.ok = true
score.rank = 2
score.index = len([]rune(name[:index]))
score.lengthDelta = max(0, nameRunes-filterRunes)
return score
}
if index := strings.Index(description, filter); index >= 0 {
score.ok = true
score.rank = 3
score.index = len([]rune(description[:index]))
score.lengthDelta = max(0, nameRunes-filterRunes)
}
return score
}
func compareSelectorInt(a, b int) int {
switch {
case a < b:
return -1
case a > b:
return 1
default:
return 0
}
}
// cursorForCurrent returns the item index matching current, or 0 if not found.
func cursorForCurrent(items []SelectItem, current string) int {
if current == "" {
@@ -451,10 +688,8 @@ func cursorForCurrent(items []SelectItem, current string) int {
// Prefer exact name matches before tag-prefix fallback so "qwen3.5" does not
// incorrectly select "qwen3.5:cloud" (and vice versa) based on list order.
for i, item := range items {
if item.Name == current {
return i
}
if i := indexOfItemName(items, current); i >= 0 {
return i
}
for i, item := range items {
@@ -700,7 +935,7 @@ func (m multiSelectorModel) Update(msg tea.Msg) (tea.Model, tea.Cmd) {
wasSet := m.width > 0
m.width = msg.Width
if wasSet {
return m, tea.EnterAltScreen
return m, tea.ClearScreen
}
return m, nil
+2 -2
View File
@@ -60,7 +60,7 @@ func (m signInModel) Update(msg tea.Msg) (tea.Model, tea.Cmd) {
wasSet := m.width > 0
m.width = msg.Width
if wasSet {
return m, tea.EnterAltScreen
return m, tea.ClearScreen
}
return m, nil
@@ -115,7 +115,7 @@ func (m upgradeModel) Update(msg tea.Msg) (tea.Model, tea.Cmd) {
wasSet := m.width > 0
m.width = msg.Width
if wasSet {
return m, tea.EnterAltScreen
return m, tea.ClearScreen
}
return m, nil
+110
View File
@@ -0,0 +1,110 @@
package tui
import (
"os"
"sync"
"time"
tea "github.com/charmbracelet/bubbletea"
"github.com/charmbracelet/lipgloss"
"github.com/ollama/ollama/cmd/launch"
"golang.org/x/term"
)
// spinnerStyle dims the spinner so it reads as ancillary status text, matching
// the sign-in/upgrade spinners in signin.go.
var spinnerStyle = lipgloss.NewStyle().
Foreground(lipgloss.AdaptiveColor{Light: "242", Dark: "246"})
type spinnerTickMsg struct{}
// spinnerQuitMsg is sent by Stop to ask the program to quit cleanly.
type spinnerQuitMsg struct{}
type spinnerModel struct {
message string
frame int
quitting bool
cancelled chan struct{}
once sync.Once
}
func (m *spinnerModel) Init() tea.Cmd {
return tea.Tick(100*time.Millisecond, func(time.Time) tea.Msg { return spinnerTickMsg{} })
}
func (m *spinnerModel) Update(msg tea.Msg) (tea.Model, tea.Cmd) {
switch msg := msg.(type) {
case spinnerTickMsg:
if m.quitting {
return m, nil
}
m.frame++
return m, tea.Tick(100*time.Millisecond, func(time.Time) tea.Msg { return spinnerTickMsg{} })
case tea.KeyMsg:
// bubbletea runs the terminal in raw mode, so Ctrl+C is delivered here
// as a key rather than as a SIGINT. Treat it as a user cancellation:
// close the cancelled channel (so the caller's wait loop can abort) and
// quit the program so bubbletea restores the terminal before control
// returns to the caller.
if msg.String() == "ctrl+c" {
m.once.Do(func() { close(m.cancelled) })
m.quitting = true
return m, tea.Quit
}
case spinnerQuitMsg:
m.quitting = true
// Returning "" from View on quit clears the spinner line, mirroring how
// confirm.go blanks its view when it quits.
return m, tea.Quit
}
return m, nil
}
func (m *spinnerModel) View() string {
if m.quitting {
return ""
}
frame := launch.SpinnerFrames[m.frame%len(launch.SpinnerFrames)]
return spinnerStyle.Render(frame + " " + m.message)
}
// RunSpinner runs a bubbletea spinner displaying message until the returned
// Spinner's Stop is called. Stop signals the program to quit and blocks until
// it has exited and cleared its line. If the user presses Ctrl+C while the
// spinner is running, Spinner.Cancelled() is closed so the caller can abort
// its wait; the program quits and the terminal is restored before Stop
// returns. RunSpinner returns nil when there is no interactive terminal, so
// launch.StartSpinner can fall back to its ANSI spinner for headless/--yes
// runs.
func RunSpinner(message string) *launch.Spinner {
if !term.IsTerminal(int(os.Stdin.Fd())) || !term.IsTerminal(int(os.Stderr.Fd())) {
return nil
}
cancelled := make(chan struct{})
m := &spinnerModel{message: message, cancelled: cancelled}
p := tea.NewProgram(m, tea.WithOutput(os.Stderr))
done := make(chan struct{})
go func() {
_, _ = p.Run()
close(done)
}()
var once sync.Once
stop := func() {
once.Do(func() {
select {
case <-done:
// Program already finished (e.g. the user cancelled), so don't
// send to it; just ensure it has exited.
return
default:
}
p.Send(spinnerQuitMsg{})
<-done
})
}
return launch.NewSpinner(stop, cancelled)
}
+22 -82
View File
@@ -42,68 +42,41 @@ type menuItem struct {
description string
integration string
isRunModel bool
isOthers bool
}
const pinnedIntegrationCount = 4
var runModelMenuItem = menuItem{
title: "Chat with a model",
description: "Start an interactive chat with a model",
title: "Chat, Code, & Work",
description: "Chat with models, code, search the web, and delegate real work",
isRunModel: true,
}
var othersMenuItem = menuItem{
title: "More...",
description: "Show additional integrations",
isOthers: true,
}
// launcherMenuIntegrations is intentionally short: the root ollama command is
// a quick path to the most common launch targets. Other registered
// integrations remain available through `ollama launch <integration>`.
var launcherMenuIntegrations = []string{"claude", "opencode", "hermes", "openclaw"}
type model struct {
state *launch.LauncherState
items []menuItem
cursor int
showOthers bool
width int
quitting bool
selected bool
action TUIAction
state *launch.LauncherState
items []menuItem
cursor int
width int
quitting bool
selected bool
action TUIAction
}
func newModel(state *launch.LauncherState) model {
m := model{
state: state,
}
m.showOthers = shouldExpandOthers(state)
m.items = buildMenuItems(state, m.showOthers)
m.items = buildMenuItems(state)
m.cursor = initialCursor(state, m.items)
return m
}
func shouldExpandOthers(state *launch.LauncherState) bool {
if state == nil {
return false
}
for _, item := range otherIntegrationItems(state) {
if item.integration == state.LastSelection {
return true
}
}
return false
}
func buildMenuItems(state *launch.LauncherState, showOthers bool) []menuItem {
func buildMenuItems(state *launch.LauncherState) []menuItem {
items := []menuItem{runModelMenuItem}
items = append(items, pinnedIntegrationItems(state)...)
otherItems := otherIntegrationItems(state)
switch {
case showOthers:
items = append(items, otherItems...)
case len(otherItems) > 0:
items = append(items, othersMenuItem)
}
items = append(items, launcherIntegrationItems(state)...)
return items
}
@@ -119,30 +92,14 @@ func integrationMenuItem(state launch.LauncherIntegrationState) menuItem {
}
}
func otherIntegrationItems(state *launch.LauncherState) []menuItem {
ordered := orderedIntegrationItems(state)
if len(ordered) <= pinnedIntegrationCount {
return nil
}
return ordered[pinnedIntegrationCount:]
}
func pinnedIntegrationItems(state *launch.LauncherState) []menuItem {
ordered := orderedIntegrationItems(state)
if len(ordered) <= pinnedIntegrationCount {
return ordered
}
return ordered[:pinnedIntegrationCount]
}
func orderedIntegrationItems(state *launch.LauncherState) []menuItem {
func launcherIntegrationItems(state *launch.LauncherState) []menuItem {
if state == nil {
return nil
}
items := make([]menuItem, 0, len(state.Integrations))
for _, info := range launch.ListIntegrationInfos() {
integrationState, ok := state.Integrations[info.Name]
items := make([]menuItem, 0, len(launcherMenuIntegrations))
for _, name := range launcherMenuIntegrations {
integrationState, ok := state.Integrations[name]
if !ok {
continue
}
@@ -151,10 +108,6 @@ func orderedIntegrationItems(state *launch.LauncherState) []menuItem {
return items
}
func primaryMenuItemCount(state *launch.LauncherState) int {
return 1 + len(pinnedIntegrationItems(state))
}
func initialCursor(state *launch.LauncherState, items []menuItem) int {
if state == nil || state.LastSelection == "" {
return 0
@@ -190,21 +143,12 @@ func (m model) Update(msg tea.Msg) (tea.Model, tea.Cmd) {
if m.cursor > 0 {
m.cursor--
}
if m.showOthers && m.cursor < primaryMenuItemCount(m.state) {
m.showOthers = false
m.items = buildMenuItems(m.state, false)
m.cursor = min(m.cursor, len(m.items)-1)
}
return m, nil
case "down", "j":
if m.cursor < len(m.items)-1 {
m.cursor++
}
if m.cursor < len(m.items) && m.items[m.cursor].isOthers && !m.showOthers {
m.showOthers = true
m.items = buildMenuItems(m.state, true)
}
return m, nil
case "enter", " ":
@@ -235,7 +179,7 @@ func (m model) selectableItem(item menuItem) bool {
if item.isRunModel {
return true
}
if item.integration == "" || item.isOthers {
if item.integration == "" {
return false
}
state, ok := m.state.Integrations[item.integration]
@@ -243,7 +187,7 @@ func (m model) selectableItem(item menuItem) bool {
}
func (m model) changeableItem(item menuItem) bool {
if item.integration == "" || item.isOthers {
if item.integration == "" {
return false
}
state, ok := m.state.Integrations[item.integration]
@@ -287,10 +231,6 @@ func (m model) renderMenuItem(index int, item menuItem) string {
if m.cursor == index {
style = menuSelectedItemStyle
}
} else if item.isOthers {
if m.cursor == index {
style = menuSelectedItemStyle
}
} else {
integrationState := m.state.Integrations[item.integration]
if !integrationState.Selectable {
+30 -82
View File
@@ -29,10 +29,10 @@ func launcherTestState() *launch.LauncherState {
Selectable: true,
Changeable: true,
},
"codex-app": {
Name: "codex-app",
DisplayName: "Codex App",
Description: "An AI agent you can delegate real work to, by OpenAI",
"chatgpt": {
Name: "chatgpt",
DisplayName: "ChatGPT",
Description: "Complete work with ChatGPT",
Selectable: true,
Changeable: true,
},
@@ -91,8 +91,6 @@ func integrationSequence(items []menuItem) []string {
switch {
case item.isRunModel:
sequence = append(sequence, "run")
case item.isOthers:
sequence = append(sequence, "more")
case item.integration != "":
sequence = append(sequence, item.integration)
}
@@ -104,81 +102,31 @@ func compareStrings(got, want []string) string {
return cmp.Diff(want, got)
}
func expectedCollapsedSequence(state *launch.LauncherState) []string {
sequence := []string{"run"}
for _, item := range pinnedIntegrationItems(state) {
sequence = append(sequence, item.integration)
}
if len(otherIntegrationItems(state)) > 0 {
sequence = append(sequence, "more")
}
return sequence
}
func expectedExpandedSequence(state *launch.LauncherState) []string {
sequence := []string{"run"}
for _, item := range pinnedIntegrationItems(state) {
sequence = append(sequence, item.integration)
}
for _, item := range otherIntegrationItems(state) {
sequence = append(sequence, item.integration)
}
return sequence
}
func TestMenuRendersPinnedItemsAndMore(t *testing.T) {
func TestMenuRendersRootLaunchChoices(t *testing.T) {
state := launcherTestState()
menu := newModel(state)
wantPrefix := []string{"run", "claude", "codex-app", "hermes", "openclaw"}
if findMenuCursorByIntegration(menu.items, "codex-app") == -1 {
wantPrefix = []string{"run", "claude", "hermes", "openclaw", "opencode"}
}
if got := integrationSequence(menu.items); len(got) < len(wantPrefix) {
t.Fatalf("expected at least %d menu items, got %v", len(wantPrefix), got)
} else if diff := compareStrings(got[:len(wantPrefix)], wantPrefix); diff != "" {
t.Fatalf("unexpected primary TUI order: %s", diff)
want := []string{"run", "claude", "opencode", "hermes", "openclaw"}
if diff := compareStrings(integrationSequence(menu.items), want); diff != "" {
t.Fatalf("unexpected root launch choices: %s", diff)
}
view := menu.View()
for _, want := range []string{"Chat with a model", "Launch Claude Code", "Launch Hermes Agent", "Launch OpenClaw", "More..."} {
for _, want := range []string{
"Chat, Code, & Work",
"Chat with models, code, search the web, and delegate real work",
"Launch Claude Code",
"Launch OpenCode",
"Launch Hermes Agent",
"Launch OpenClaw",
} {
if !strings.Contains(view, want) {
t.Fatalf("expected menu view to contain %q\n%s", want, view)
}
}
if findMenuCursorByIntegration(menu.items, "codex-app") != -1 && !strings.Contains(view, "Launch Codex App") {
t.Fatalf("expected menu view to contain Codex App\n%s", view)
}
if strings.Contains(view, "Launch Claude Desktop") {
t.Fatalf("expected hidden Claude Desktop to be absent\n%s", view)
}
wantOrder := expectedCollapsedSequence(state)
if diff := compareStrings(integrationSequence(menu.items), wantOrder); diff != "" {
t.Fatalf("unexpected pinned order: %s", diff)
}
}
func TestMenuExpandsOthersFromLastSelection(t *testing.T) {
state := launcherTestState()
overflow := otherIntegrationItems(state)
if len(overflow) == 0 {
t.Fatal("expected at least one overflow integration")
}
state.LastSelection = overflow[0].integration
menu := newModel(state)
if !menu.showOthers {
t.Fatal("expected others section to expand when last selection is in the overflow list")
}
view := menu.View()
if !strings.Contains(view, overflow[0].title) {
t.Fatalf("expected expanded view to contain overflow integration\n%s", view)
}
if strings.Contains(view, "More...") {
t.Fatalf("expected expanded view to replace More... item\n%s", view)
}
wantOrder := expectedExpandedSequence(state)
if diff := compareStrings(integrationSequence(menu.items), wantOrder); diff != "" {
t.Fatalf("unexpected expanded order: %s", diff)
for _, hidden := range []string{"Launch ChatGPT", "Launch Codex", "Launch Droid", "Launch Pi", "More..."} {
if strings.Contains(view, hidden) {
t.Fatalf("expected root menu to omit %q\n%s", hidden, view)
}
}
}
@@ -273,24 +221,24 @@ func TestMenuShowsCurrentModelSuffixes(t *testing.T) {
func TestMenuShowsInstallStatusAndHint(t *testing.T) {
state := launcherTestState()
codex := state.Integrations["codex"]
codex.Installed = false
codex.Selectable = false
codex.Changeable = false
codex.InstallHint = "Install from https://example.com/codex"
state.Integrations["codex"] = codex
opencode := state.Integrations["opencode"]
opencode.Installed = false
opencode.Selectable = false
opencode.Changeable = false
opencode.InstallHint = "Install from https://example.com/opencode"
state.Integrations["opencode"] = opencode
state.LastSelection = "codex"
state.LastSelection = "opencode"
menu := newModel(state)
menu.cursor = findMenuCursorByIntegration(menu.items, "codex")
menu.cursor = findMenuCursorByIntegration(menu.items, "opencode")
if menu.cursor == -1 {
t.Fatal("expected codex menu item in overflow section")
t.Fatal("expected opencode menu item")
}
view := menu.View()
if !strings.Contains(view, "(not installed)") {
t.Fatalf("expected not-installed marker\n%s", view)
}
if !strings.Contains(view, codex.InstallHint) {
if !strings.Contains(view, opencode.InstallHint) {
t.Fatalf("expected install hint in description\n%s", view)
}
}
+8 -5
View File
@@ -39,10 +39,6 @@ func (q *qwen25VLModel) KV(t *Tokenizer) KV {
}
}
if q.VisionModel.FullAttentionBlocks == nil {
kv["qwen25vl.vision.fullatt_block_indexes"] = []int32{7, 15, 23, 31}
}
kv["qwen25vl.vision.block_count"] = cmp.Or(q.VisionModel.Depth, 32)
kv["qwen25vl.vision.embedding_length"] = q.VisionModel.HiddenSize
kv["qwen25vl.vision.attention.head_count"] = cmp.Or(q.VisionModel.NumHeads, 16)
@@ -53,12 +49,19 @@ func (q *qwen25VLModel) KV(t *Tokenizer) KV {
kv["qwen25vl.vision.window_size"] = cmp.Or(q.VisionModel.WindowSize, 112)
kv["qwen25vl.vision.attention.layer_norm_epsilon"] = cmp.Or(q.VisionModel.RMSNormEps, 1e-6)
kv["qwen25vl.vision.rope.freq_base"] = cmp.Or(q.VisionModel.RopeTheta, 1e4)
kv["qwen25vl.vision.fullatt_block_indexes"] = q.VisionModel.FullAttentionBlocks
kv["qwen25vl.vision.fullatt_block_indexes"] = q.fullAttentionBlocks()
kv["qwen25vl.vision.temporal_patch_size"] = cmp.Or(q.VisionModel.TemporalPatchSize, 2)
return kv
}
func (q *qwen25VLModel) fullAttentionBlocks() []int32 {
if len(q.VisionModel.FullAttentionBlocks) > 0 {
return q.VisionModel.FullAttentionBlocks
}
return []int32{7, 15, 23, 31}
}
func (q *qwen25VLModel) Tensors(ts []Tensor) []*ggml.Tensor {
var out []*ggml.Tensor
+44
View File
@@ -0,0 +1,44 @@
package convert
import (
"slices"
"testing"
)
func TestQwen25VLFullAttentionBlockDefaults(t *testing.T) {
tokenizer := &Tokenizer{Vocabulary: &Vocabulary{}}
for _, tt := range []struct {
name string
blocks []int32
want []int32
}{
{
name: "nil",
want: []int32{7, 15, 23, 31},
},
{
name: "empty",
blocks: []int32{},
want: []int32{7, 15, 23, 31},
},
{
name: "custom",
blocks: []int32{5, 17},
want: []int32{5, 17},
},
} {
t.Run(tt.name, func(t *testing.T) {
model := &qwen25VLModel{}
model.VisionModel.FullAttentionBlocks = tt.blocks
got, ok := model.KV(tokenizer)["qwen25vl.vision.fullatt_block_indexes"].([]int32)
if !ok {
t.Fatalf("fullatt_block_indexes has unexpected type %T", got)
}
if !slices.Equal(got, tt.want) {
t.Fatalf("fullatt_block_indexes = %v, want %v", got, tt.want)
}
})
}
}
+48 -7
View File
@@ -7,34 +7,66 @@ import (
"github.com/ollama/ollama/ml"
)
const (
cudaV12RuntimeMajor = 12
minFatbinCompressionCUDARuntimeMinor = 4
minFatbinCompressionNVIDIADriverMajor = 550
minLegacyComputeJITCUDARuntimeMinor = 8
// Older CUDA compute targets need newer drivers when they are JITed from PTX.
minLegacyComputeJITNVIDIADriverMajor = 570
)
func filterOldCUDADriver(_ context.Context, devices []ml.DeviceInfo) []ml.DeviceInfo {
oldCUDA := func(dev ml.DeviceInfo) bool {
return dev.Library == "CUDA" && dev.ComputeMajor > 0 && dev.ComputeMajor < 7
}
needsCheck := false
hasCUDA := false
for _, dev := range devices {
if oldCUDA(dev) {
needsCheck = true
if dev.Library == "CUDA" {
hasCUDA = true
break
}
}
if !needsCheck {
if !hasCUDA {
return devices
}
driver := nvidiaDriverMajorFromDevices(devices)
if driver == 0 {
slog.Warn("could not verify NVIDIA driver compatibility for an older NVIDIA GPU")
slog.Warn("could not verify NVIDIA driver compatibility for CUDA")
return devices
}
if driver >= 570 {
// Match the driver floor to the CUDA runtime we are about to load, so source
// builds with older CUDA runtimes can still run on matching older drivers.
runtimeMajor, runtimeMinor, hasRuntime := cudaRuntimeVersionFromDevices(devices)
runtimeMayUseCompressedFatbins := hasRuntime &&
runtimeMajor == cudaV12RuntimeMajor &&
runtimeMinor >= minFatbinCompressionCUDARuntimeMinor
// CUDA v12.8+ source builds are expected to either use Ollama's PTX packaging
// for older compute targets or be built against a matching local driver/toolkit.
runtimeMayJITLegacyCompute := hasRuntime &&
runtimeMajor == cudaV12RuntimeMajor &&
runtimeMinor >= minLegacyComputeJITCUDARuntimeMinor
if driver >= minLegacyComputeJITNVIDIADriverMajor || (!runtimeMayUseCompressedFatbins && !runtimeMayJITLegacyCompute) {
return devices
}
filtered := devices[:0]
for _, dev := range devices {
if oldCUDA(dev) {
if dev.Library != "CUDA" {
filtered = append(filtered, dev)
continue
}
if runtimeMayUseCompressedFatbins && driver < minFatbinCompressionNVIDIADriverMajor {
slog.Warn("NVIDIA driver too old",
"device", dev.Description, "compute", dev.Compute(), "driver", driver, "required_driver", "550 or newer")
continue
}
if runtimeMayJITLegacyCompute && oldCUDA(dev) {
slog.Warn("NVIDIA driver too old",
"device", dev.Description, "compute", dev.Compute(), "driver", driver, "required_driver", "570 or newer")
continue
@@ -52,3 +84,12 @@ func nvidiaDriverMajorFromDevices(devices []ml.DeviceInfo) int {
}
return 0
}
func cudaRuntimeVersionFromDevices(devices []ml.DeviceInfo) (int, int, bool) {
for _, dev := range devices {
if dev.Library == "CUDA" {
return cudaRuntimeVersion(dev.LibraryPath)
}
}
return 0, 0, false
}
Loaded 100 of 364 files, more files were not shown because too many files have changed in this diff. Show more