Compare commits

...

8 Commits

Author SHA1 Message Date
Parth Sareen
fce745fe5e agent: import skills from coding agents (#17294) 2026-07-22 22:40:25 -07:00
Daniel Hiltgen
1fd1ccf7ad model: align Laguna with upstream llama.cpp (#17335)
Update llama.cpp to pick up upstream Laguna implementation and remove Ollama's local Laguna implementation. Retain a narrow Metal-only scaling workaround for routed-MoE prompt overflow.

Translate older Ollama GGUF attention-gate and SWA metadata names so existing models continue to load.
2026-07-22 17:09:18 -07:00
Michael Yang
efb7e3c55e docs: update retirements (#17289) 2026-07-22 14:24:11 -07:00
Daniel Hiltgen
b517b9bd01 model/parsers: finalize incomplete GLM tool calls (#17250)
The GLM parser buffered tool calls until it observed </tool_call>, but ignored the terminal done signal. If the model omitted or partially emitted the outer closing tag, Ollama returned a successful empty response instead of a tool call or an actionable error, leaving coding agents unable to continue.

On end-of-stream, finalize only structurally complete calls for declared tools with all required arguments. Complete calls missing only the outer delimiter now proceed through the existing parser, while genuinely truncated calls return an explicit error rather than being silently dropped.

Fixes #16497
2026-07-22 13:53:47 -07:00
Daniel Hiltgen
479664e7aa mlx update (#17332) 2026-07-22 13:36:49 -07:00
Daniel Hiltgen
a51df81573 test: revamp integration test entrpoints (#16560)
This refactors the existing integration tests into 3 priumary groups: fast,
release, and library.  It also refines some of the release tests to drop some
of the older models and pick up newer models, while retaining the broad
coverage in the library group.
2026-07-21 16:06:38 -07:00
Daniel Hiltgen
a18c230189 model: add Laguna v8 chat support and fix Metal inference (#17291)
Add a laguna-v8 renderer/parser matching the Laguna XS 2.1 template, and fix v2 handling of embedded thinking and structured tool arguments.

Prevent FP16 overflow in Metal's quantized routed-MoE prefill path by scaling the linear branch and folding the inverse into the routing scale. Other backends and token-generation paths are unchanged.

Add comprehensive v2/v8 Jinja parity and parser tests.
2026-07-21 16:06:29 -07:00
Daniel Hiltgen
e21d5327b0 CI: fix missing CUDA v13.4 sub-package (#17288)
Needed for cross-compiling WoA
2026-07-21 12:25:10 -07:00
58 changed files with 3905 additions and 1780 deletions

View File

@@ -130,6 +130,7 @@ jobs:
build-steps: cuda13Arm64Cross
install: https://packages.nvidia.com/prerelease/cuda/13.4.0/local_installers/cuda_13.4.0_windows_x86_64.exe
cuda-components:
- '"cudart"'
- '"cudart_cross"'
- '"nvcc"'
- '"nvcc_cross"'

View File

@@ -327,6 +327,7 @@ jobs:
build-steps: cuda13Arm64Cross
install: https://packages.nvidia.com/prerelease/cuda/13.4.0/local_installers/cuda_13.4.0_windows_x86_64.exe
cuda-components:
- '"cudart"'
- '"cudart_cross"'
- '"nvcc"'
- '"nvcc_cross"'

View File

@@ -1 +1 @@
b10069
b10091

View File

@@ -1 +1 @@
b7c3dd6d27f45b5365b08a840310187dc503f1db
33c03c486c34a7dadab5339563612c9933c4a406

View File

@@ -1,8 +1,10 @@
package agent
import (
"bytes"
"errors"
"fmt"
"io"
"io/fs"
"os"
"path/filepath"
@@ -271,6 +273,354 @@ type skillRoot struct {
path string
}
// SkillImportResult describes one import attempt. Failed skills do not prevent
// other valid skills in the same source root from being imported.
type SkillImportResult struct {
Source string
SourceDir string
Destination string
Imported []string
Existing []string
Failures []SkillImportFailure
}
// SkillImportFailure identifies a source skill that was deliberately skipped.
// The destination is never changed for a failed skill.
type SkillImportFailure struct {
Name string
Err error
}
// ImportSkills imports skills from a conventional coding-agent source into the
// canonical Ollama skills directory. Supported sources are codex, claude, and
// pi. Existing skills are left untouched: an identical directory is reported
// as existing, and a differing one is reported as a conflict.
func ImportSkills(source string) (SkillImportResult, error) {
home, err := os.UserHomeDir()
if err != nil {
return SkillImportResult{}, fmt.Errorf("resolve home directory: %w", err)
}
destination, err := SkillsDir()
if err != nil {
return SkillImportResult{}, fmt.Errorf("resolve Ollama skills directory: %w", err)
}
return importSkillsFromRoots(source, conventionalSkillImportRoots(home), destination)
}
func conventionalSkillImportRoots(home string) map[string]string {
return map[string]string{
"codex": filepath.Join(home, ".codex", "skills"),
"claude": filepath.Join(home, ".claude", "skills"),
"pi": filepath.Join(home, ".pi", "agent", "skills"),
}
}
func importSkillsFromRoots(source string, roots map[string]string, destination string) (SkillImportResult, error) {
source = strings.ToLower(strings.TrimSpace(source))
sourceDir, ok := roots[source]
if !ok {
return SkillImportResult{}, fmt.Errorf("unknown skill source %q", source)
}
return importSkillsFromDir(source, sourceDir, destination)
}
func importSkillsFromDir(source, sourceDir, destination string) (SkillImportResult, error) {
result := SkillImportResult{Source: source, SourceDir: sourceDir, Destination: destination}
info, err := os.Lstat(sourceDir)
if errors.Is(err, fs.ErrNotExist) {
return result, nil
}
if err != nil {
return result, fmt.Errorf("inspect %s skills directory: %w", source, err)
}
if info.Mode()&os.ModeSymlink != 0 {
return result, fmt.Errorf("inspect %s skills directory: symlinks are not supported", source)
}
if !info.IsDir() {
return result, fmt.Errorf("inspect %s skills directory: not a directory", source)
}
entries, err := os.ReadDir(sourceDir)
if err != nil {
return result, fmt.Errorf("read %s skills directory: %w", source, err)
}
for _, entry := range entries {
name := entry.Name()
path := filepath.Join(sourceDir, name)
if entry.Type()&os.ModeSymlink != 0 {
result.Failures = append(result.Failures, SkillImportFailure{Name: name, Err: errors.New("symlinked skill directories are not supported")})
continue
}
info, err := entry.Info()
if err != nil {
result.Failures = append(result.Failures, SkillImportFailure{Name: name, Err: fmt.Errorf("inspect source: %w", err)})
continue
}
if !info.IsDir() {
continue
}
if !skillName.MatchString(name) {
result.Failures = append(result.Failures, SkillImportFailure{Name: name, Err: errors.New("invalid skill directory name")})
continue
}
if err := validateImportSkill(path, name); err != nil {
result.Failures = append(result.Failures, SkillImportFailure{Name: name, Err: err})
continue
}
state, err := importSkillDirectory(path, filepath.Join(destination, name))
if err != nil {
result.Failures = append(result.Failures, SkillImportFailure{Name: name, Err: err})
continue
}
if state == skillImportExisting {
result.Existing = append(result.Existing, name)
} else {
result.Imported = append(result.Imported, name)
}
}
return result, nil
}
func validateImportSkill(dir, name string) error {
manifest := filepath.Join(dir, skillFilename)
info, err := os.Lstat(manifest)
if err != nil {
return fmt.Errorf("inspect %s: %w", skillFilename, err)
}
if info.Mode()&os.ModeSymlink != 0 || !info.Mode().IsRegular() {
return fmt.Errorf("%s must be a regular, non-symlinked file", skillFilename)
}
if _, err := parseSkill(manifest, name); err != nil {
return err
}
return walkImportTree(dir, func(path string, entry fs.DirEntry, info fs.FileInfo) error {
if info.IsDir() || path == dir {
return nil
}
if !info.Mode().IsRegular() {
return fmt.Errorf("only regular files may be imported: %s", path)
}
file, err := os.Open(path)
if err != nil {
return fmt.Errorf("read %s: %w", path, err)
}
return file.Close()
})
}
func walkImportTree(root string, visit func(string, fs.DirEntry, fs.FileInfo) error) error {
return filepath.WalkDir(root, func(path string, entry fs.DirEntry, err error) error {
if err != nil {
return err
}
rel, err := filepath.Rel(root, path)
if err != nil || rel == ".." || strings.HasPrefix(rel, ".."+string(filepath.Separator)) {
return fmt.Errorf("unsafe skill path %q", path)
}
if entry.Type()&os.ModeSymlink != 0 {
return fmt.Errorf("symlinks may not be imported: %s", path)
}
info, err := entry.Info()
if err != nil {
return err
}
return visit(path, entry, info)
})
}
type skillImportState int
const (
skillImportCopied skillImportState = iota
skillImportExisting
)
func importSkillDirectory(source, destination string) (skillImportState, error) {
if info, err := os.Lstat(destination); err == nil {
if info.Mode()&os.ModeSymlink != 0 || !info.IsDir() {
return 0, errors.New("destination exists but is not a regular directory")
}
same, err := sameImportTree(source, destination)
if err != nil {
return 0, fmt.Errorf("inspect existing destination: %w", err)
}
if same {
return skillImportExisting, nil
}
return 0, errors.New("destination skill already exists with different contents")
} else if !errors.Is(err, fs.ErrNotExist) {
return 0, fmt.Errorf("inspect destination: %w", err)
}
if err := ensureImportDestination(filepath.Dir(destination)); err != nil {
return 0, err
}
stage, err := os.MkdirTemp(filepath.Dir(destination), "."+filepath.Base(destination)+".import-")
if err != nil {
return 0, fmt.Errorf("create import staging directory: %w", err)
}
defer os.RemoveAll(stage)
if err := copyImportTree(source, stage); err != nil {
return 0, err
}
if _, err := os.Lstat(destination); err == nil {
return 0, errors.New("destination skill was created during import")
} else if !errors.Is(err, fs.ErrNotExist) {
return 0, fmt.Errorf("inspect destination before install: %w", err)
}
if err := os.Rename(stage, destination); err != nil {
return 0, fmt.Errorf("install imported skill: %w", err)
}
return skillImportCopied, nil
}
func ensureImportDestination(dir string) error {
if err := os.MkdirAll(dir, 0o755); err != nil {
return fmt.Errorf("create Ollama skills directory: %w", err)
}
info, err := os.Lstat(dir)
if err != nil {
return fmt.Errorf("inspect Ollama skills directory: %w", err)
}
if info.Mode()&os.ModeSymlink != 0 || !info.IsDir() {
return errors.New("Ollama skills directory must be a regular, non-symlinked directory")
}
return nil
}
func copyImportTree(source, destination string) error {
return walkImportTree(source, func(path string, entry fs.DirEntry, info fs.FileInfo) error {
rel, err := filepath.Rel(source, path)
if err != nil {
return err
}
target := destination
if rel != "." {
target = filepath.Join(destination, rel)
}
if info.IsDir() {
if rel == "." {
return nil
}
return os.Mkdir(target, info.Mode().Perm())
}
if !info.Mode().IsRegular() {
return fmt.Errorf("only regular files may be imported: %s", path)
}
return copyImportFile(path, target, info.Mode().Perm())
})
}
func copyImportFile(source, destination string, mode fs.FileMode) error {
in, err := os.Open(source)
if err != nil {
return fmt.Errorf("read %s: %w", source, err)
}
defer in.Close()
out, err := os.OpenFile(destination, os.O_WRONLY|os.O_CREATE|os.O_EXCL, mode)
if err != nil {
return fmt.Errorf("create %s: %w", destination, err)
}
_, copyErr := io.Copy(out, in)
closeErr := out.Close()
if copyErr != nil {
return fmt.Errorf("copy %s: %w", source, copyErr)
}
if closeErr != nil {
return fmt.Errorf("write %s: %w", destination, closeErr)
}
return nil
}
func sameImportTree(source, destination string) (bool, error) {
seen := make(map[string]struct{})
same := true
err := walkImportTree(source, func(path string, entry fs.DirEntry, info fs.FileInfo) error {
rel, err := filepath.Rel(source, path)
if err != nil {
return err
}
seen[rel] = struct{}{}
other := destination
if rel != "." {
other = filepath.Join(destination, rel)
}
otherInfo, err := os.Lstat(other)
if errors.Is(err, fs.ErrNotExist) {
same = false
return nil
}
if err != nil {
return err
}
if otherInfo.Mode()&os.ModeSymlink != 0 || otherInfo.IsDir() != info.IsDir() || (!info.IsDir() && !otherInfo.Mode().IsRegular()) {
same = false
return nil
}
if info.Mode().IsRegular() {
equal, err := sameImportFile(path, other)
if err != nil {
return err
}
if !equal {
same = false
}
}
return nil
})
if err != nil || !same {
return same, err
}
err = walkImportTree(destination, func(path string, entry fs.DirEntry, info fs.FileInfo) error {
rel, err := filepath.Rel(destination, path)
if err != nil {
return err
}
if _, ok := seen[rel]; !ok {
same = false
}
return nil
})
return same, err
}
func sameImportFile(first, second string) (bool, error) {
a, err := os.Open(first)
if err != nil {
return false, err
}
defer a.Close()
b, err := os.Open(second)
if err != nil {
return false, err
}
defer b.Close()
left := make([]byte, 32*1024)
right := make([]byte, len(left))
for {
n, errA := a.Read(left)
m, errB := b.Read(right)
if n != m || !bytes.Equal(left[:n], right[:m]) {
return false, nil
}
if errA == io.EOF && errB == io.EOF {
return true, nil
}
if errA != nil && errA != io.EOF {
return false, errA
}
if errB != nil && errB != io.EOF {
return false, errB
}
if errA == io.EOF || errB == io.EOF {
return false, nil
}
}
}
// defaultSkillRoots returns skill directories ordered lowest- to
// highest-precedence. Non-existent directories are scanned harmlessly
// (DiscoverSkills skips them).

View File

@@ -21,6 +21,21 @@ func writeCatalogSkill(t *testing.T, dir, name, content string) {
}
}
func writeImportFixtureSkill(t *testing.T, dir string) {
t.Helper()
contents, err := os.ReadFile(filepath.Join("testdata", "import", "release-notes", skillFilename))
if err != nil {
t.Fatal(err)
}
path := filepath.Join(dir, "release-notes", skillFilename)
if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(path, contents, 0o644); err != nil {
t.Fatal(err)
}
}
func TestDiscoverAndLoadSkills(t *testing.T) {
dir := t.TempDir()
writeCatalogSkill(t, dir, "release-notes", "---\nname: release-notes\ndescription: Draft concise release notes.\nmetadata:\n author: Ollama\n labels:\n - release\n - docs\n---\n# Release notes\n\nUse short bullets.")
@@ -291,3 +306,187 @@ func TestSkillContentListsDirectoryAndResources(t *testing.T) {
t.Fatalf("content missing resource listing: %q", content)
}
}
func TestImportSkillsCopiesFixtureAndIsIdempotent(t *testing.T) {
source := t.TempDir()
destination := t.TempDir()
writeImportFixtureSkill(t, source)
writeCatalogSkill(t, source, "broken", "---\nname: another-skill\ndescription: Deliberately invalid.\n---\nIgnore this.")
if err := os.MkdirAll(filepath.Join(source, "release-notes", "references"), 0o755); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(filepath.Join(source, "release-notes", "references", "style.txt"), []byte("Keep it short.\n"), 0o644); err != nil {
t.Fatal(err)
}
if err := os.MkdirAll(filepath.Join(source, "release-notes", "scripts"), 0o755); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(filepath.Join(source, "release-notes", "scripts", "prepare.sh"), []byte("#!/bin/sh\n"), 0o755); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(filepath.Join(source, "ignored.md"), []byte("Ignored root file.\n"), 0o644); err != nil {
t.Fatal(err)
}
result, err := importSkillsFromDir("codex", source, destination)
if err != nil {
t.Fatal(err)
}
if got, want := strings.Join(result.Imported, ","), "release-notes"; got != want {
t.Fatalf("imported = %q, want %q", got, want)
}
catalog, err := DiscoverSkills(destination)
if err != nil {
t.Fatal(err)
}
skill, err := catalog.Load("release-notes")
if err != nil || skill.Description != "Draft concise release notes." {
t.Fatalf("imported skill = %#v, %v", skill, err)
}
if got := len(result.Failures); got != 1 || result.Failures[0].Name != "broken" {
t.Fatalf("failures = %#v, want broken fixture failure", result.Failures)
}
for _, file := range []string{skillFilename, filepath.Join("references", "style.txt"), filepath.Join("scripts", "prepare.sh")} {
if _, err := os.Stat(filepath.Join(destination, "release-notes", file)); err != nil {
t.Fatalf("imported fixture file %q: %v", file, err)
}
}
result, err = importSkillsFromDir("codex", source, destination)
if err != nil {
t.Fatal(err)
}
if got, want := strings.Join(result.Existing, ","), "release-notes"; got != want {
t.Fatalf("existing = %q, want %q", got, want)
}
if len(result.Imported) != 0 {
t.Fatalf("repeated import copied skills: %#v", result.Imported)
}
}
func TestImportSkillsLeavesConflictsAndUnsafeSourcesUntouched(t *testing.T) {
source := t.TempDir()
destination := t.TempDir()
writeCatalogSkill(t, source, "release-notes", "source instructions")
writeCatalogSkill(t, destination, "release-notes", "existing instructions")
writeCatalogSkill(t, source, "nested-link", "safe manifest")
if err := os.Symlink(filepath.Join(source, "release-notes", skillFilename), filepath.Join(source, "nested-link", "reference")); err != nil {
t.Skipf("symlink not supported: %v", err)
}
if err := os.Symlink(filepath.Join(source, "release-notes"), filepath.Join(source, "linked-skill")); err != nil {
t.Skipf("symlink not supported: %v", err)
}
result, err := importSkillsFromDir("codex", source, destination)
if err != nil {
t.Fatal(err)
}
if len(result.Imported) != 0 || len(result.Existing) != 0 {
t.Fatalf("unexpected successful import: %#v", result)
}
if got, err := os.ReadFile(filepath.Join(destination, "release-notes", skillFilename)); err != nil || !strings.Contains(string(got), "existing instructions") {
t.Fatalf("conflicting destination changed: %q, %v", got, err)
}
failed := make(map[string]bool)
for _, failure := range result.Failures {
failed[failure.Name] = true
}
for _, name := range []string{"release-notes", "nested-link", "linked-skill"} {
if !failed[name] {
t.Fatalf("missing failure for %q: %#v", name, result.Failures)
}
}
}
func TestImportSkillsRejectsSymlinkedRoot(t *testing.T) {
root := t.TempDir()
source := filepath.Join(t.TempDir(), "codex-skills")
if err := os.Symlink(root, source); err != nil {
t.Skipf("symlink not supported: %v", err)
}
result, err := importSkillsFromDir("codex", source, t.TempDir())
if err == nil || !strings.Contains(err.Error(), "symlinks are not supported") {
t.Fatalf("symlinked root error = %v", err)
}
if len(result.Imported) != 0 || len(result.Existing) != 0 || len(result.Failures) != 0 {
t.Fatalf("symlinked root result = %#v", result)
}
}
func TestImportSkillsMissingRootAndConfiguredRoots(t *testing.T) {
result, err := importSkillsFromDir("codex", filepath.Join(t.TempDir(), "missing"), t.TempDir())
if err != nil {
t.Fatal(err)
}
if len(result.Imported) != 0 || len(result.Existing) != 0 || len(result.Failures) != 0 {
t.Fatalf("missing root result = %#v", result)
}
destination := t.TempDir()
rootBase := t.TempDir()
roots := map[string]string{
"codex": filepath.Join(rootBase, "codex"),
"claude": filepath.Join(rootBase, "claude"),
"pi": filepath.Join(rootBase, "pi"),
}
for _, test := range []struct {
source string
root string
name string
}{
{source: "codex", root: roots["codex"], name: "from-codex"},
{source: "claude", root: roots["claude"], name: "from-claude"},
{source: "pi", root: roots["pi"], name: "from-pi"},
} {
t.Run(test.source, func(t *testing.T) {
writeCatalogSkill(t, test.root, test.name, "from "+test.source)
result, err = importSkillsFromRoots(test.source, roots, destination)
if err != nil {
t.Fatal(err)
}
if result.SourceDir != test.root {
t.Fatalf("source dir = %q, want %q", result.SourceDir, test.root)
}
if _, err := os.Stat(filepath.Join(destination, test.name, skillFilename)); err != nil {
t.Fatalf("conventional source was not imported: %v", err)
}
})
}
if _, err := importSkillsFromRoots("unknown", roots, destination); err == nil || !strings.Contains(err.Error(), "unknown skill source") {
t.Fatalf("unknown source error = %v", err)
}
}
func TestConventionalSkillImportRoots(t *testing.T) {
home := t.TempDir()
roots := conventionalSkillImportRoots(home)
for source, want := range map[string]string{
"codex": filepath.Join(home, ".codex", "skills"),
"claude": filepath.Join(home, ".claude", "skills"),
"pi": filepath.Join(home, ".pi", "agent", "skills"),
} {
if got := roots[source]; got != want {
t.Fatalf("%s root = %q, want %q", source, got, want)
}
}
}
func TestImportSkillsRejectsUnreadableManifest(t *testing.T) {
source := t.TempDir()
writeCatalogSkill(t, source, "private", "do not read")
manifest := filepath.Join(source, "private", skillFilename)
if err := os.Chmod(manifest, 0); err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = os.Chmod(manifest, 0o644) })
if _, err := os.ReadFile(manifest); err == nil {
t.Skip("test user can read a mode-000 file")
}
result, err := importSkillsFromDir("codex", source, t.TempDir())
if err != nil {
t.Fatal(err)
}
if len(result.Failures) != 1 || result.Failures[0].Name != "private" {
t.Fatalf("failures = %#v", result.Failures)
}
}

View File

@@ -0,0 +1,8 @@
---
name: release-notes
description: Draft concise release notes.
---
# Release notes
Use short bullets.

View File

@@ -83,12 +83,20 @@ func GenerateAgentTUI(cmd *cobra.Command, client *api.Client, opts agentTUIOptio
return agentContextWindowForModel(ctx, client, model, fallback)
}
skillCatalog, err := coreagent.LoadDefaultSkills(cwd)
if err != nil {
return fmt.Errorf("load agent skills: %w", err)
var skillCatalog *coreagent.SkillCatalog
reloadSkills := func() (*coreagent.SkillCatalog, error) {
catalog, err := coreagent.LoadDefaultSkills(cwd)
if err != nil {
return nil, err
}
for _, diagnostic := range catalog.Diagnostics() {
fmt.Fprintf(os.Stderr, "\033[1mwarning:\033[0m ignored invalid agent skill: %v\n", diagnostic)
}
skillCatalog = catalog
return catalog, nil
}
for _, diagnostic := range skillCatalog.Diagnostics() {
fmt.Fprintf(os.Stderr, "\033[1mwarning:\033[0m ignored invalid agent skill: %v\n", diagnostic)
if _, err := reloadSkills(); err != nil {
return fmt.Errorf("load agent skills: %w", err)
}
var registry *coreagent.Registry
registryForModel := func(ctx context.Context, model string) *coreagent.Registry {
@@ -99,7 +107,7 @@ func GenerateAgentTUI(cmd *cobra.Command, client *api.Client, opts agentTUIOptio
}
systemPrompt := agentSystemPromptWithWorkingDir(opts.Model, opts.System, agentSkillSystemContext(skillCatalog, registry, opts.ToolsDisabled), cwd)
_, err = agentchat.Run(cmd.Context(), agentchat.Options{
_, err := agentchat.Run(cmd.Context(), agentchat.Options{
Model: opts.Model,
Client: client,
Tools: registry,
@@ -118,6 +126,8 @@ func GenerateAgentTUI(cmd *cobra.Command, client *api.Client, opts agentTUIOptio
return agentSystemPromptWithWorkingDir(model, agentSystemFromShow(ctx, client, model), agentSkillSystemContext(skillCatalog, registry, toolsDisabled), cwd)
},
Skills: skillCatalog,
ImportSkills: coreagent.ImportSkills,
ReloadSkills: reloadSkills,
SystemPrompt: systemPrompt,
WorkingDir: cwd,
Format: opts.Format,

View File

@@ -55,6 +55,8 @@ type Options struct {
Client coreagent.ChatClient
Tools *coreagent.Registry
Skills *coreagent.SkillCatalog
ImportSkills func(string) (coreagent.SkillImportResult, error)
ReloadSkills func() (*coreagent.SkillCatalog, error)
ToolRegistryForModel func(context.Context, string) *coreagent.Registry
ToolsDisabled bool
MultiModalForModel func(context.Context, string) bool

View File

@@ -55,7 +55,7 @@ var chatSlashCommands = []chatSlashCommand{
{name: "/new", description: "start a new chat"},
{name: "/think", description: "set thinking mode"},
{name: "/tools", description: "toggle tools on or off"},
{name: "/skills", description: "list available skills"},
{name: "/skills", usage: "/skills [import codex|claude|pi]", description: "list or import skills"},
{name: "/compact", description: "summarize older context"},
{name: "/help", description: "show commands", aliases: []string{"/?"}},
{name: "/bye", description: "exit", aliases: []string{"/exit"}},
@@ -63,6 +63,12 @@ var chatSlashCommands = []chatSlashCommand{
{name: "/save", usage: "/save <filename>", description: "save request JSON; saved as <filename>.json"},
}
var skillsImportCompletions = []chatCompletion{
{value: "/skills import codex", label: "/skills import codex", description: "import from ~/.codex/skills"},
{value: "/skills import claude", label: "/skills import claude", description: "import from ~/.claude/skills"},
{value: "/skills import pi", label: "/skills import pi", description: "import from ~/.pi/agent/skills"},
}
func (m *chatModel) handleSubmit() (tea.Model, tea.Cmd) {
m.syncInputPlaceholders()
input := strings.TrimSpace(string(m.input))
@@ -159,8 +165,10 @@ func (m *chatModel) submitInput(input string) (tea.Model, tea.Cmd) {
}
func (m *chatModel) handleSkillsCommand(args string) (tea.Model, tea.Cmd) {
if strings.TrimSpace(args) != "" {
m.entries = append(m.entries, newChatEntry(chatEntry{role: "error", content: "usage: /skills"}))
if fields := strings.Fields(args); len(fields) == 2 && fields[0] == "import" {
return m.handleSkillsImport(fields[1])
} else if len(fields) != 0 {
m.entries = append(m.entries, newChatEntry(chatEntry{role: "error", content: "usage: /skills [import codex|claude|pi]"}))
return *m, nil
}
skills := m.opts.Skills.List()
@@ -181,6 +189,61 @@ func (m *chatModel) handleSkillsCommand(args string) (tea.Model, tea.Cmd) {
return *m, nil
}
func (m *chatModel) handleSkillsImport(source string) (tea.Model, tea.Cmd) {
importSkills := m.opts.ImportSkills
if importSkills == nil {
importSkills = coreagent.ImportSkills
}
result, err := importSkills(source)
if err != nil {
m.status = "error"
m.entries = append(m.entries, newChatEntry(chatEntry{role: "error", content: fmt.Sprintf("Could not import %s skills: %v", source, err)}))
return *m, nil
}
if len(result.Imported) != 0 || len(result.Existing) != 0 {
reload := m.opts.ReloadSkills
if reload == nil {
reload = func() (*coreagent.SkillCatalog, error) {
return coreagent.LoadDefaultSkills(m.currentWorkingDir())
}
}
catalog, err := reload()
if err != nil {
m.status = "error"
m.entries = append(m.entries, newChatEntry(chatEntry{role: "error", content: fmt.Sprintf("%s\n\nCould not reload skills: %v", skillsImportSummary(result), err)}))
return *m, nil
}
m.opts.Skills = catalog
if m.opts.ToolRegistryForModel != nil && m.opts.Model != "" {
m.opts.Tools = m.opts.ToolRegistryForModel(m.ctx, m.opts.Model)
}
if m.opts.SystemPromptForModel != nil {
m.opts.SystemPrompt = m.opts.SystemPromptForModel(m.ctx, m.opts.Model, m.opts.Tools, m.opts.ToolsDisabled)
}
m.status = "skills reloaded"
}
m.entries = append(m.entries, newSlashEntry(skillsImportSummary(result)))
return *m, nil
}
func skillsImportSummary(result coreagent.SkillImportResult) string {
if len(result.Imported) == 0 && len(result.Existing) == 0 && len(result.Failures) == 0 {
return fmt.Sprintf("No %s skills found at %s.", result.Source, result.SourceDir)
}
var lines []string
if len(result.Imported) != 0 {
lines = append(lines, fmt.Sprintf("Imported %d skill%s from %s.", len(result.Imported), pluralSuffix(len(result.Imported)), result.SourceDir))
}
if len(result.Existing) != 0 {
lines = append(lines, "Already present (left unchanged): "+strings.Join(result.Existing, ", ")+".")
}
for _, failure := range result.Failures {
lines = append(lines, fmt.Sprintf("Skipped %s: %v.", failure.Name, failure.Err))
}
return strings.Join(lines, "\n")
}
func skillsDirForDisplay(catalog *coreagent.SkillCatalog) string {
if catalog != nil && catalog.Dir() != "" {
return catalog.Dir()
@@ -1117,13 +1180,16 @@ func (m chatModel) completions() []chatCompletion {
func (m chatModel) slashCompletions() []chatCompletion {
rawInput := string(m.input)
input := strings.TrimSpace(rawInput)
input := strings.TrimLeftFunc(rawInput, unicode.IsSpace)
if !strings.HasPrefix(input, "/") {
return nil
}
if m.skillSlashPromptStarted(rawInput) {
return nil
}
if completions := matchingSkillsImportCompletions(input); completions != nil {
return completions
}
commands := matchingSlashCommands(input)
completions := make([]chatCompletion, 0, len(commands))
@@ -1134,6 +1200,13 @@ func (m chatModel) slashCompletions() []chatCompletion {
description: command.description,
})
}
if strings.EqualFold(input, "/skills") {
completions = append(completions, chatCompletion{
value: "/skills import",
label: "/skills import",
description: "import skills from Codex, Claude, or Pi",
})
}
// Each catalog skill is also invocable as "/<skill-name>"; surface them as
// completions so they are discoverable by typing.
if m.opts.Skills != nil {
@@ -1163,6 +1236,38 @@ func (m chatModel) slashCompletions() []chatCompletion {
return completions
}
func matchingSkillsImportCompletions(input string) []chatCompletion {
const importCommand = "/skills import"
lower := strings.ToLower(input)
if lower == "/skills" {
return nil // Preserve Enter on /skills as the listing command.
}
if !strings.HasPrefix(lower, "/skills ") {
return nil
}
if strings.HasPrefix(importCommand, lower) {
return []chatCompletion{{
value: importCommand,
label: importCommand,
description: "import skills from Codex, Claude, or Pi",
}}
}
if !strings.HasPrefix(lower, importCommand) {
return nil
}
prefix := strings.TrimSpace(strings.TrimPrefix(lower, importCommand))
completions := make([]chatCompletion, 0, len(skillsImportCompletions))
for _, completion := range skillsImportCompletions {
if strings.HasPrefix(strings.TrimPrefix(completion.value, importCommand+" "), prefix) {
completions = append(completions, completion)
}
}
if len(completions) == 0 {
return []chatCompletion{{label: "No matching skill sources"}}
}
return completions
}
func (m chatModel) skillSlashPromptStarted(input string) bool {
input = strings.TrimLeftFunc(input, unicode.IsSpace)
end := strings.IndexFunc(input, unicode.IsSpace)

View File

@@ -570,6 +570,76 @@ func TestSkillCommandsListAndPersistSyntheticToolCall(t *testing.T) {
}
}
func TestSkillsImportReloadsCatalogRegistryAndSystemPrompt(t *testing.T) {
before := writeTestSkillCatalog(t)
dir := t.TempDir()
if err := os.Mkdir(filepath.Join(dir, "from-codex"), 0o755); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(filepath.Join(dir, "from-codex", "SKILL.md"), []byte("---\nname: from-codex\ndescription: Imported skill.\n---\nImported instructions."), 0o644); err != nil {
t.Fatal(err)
}
after, err := coreagent.DiscoverSkills(dir)
if err != nil {
t.Fatal(err)
}
registry := &coreagent.Registry{}
var reloaded, rebuilt, prompted bool
m := chatModel{
ctx: context.Background(),
opts: Options{
Model: "test",
Skills: before,
ImportSkills: func(source string) (coreagent.SkillImportResult, error) {
if source != "codex" {
t.Fatalf("source = %q", source)
}
return coreagent.SkillImportResult{Source: source, SourceDir: "/source", Imported: []string{"from-codex"}}, nil
},
ReloadSkills: func() (*coreagent.SkillCatalog, error) {
reloaded = true
return after, nil
},
ToolRegistryForModel: func(context.Context, string) *coreagent.Registry {
rebuilt = true
return registry
},
SystemPromptForModel: func(_ context.Context, _ string, got *coreagent.Registry, _ bool) string {
prompted = got == registry
return after.SystemContext()
},
},
input: []rune("/skills import codex"),
}
updated, cmd := m.handleSubmit()
if cmd != nil {
t.Fatal("skills import should not start a model run")
}
m = updated.(chatModel)
if !reloaded || !rebuilt || !prompted {
t.Fatalf("reload=%v rebuilt=%v prompted=%v", reloaded, rebuilt, prompted)
}
if m.opts.Skills != after || m.opts.Tools != registry || !strings.Contains(m.opts.SystemPrompt, "from-codex") {
t.Fatalf("reloaded options = %#v", m.opts)
}
if m.status != "skills reloaded" || len(m.entries) != 1 || !strings.Contains(m.entries[0].content, "Imported 1 skill") {
t.Fatalf("import result = status %q entries %#v", m.status, m.entries)
}
}
func TestSkillsImportUsage(t *testing.T) {
m := chatModel{input: []rune("/skills import")}
updated, cmd := m.handleSubmit()
if cmd != nil {
t.Fatal("invalid skills import should not start a model run")
}
m = updated.(chatModel)
if len(m.entries) != 1 || m.entries[0].role != "error" || !strings.Contains(m.entries[0].content, "usage: /skills [import codex|claude|pi]") {
t.Fatalf("entries = %#v", m.entries)
}
}
func TestSkillSlashCommandPromptBecomesUserMessage(t *testing.T) {
catalog := writeTestSkillCatalog(t)
m := chatModel{ctx: context.Background(), opts: Options{Model: "test", Skills: catalog, Client: chatTestClient{}}, input: []rune("/release-notes draft the v1.2 notes")}
@@ -665,6 +735,31 @@ func TestSkillSlashCommandAppearsInCompletions(t *testing.T) {
}
}
func TestSkillsImportSlashCompletions(t *testing.T) {
for _, test := range []struct {
input string
want []string
}{
{input: "/skills", want: []string{"/skills", "/skills import"}},
{input: "/skills impo", want: []string{"/skills import"}},
{input: "/skills import ", want: []string{"/skills import codex", "/skills import claude", "/skills import pi"}},
{input: "/skills import c", want: []string{"/skills import codex", "/skills import claude"}},
{input: "/skills import pi", want: []string{"/skills import pi"}},
} {
t.Run(test.input, func(t *testing.T) {
m := chatModel{input: []rune(test.input)}
completions := m.slashCompletions()
got := make([]string, 0, len(completions))
for _, completion := range completions {
got = append(got, completion.value)
}
if strings.Join(got, "\n") != strings.Join(test.want, "\n") {
t.Fatalf("completions = %#v, want %#v", got, test.want)
}
})
}
}
func TestSkillSlashPromptHidesCommandCompletions(t *testing.T) {
catalog := writeTestSkillCatalog(t)
for _, input := range []string{"/release-notes ", "/release-notes draft the release notes"} {

View File

@@ -244,26 +244,33 @@ Ollama Cloud model retirement does not affect local models.
| Retirement date | Model | Recommended alternative |
| --- | --- | --- |
| July 15, 2026 | `deepseek-v3.1:671b` | `deepseek-v4-flash` |
| July 15, 2026 | `deepseek-v3.2` | `deepseek-v4-flash` |
| July 15, 2026 | `devstral-2:123b` | `mistral-large-3:675b` |
| July 15, 2026 | `devstral-small-2:24b` | |
| July 15, 2026 | `ministral-3:14b` | |
| July 15, 2026 | `ministral-3:3b` | |
| July 15, 2026 | `ministral-3:8b` | |
| July 15, 2026 | `gemini-3-flash-preview` | `minimax-m3` |
| July 15, 2026 | `gemma3:12b` | `gemma4:31b` |
| July 15, 2026 | `gemma3:27b` | `gemma4:31b` |
| July 15, 2026 | `gemma3:4b` | `gemma4:31b` |
| July 15, 2026 | `glm-4.7` | `glm-5.2` |
| July 15, 2026 | `glm-5` | `glm-5.2` |
| July 15, 2026 | `minimax-m2.1` | `minimax-m3` |
| July 15, 2026 | `qwen3-coder-next` | `qwen3.5:397b` |
| July 15, 2026 | `qwen3-coder:480b` | `qwen3.5:397b` |
| July 31, 2026 | `minimax-m2.5` | `minimax-m2.7` |
| July 31, 2026 | `kimi-k2.5` | `kimi-k2.6` |
### Past retirements
<AccordionGroup>
<Accordion title="July 15, 2026">
| Model | Recommended alternative |
| --- | --- |
| `deepseek-v3.1:671b` | `deepseek-v4-flash` |
| `deepseek-v3.2` | `deepseek-v4-flash` |
| `devstral-2:123b` | `mistral-large-3:675b` |
| `devstral-small-2:24b` | |
| `ministral-3:14b` | |
| `ministral-3:3b` | |
| `ministral-3:8b` | |
| `gemini-3-flash-preview` | `minimax-m3` |
| `gemma3:12b` | `gemma4:31b` |
| `gemma3:27b` | `gemma4:31b` |
| `gemma3:4b` | `gemma4:31b` |
| `glm-4.7` | `glm-5.2` |
| `glm-5` | `glm-5.2` |
| `minimax-m2.1` | `minimax-m3` |
| `qwen3-coder-next` | `qwen3.5:397b` |
| `qwen3-coder:480b` | `qwen3.5:397b` |
</Accordion>
<Accordion title="June 30, 2026">
| Model | Recommended alternative |
| --- | --- |

View File

@@ -2,8 +2,21 @@
This directory contains integration tests to exercise Ollama end-to-end to verify behavior
By default, these tests are disabled so `go test ./...` will exercise only unit tests. To run integration tests you must pass the integration tag. `go test -tags=integration ./...` Some tests require additional tags to enable to allow scoped testing to keep the duration reasonable. For example, testing a broad set of models requires `-tags=integration,models` and a longer timeout (~60m or more depending on the speed of your GPU.). To view the current set of tag combinations use `find integration -type f | xargs grep "go:build"`
By default, these tests are disabled so `go test ./...` will exercise only unit tests. To run integration tests, pass the `integration` tag and one of the scoped tags:
```bash
go test -tags=integration,fast -v -count 1 ./integration/
go test -tags=integration,release -v -count 1 -timeout 30m ./integration/
go test -tags=integration,library -v -count 1 -timeout 120m ./integration/
```
Tags:
- `fast`: quick runner/model smoke coverage.
- `release`: release regression coverage.
- `library`: broad library coverage requiring about 2.5 TiB of disk space.
Scope wiring and model selections live in `integration/reg_fast_test.go`, `integration/reg_release_test.go`, and `integration/reg_library_test.go`.
The integration tests have 2 modes of operating.
@@ -21,12 +34,12 @@ harness starts the server.
## Testing a New Model
When implementing new model architecture, use `OLLAMA_TEST_MODEL` to run the
integration suite against your model.
integration suite against your model with either the `fast` or `release` coverage.
```bash
# Build the binary first
go build .
# Run integration tests against it
OLLAMA_TEST_MODEL=mymodel go test -tags integration -v -count 1 -timeout 15m ./integration/
OLLAMA_TEST_MODEL=mymodel go test -tags=integration,fast -v -count 1 ./integration/
```

View File

@@ -14,6 +14,13 @@ import (
"github.com/ollama/ollama/api"
)
const (
apiTestTimeout = 4 * time.Minute
apiInitialResponseTimeout = time.Minute
apiOverrideInitialResponseTimeout = 2 * time.Minute
apiStreamResponseTimeout = 30 * time.Second
)
func assertBytesMatchToken(t *testing.T, label, token string, ints []int) {
t.Helper()
@@ -31,10 +38,12 @@ func assertBytesMatchToken(t *testing.T, label, token string, ints []int) {
}
}
func TestAPIGenerate(t *testing.T) {
initialTimeout := 60 * time.Second
streamTimeout := 30 * time.Second
ctx, cancel := context.WithTimeout(context.Background(), 1*time.Minute)
func runAPIGenerate(t *testing.T) {
initialTimeout := apiInitialResponseTimeout
if testModel != "" {
initialTimeout = apiOverrideInitialResponseTimeout
}
ctx, cancel := context.WithTimeout(context.Background(), apiTestTimeout)
defer cancel()
// Set up the test data
req := api.GenerateRequest{
@@ -45,7 +54,6 @@ func TestAPIGenerate(t *testing.T) {
"seed": 123,
},
}
client, _, cleanup := InitServerConnection(ctx, t)
defer cleanup()
pullOrSkip(ctx, t, client, req.Model)
@@ -105,7 +113,7 @@ func TestAPIGenerate(t *testing.T) {
} // else incremental response, nothing to check right now...
buf.Write([]byte(response.Response))
if !stallTimer.Reset(streamTimeout) {
if !stallTimer.Reset(apiStreamResponseTimeout) {
return fmt.Errorf("stall was detected while streaming response, aborting")
}
return nil
@@ -188,10 +196,12 @@ func TestAPIGenerate(t *testing.T) {
}
}
func TestAPIChat(t *testing.T) {
initialTimeout := 60 * time.Second
streamTimeout := 30 * time.Second
ctx, cancel := context.WithTimeout(context.Background(), 1*time.Minute)
func runAPIChat(t *testing.T) {
initialTimeout := apiInitialResponseTimeout
if testModel != "" {
initialTimeout = apiOverrideInitialResponseTimeout
}
ctx, cancel := context.WithTimeout(context.Background(), apiTestTimeout)
defer cancel()
// Set up the test data
req := api.ChatRequest{
@@ -207,7 +217,6 @@ func TestAPIChat(t *testing.T) {
"seed": 123,
},
}
client, _, cleanup := InitServerConnection(ctx, t)
defer cleanup()
pullOrSkip(ctx, t, client, req.Model)
@@ -265,7 +274,7 @@ func TestAPIChat(t *testing.T) {
}
} // else incremental response, nothing to check right now...
buf.Write([]byte(response.Message.Content))
if !stallTimer.Reset(streamTimeout) {
if !stallTimer.Reset(apiStreamResponseTimeout) {
return fmt.Errorf("stall was detected while streaming response, aborting")
}
return nil
@@ -310,11 +319,11 @@ func TestAPIChat(t *testing.T) {
}
}
func TestAPIListModels(t *testing.T) {
func runAPIListModels(t *testing.T) {
if testModel != "" {
t.Skip("skipping metadata test with model override")
}
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
ctx, cancel := context.WithTimeout(context.Background(), apiTestTimeout)
defer cancel()
client, _, cleanup := InitServerConnection(ctx, t)
defer cleanup()
@@ -331,44 +340,54 @@ func TestAPIListModels(t *testing.T) {
if len(resp.Models) == 0 {
t.Fatalf("list should not be empty")
}
model := resp.Models[0]
var model *api.ListModelResponse
for i := range resp.Models {
if resp.Models[i].Name == smol || resp.Models[i].Model == smol || strings.Contains(resp.Models[i].Name, smol) || strings.Contains(resp.Models[i].Model, smol) {
model = &resp.Models[i]
break
}
}
if model == nil {
t.Fatalf("list should include pulled model %s: %#v", smol, resp.Models)
}
if model.Name == "" {
t.Errorf("first model name empty: %#v", model)
t.Errorf("model name empty: %#v", model)
}
var nilTime time.Time
if model.ModifiedAt == nilTime {
t.Errorf("first model modified_at empty: %#v", model)
t.Errorf("model modified_at empty: %#v", model)
}
if model.Size == 0 {
t.Errorf("first model size empty: %#v", model)
t.Errorf("model size empty: %#v", model)
}
if model.Digest == "" {
t.Errorf("first model digest empty: %#v", model)
t.Errorf("model digest empty: %#v", model)
}
verifyModelDetails(t, model.Details)
}
func verifyModelDetails(t *testing.T, details api.ModelDetails) {
if details.Format == "" {
t.Errorf("first model details.format empty: %#v", details)
t.Errorf("model details.format empty: %#v", details)
}
if details.Family == "" {
t.Errorf("first model details.family empty: %#v", details)
t.Errorf("model details.family empty: %#v", details)
}
if details.ParameterSize == "" {
t.Errorf("first model details.parameter_size empty: %#v", details)
t.Errorf("model details.parameter_size empty: %#v", details)
}
if details.QuantizationLevel == "" {
t.Errorf("first model details.quantization_level empty: %#v", details)
t.Errorf("model details.quantization_level empty: %#v", details)
}
}
func TestAPIShowModel(t *testing.T) {
func runAPIShowModel(t *testing.T) {
if testModel != "" {
t.Skip("skipping metadata test with model override")
}
modelName := "llama3.2"
ctx, cancel := context.WithTimeout(context.Background(), 1*time.Minute)
ctx, cancel := context.WithTimeout(context.Background(), apiTestTimeout)
defer cancel()
client, _, cleanup := InitServerConnection(ctx, t)
defer cleanup()
@@ -405,8 +424,8 @@ func TestAPIShowModel(t *testing.T) {
}
}
func TestAPIGenerateLogprobs(t *testing.T) {
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Minute)
func runAPIGenerateLogprobs(t *testing.T) {
ctx, cancel := context.WithTimeout(context.Background(), apiTestTimeout)
defer cancel()
client, _, cleanup := InitServerConnection(ctx, t)
@@ -518,8 +537,8 @@ func TestAPIGenerateLogprobs(t *testing.T) {
}
}
func TestAPIChatLogprobs(t *testing.T) {
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Minute)
func runAPIChatLogprobs(t *testing.T) {
ctx, cancel := context.WithTimeout(context.Background(), apiTestTimeout)
defer cancel()
client, _, cleanup := InitServerConnection(ctx, t)

View File

@@ -18,12 +18,6 @@ import (
"github.com/ollama/ollama/api"
)
var defaultAudioModels = []string{
"nemotron3:33b",
"gemma4:e2b",
"gemma4:e4b",
}
// decodeTestAudio returns the test audio clip ("Why is the sky blue?", 16kHz mono WAV).
func decodeTestAudio(t *testing.T) api.ImageData {
t.Helper()
@@ -37,61 +31,56 @@ func decodeTestAudio(t *testing.T) api.ImageData {
// setupAudioModel pulls the model, preloads it, and skips if it doesn't support audio.
func setupAudioModel(ctx context.Context, t *testing.T, client *api.Client, model string) {
t.Helper()
if testModel == "" {
pullOrSkip(ctx, t, client, model)
}
pullOrSkip(ctx, t, client, model)
skipIfModelTooLargeForVRAM(ctx, t, client, model)
requireCapability(ctx, t, client, model, "audio")
err := client.Generate(ctx, &api.GenerateRequest{Model: model}, func(response api.GenerateResponse) error { return nil })
if err != nil {
t.Fatalf("failed to load model %s: %s", model, err)
}
preloadGenerateModel(ctx, t, client, api.GenerateRequest{Model: model})
}
// TestAudioTranscription tests that the model can transcribe audio to text.
func TestAudioTranscription(t *testing.T) {
for _, model := range testModels(defaultAudioModels) {
t.Run(model, func(t *testing.T) {
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Minute)
defer cancel()
client, _, cleanup := InitServerConnection(ctx, t)
defer cleanup()
setupAudioModel(ctx, t, client, model)
audio := decodeTestAudio(t)
noThink := &api.ThinkValue{Value: false}
req := api.ChatRequest{
Model: model,
Think: noThink,
Messages: []api.Message{
{
Role: "system",
Content: "Transcribe the audio exactly as spoken. Output only the spoken words. Do not answer any question in the audio.",
},
{
Role: "user",
Content: "What exact words are spoken in this audio?",
Images: []api.ImageData{audio},
},
},
Stream: &stream,
Options: map[string]any{
"temperature": 0,
"seed": 123,
"num_predict": 50,
},
}
// The audio says "Why is the sky blue?" — expect key words in transcription.
DoChat(ctx, t, client, req, []string{"sky", "blue"}, 60*time.Second, 10*time.Second)
})
}
func registerAudioTranscriptionCases(models []string) {
registerModelIntegrationCases("audio-transcription", models, runAudioTranscriptionModel)
}
// TestAudioResponse tests that the model can respond to a spoken question.
func TestAudioResponse(t *testing.T) {
for _, model := range testModels(defaultAudioModels) {
func runAudioTranscriptionModel(t *testing.T, model string) {
t.Helper()
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Minute)
defer cancel()
client, _, cleanup := InitServerConnection(ctx, t)
defer cleanup()
setupAudioModel(ctx, t, client, model)
audio := decodeTestAudio(t)
noThink := &api.ThinkValue{Value: false}
req := api.ChatRequest{
Model: model,
Think: noThink,
Messages: []api.Message{
{
Role: "system",
Content: "Transcribe the audio exactly as spoken. Output only the spoken words. Do not answer any question in the audio.",
},
{
Role: "user",
Content: "What exact words are spoken in this audio?",
Images: []api.ImageData{audio},
},
},
Stream: &stream,
Options: map[string]any{
"temperature": 0,
"seed": 123,
"num_predict": 50,
},
}
// The audio says "Why is the sky blue?" - expect key words in transcription.
DoChat(ctx, t, client, req, []string{"sky", "blue"}, 60*time.Second, 10*time.Second)
}
// runAudioResponse tests that the model can respond to a spoken question.
func runAudioResponse(t *testing.T, models []string) {
for _, model := range testModels(models) {
t.Run(model, func(t *testing.T) {
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Minute)
defer cancel()
@@ -128,9 +117,9 @@ func TestAudioResponse(t *testing.T) {
}
}
// TestOpenAIAudioTranscription tests the /v1/audio/transcriptions endpoint.
func TestOpenAIAudioTranscription(t *testing.T) {
for _, model := range testModels(defaultAudioModels) {
// runOpenAIAudioTranscription tests the /v1/audio/transcriptions endpoint.
func runOpenAIAudioTranscription(t *testing.T, models []string) {
for _, model := range testModels(models) {
t.Run(model, func(t *testing.T) {
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Minute)
defer cancel()
@@ -182,9 +171,9 @@ func TestOpenAIAudioTranscription(t *testing.T) {
}
}
// TestOpenAIChatWithAudio tests /v1/chat/completions with input_audio content.
func TestOpenAIChatWithAudio(t *testing.T) {
for _, model := range testModels(defaultAudioModels) {
// runOpenAIChatWithAudio tests /v1/chat/completions with input_audio content.
func runOpenAIChatWithAudio(t *testing.T, models []string) {
for _, model := range testModels(models) {
t.Run(model, func(t *testing.T) {
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Minute)
defer cancel()

View File

@@ -13,8 +13,8 @@ import (
"github.com/ollama/ollama/api"
)
func TestBlueSky(t *testing.T) {
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Minute)
func runBlueSky(t *testing.T) {
ctx, cancel := context.WithTimeout(context.Background(), 4*time.Minute)
defer cancel()
// Set up the test data
req := api.ChatRequest{
@@ -34,17 +34,17 @@ func TestBlueSky(t *testing.T) {
ChatTestHelper(ctx, t, req, blueSkyExpected)
}
func TestUnicode(t *testing.T) {
func runUnicode(t *testing.T, model string) {
if testModel != "" {
t.Skip("uses hardcoded model, not applicable with model override")
}
skipUnderMinVRAM(t, 12) // Actual model load is ~26G
skipRegisteredMinVRAM(t, model)
ctx, cancel := context.WithTimeout(context.Background(), 4*time.Minute)
defer cancel()
// Set up the test data
req := api.ChatRequest{
// DeepSeek has a Unicode tokenizer regex, making it a unicode torture test
Model: "deepseek-coder-v2:16b-lite-instruct-q2_K", // TODO is there an ollama-engine model we can switch to and keep the coverage?
Model: model, // TODO is there an ollama-engine model we can switch to and keep the coverage?
Messages: []api.Message{
{
Role: "user",
@@ -63,11 +63,7 @@ func TestUnicode(t *testing.T) {
client, _, cleanup := InitServerConnection(ctx, t)
defer cleanup()
pullOrSkip(ctx, t, client, req.Model)
slog.Info("loading", "model", req.Model)
err := client.Generate(ctx, &api.GenerateRequest{Model: req.Model}, func(response api.GenerateResponse) error { return nil })
if err != nil {
t.Fatalf("failed to load model %s: %s", req.Model, err)
}
preloadGenerateModel(ctx, t, client, api.GenerateRequest{Model: req.Model})
defer func() {
// best effort unload once we're done with the model
client.Generate(ctx, &api.GenerateRequest{Model: req.Model, KeepAlive: &api.Duration{Duration: 0}}, func(rsp api.GenerateResponse) error { return nil })
@@ -81,15 +77,15 @@ func TestUnicode(t *testing.T) {
}, 180*time.Second, 30*time.Second)
}
func TestExtendedUnicodeOutput(t *testing.T) {
func runExtendedUnicodeOutput(t *testing.T, model string) {
if testModel != "" {
t.Skip("uses hardcoded model, not applicable with model override")
}
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Minute)
ctx, cancel := context.WithTimeout(context.Background(), 4*time.Minute)
defer cancel()
// Set up the test data
req := api.ChatRequest{
Model: "gemma2:2b",
Model: model,
Messages: []api.Message{
{
Role: "user",
@@ -108,14 +104,14 @@ func TestExtendedUnicodeOutput(t *testing.T) {
DoChat(ctx, t, client, req, []string{"😀", "😊", "😁", "😂", "😄", "😃"}, 120*time.Second, 120*time.Second)
}
func TestUnicodeModelDir(t *testing.T) {
func runUnicodeModelDir(t *testing.T) {
// This is only useful for Windows with utf-16 characters, so skip this test for other platforms
if runtime.GOOS != "windows" {
t.Skip("Unicode test only applicable to windows")
}
// Only works for local testing
if os.Getenv("OLLAMA_TEST_EXISTING") != "" {
t.Skip("TestUnicodeModelDir only works for local testing, skipping")
t.Skip("runUnicodeModelDir only works for local testing, skipping")
}
modelDir, err := os.MkdirTemp("", "ollama_埃")
@@ -127,7 +123,7 @@ func TestUnicodeModelDir(t *testing.T) {
t.Setenv("OLLAMA_MODELS", modelDir)
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Minute)
ctx, cancel := context.WithTimeout(context.Background(), 4*time.Minute)
defer cancel()
req := api.ChatRequest{
@@ -147,22 +143,22 @@ func TestUnicodeModelDir(t *testing.T) {
ChatTestHelper(ctx, t, req, blueSkyExpected)
}
// TestNumPredict verifies that when num_predict is set, the model generates
// runNumPredict verifies that when num_predict is set, the model generates
// exactly that many tokens. It uses logprobs to count the actual tokens output.
func TestNumPredict(t *testing.T) {
func runNumPredict(t *testing.T, model string) {
if testModel != "" {
t.Skip("uses hardcoded model, not applicable with model override")
}
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Minute)
ctx, cancel := context.WithTimeout(context.Background(), 4*time.Minute)
defer cancel()
client, _, cleanup := InitServerConnection(ctx, t)
defer cleanup()
pullOrSkip(ctx, t, client, "qwen3:0.6b")
pullOrSkip(ctx, t, client, model)
req := api.GenerateRequest{
Model: "qwen3:0.6b",
Model: model,
Prompt: "Write a long story.",
Stream: &stream,
Logprobs: true,

View File

@@ -0,0 +1,144 @@
//go:build integration
package integration
import (
"context"
"fmt"
"log/slog"
"os"
"strconv"
"strings"
"sync"
"testing"
"time"
"github.com/ollama/ollama/api"
"github.com/ollama/ollama/format"
)
var sweepVRAMWarning sync.Once
func registerChatCases(models []string) {
registerModelIntegrationCases("chat", models, runChatModel)
}
func runChatModel(t *testing.T, model string) {
t.Helper()
softTimeout, hardTimeout := getTimeouts(t)
slog.Info("Setting timeouts", "soft", softTimeout, "hard", hardTimeout)
ctx, cancel := context.WithTimeout(context.Background(), hardTimeout)
defer cancel()
client, _, cleanup := InitServerConnection(ctx, t)
defer cleanup()
if time.Since(started) > softTimeout {
t.Skip("skipping remaining tests to avoid excessive runtime")
}
skipRegisteredMinVRAM(t, model)
requireCapability(ctx, t, client, model, "completion")
skipIfTargetArchitecture(ctx, t, client, model)
skipIfModelTooLargeForSweepVRAM(ctx, t, client, model)
initialTimeout := 120 * time.Second
streamTimeout := 30 * time.Second
preloadGenerateModel(ctx, t, client, api.GenerateRequest{Model: model, KeepAlive: &api.Duration{Duration: 10 * time.Second}})
defer func() {
client.Generate(ctx, &api.GenerateRequest{Model: model, KeepAlive: &api.Duration{Duration: 0}}, func(rsp api.GenerateResponse) error { return nil })
}()
gpuPercent := getGPUPercent(ctx, t, client, model)
if gpuPercent < 80 {
slog.Warn("Low GPU percentage - increasing timeouts", "percent", gpuPercent)
initialTimeout = 240 * time.Second
streamTimeout = 40 * time.Second
}
req, anyResp := chatModelRequest(model)
DoChat(ctx, t, client, req, anyResp, initialTimeout, streamTimeout)
}
func chatModelRequest(model string) (api.ChatRequest, []string) {
req := api.ChatRequest{
Model: model,
Messages: []api.Message{
{
Role: "user",
Content: blueSkyPrompt,
},
},
KeepAlive: &api.Duration{Duration: 10 * time.Second},
Options: map[string]any{
"temperature": 0.1,
"seed": 123,
},
}
anyResp := blueSkyExpected
// Special cases
if model == "duckdb-nsql" {
anyResp = []string{"select", "from"}
} else if model == "granite3-guardian" || model == "shieldgemma" || model == "llama-guard3" || model == "bespoke-minicheck" {
anyResp = []string{"yes", "no", "safe", "unsafe"}
} else if model == "openthinker" {
anyResp = []string{"plugin", "im_sep", "components", "function call"}
} else if model == "starcoder" || model == "starcoder2" || model == "magicoder" || model == "deepseek-coder" {
req.Messages[0].Content = "def fibonacci():"
anyResp = []string{"f(n)", "sequence", "n-1", "main()", "__main__", "while"}
}
return req, anyResp
}
func skipIfTargetArchitecture(ctx context.Context, t *testing.T, client *api.Client, model string) {
t.Helper()
targetArch := os.Getenv("OLLAMA_TEST_ARCHITECTURE")
if targetArch == "" {
return
}
resp, err := client.Show(ctx, &api.ShowRequest{Name: model})
if err != nil {
t.Fatalf("unable to show model: %s", err)
}
arch := resp.ModelInfo["general.architecture"].(string)
if arch != targetArch {
t.Skip(fmt.Sprintf("Skipping %s architecture %s != %s", model, arch, targetArch))
}
}
func skipIfModelTooLargeForSweepVRAM(ctx context.Context, t *testing.T, client *api.Client, model string) {
t.Helper()
s := os.Getenv("OLLAMA_MAX_VRAM")
if s == "" {
sweepVRAMWarning.Do(func() {
slog.Warn("No VRAM info available, testing all models, so larger ones might timeout...")
})
return
}
maxVram, err := strconv.ParseUint(s, 10, 64)
if err != nil {
t.Fatalf("invalid OLLAMA_MAX_VRAM %v", err)
}
resp, err := client.List(ctx)
if err != nil {
t.Fatalf("list models failed %v", err)
}
for _, m := range resp.Models {
if modelNameMatches(model, m.Name) && float32(m.Size)*1.2 > float32(maxVram) {
t.Skipf("model %s is too large for available VRAM: %s > %s", model, format.HumanBytes(m.Size), format.HumanBytes(int64(maxVram)))
}
}
}
func modelNameMatches(model, name string) bool {
if name == model {
return true
}
return !strings.Contains(model, ":") && strings.HasPrefix(name, model+":")
}

View File

@@ -20,7 +20,7 @@ import (
)
// Send multiple requests in parallel (concurrently) to a single model and ensure responses are expected
func TestConcurrentChat(t *testing.T) {
func runConcurrentChat(t *testing.T) {
// Assumes all requests have the same model
req, resp := ChatRequests()
numParallel := int(envconfig.NumParallel() + 1)
@@ -31,16 +31,10 @@ func TestConcurrentChat(t *testing.T) {
defer cancel()
client, _, cleanup := InitServerConnection(ctx, t)
defer cleanup()
pullOrSkip(ctx, t, client, req[0].Model)
// Get the server running (if applicable) warm the model up with a single initial request
slog.Info("loading", "model", req[0].Model)
err := client.Generate(ctx,
&api.GenerateRequest{Model: req[0].Model, KeepAlive: &api.Duration{Duration: 10 * time.Second}},
func(response api.GenerateResponse) error { return nil },
)
if err != nil {
t.Fatalf("failed to load model %s: %s", req[0].Model, err)
}
preloadGenerateModel(ctx, t, client, api.GenerateRequest{Model: req[0].Model, KeepAlive: &api.Duration{Duration: 10 * time.Second}})
var wg sync.WaitGroup
r := rand.New(rand.NewSource(0))
@@ -66,7 +60,7 @@ func TestConcurrentChat(t *testing.T) {
// Stress the scheduler and attempt to load more models than will fit to cause thrashing
// This test will always load at least 2 models even on CPU based systems
func TestMultiModelStress(t *testing.T) {
func runMultiModelStress(t *testing.T) {
if testModel != "" {
t.Skip("uses hardcoded models, not applicable with model override")
}
@@ -85,7 +79,6 @@ func TestMultiModelStress(t *testing.T) {
"llama3.2:1b",
"qwen3:0.6b",
"gemma2:2b",
"deepseek-r1:1.5b", // qwen2 arch
"gemma3:270m",
}
mediumModels := []string{
@@ -126,12 +119,8 @@ func TestMultiModelStress(t *testing.T) {
slog.Info("Loading models to find how many can fit in VRAM before overflowing")
chooseModels:
for i, model := range chosenModels {
req := &api.GenerateRequest{Model: model} // Leave KeepAlive unset so they stay loaded until the scheduler decides to unload them
slog.Info("loading", "model", model)
err = client.Generate(ctx, req, func(response api.GenerateResponse) error { return nil })
if err != nil {
t.Fatalf("failed to load model %s: %s", model, err)
}
// Leave KeepAlive unset so they stay loaded until the scheduler decides to unload them.
preloadGenerateModel(ctx, t, client, api.GenerateRequest{Model: model})
targetLoadCount++
if i > 0 {
models, err := client.ListRunning(ctx)

View File

@@ -14,7 +14,12 @@ import (
"github.com/ollama/ollama/api"
)
func TestLongInputContext(t *testing.T) {
const (
longInputTimeout = 2 * time.Minute
longInputModelOverrideTimeout = 3 * time.Minute
)
func runLongInputContext(t *testing.T) {
// Setting NUM_PARALLEL to 1 ensures the allocated context is exactly what
// we asked for and there is nothing extra that we could spill over into.
// Context shift happens after a prompt has been admitted to a slot. Initial
@@ -23,7 +28,11 @@ func TestLongInputContext(t *testing.T) {
// prompt while llama-server reports it as too large to admit.
t.Setenv("OLLAMA_NUM_PARALLEL", "1")
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Minute)
timeout := longInputTimeout
if testModel != "" {
timeout = longInputModelOverrideTimeout
}
ctx, cancel := context.WithTimeout(context.Background(), timeout)
defer cancel()
req := api.ChatRequest{
Model: smol,
@@ -79,7 +88,7 @@ func isContextLimitError(err string) bool {
strings.Contains(err, "too long"))
}
func TestContextExhaustion(t *testing.T) {
func runContextExhaustion(t *testing.T) {
// Setting NUM_PARALLEL to 1 ensures the allocated context is exactly what
// we asked for and there is nothing extra that we could spill over into
t.Setenv("OLLAMA_NUM_PARALLEL", "1")
@@ -128,11 +137,13 @@ func containsEmoji(s string) bool {
}
// Send multiple generate requests with prior context and ensure the response is coherant and expected
func TestParallelGenerateWithHistory(t *testing.T) {
func runParallelGenerateWithHistory(t *testing.T, modelName string) {
if testModel != "" {
t.Skip("uses hardcoded model, not applicable with model override")
// The Generate API's Context field (token array continuation) is not
// supported by all runners (e.g. MLX). Chat history works; this is
// the only generate-specific continuation path.
t.Skip("generate context continuation not supported by all runners")
}
modelName := "gpt-oss:20b"
req, resp := GenerateRequests()
numParallel := 2
iterLimit := 2
@@ -144,16 +155,10 @@ func TestParallelGenerateWithHistory(t *testing.T) {
defer cleanup()
initialTimeout := 120 * time.Second
streamTimeout := 20 * time.Second
prepareParallelHistoryModel(ctx, t, client, modelName)
// Get the server running (if applicable) warm the model up with a single initial request
slog.Info("loading", "model", modelName)
err := client.Generate(ctx,
&api.GenerateRequest{Model: modelName, KeepAlive: &api.Duration{Duration: 10 * time.Second}},
func(response api.GenerateResponse) error { return nil },
)
if err != nil {
t.Fatalf("failed to load model %s: %s", modelName, err)
}
preloadGenerateModel(ctx, t, client, api.GenerateRequest{Model: modelName, KeepAlive: &api.Duration{Duration: 10 * time.Second}})
gpuPercent := getGPUPercent(ctx, t, client, modelName)
if gpuPercent < 80 && gpuPercent > 50 {
slog.Warn("Low GPU percentage - increasing timeouts", "percent", gpuPercent)
@@ -190,7 +195,7 @@ func TestParallelGenerateWithHistory(t *testing.T) {
}
// Send generate requests with prior context and ensure the response is coherant and expected
func TestGenerateWithHistory(t *testing.T) {
func runGenerateWithHistory(t *testing.T) {
if testModel != "" {
// The Generate API's Context field (token array continuation) is not
// supported by all runners (e.g. MLX). Chat history works; this is
@@ -212,16 +217,10 @@ func TestGenerateWithHistory(t *testing.T) {
defer cancel()
client, _, cleanup := InitServerConnection(ctx, t)
defer cleanup()
pullOrSkip(ctx, t, client, req.Model)
// Get the server running (if applicable) warm the model up with a single initial request
slog.Info("loading", "model", req.Model)
err := client.Generate(ctx,
&api.GenerateRequest{Model: req.Model, KeepAlive: &api.Duration{Duration: 10 * time.Second}, Options: req.Options},
func(response api.GenerateResponse) error { return nil },
)
if err != nil {
t.Fatalf("failed to load model %s: %s", req.Model, err)
}
preloadGenerateModel(ctx, t, client, api.GenerateRequest{Model: req.Model, KeepAlive: &api.Duration{Duration: 10 * time.Second}, Options: req.Options})
req.Context = DoGenerate(ctx, t, client, req, rainbowExpected, 30*time.Second, 20*time.Second)
@@ -236,11 +235,7 @@ func TestGenerateWithHistory(t *testing.T) {
}
// Send multiple chat requests with prior context and ensure the response is coherant and expected
func TestParallelChatWithHistory(t *testing.T) {
if testModel != "" {
t.Skip("uses hardcoded model, not applicable with model override")
}
modelName := "gpt-oss:20b"
func runParallelChatWithHistory(t *testing.T, modelName string) {
req, resp := ChatRequests()
numParallel := 2
iterLimit := 2
@@ -252,16 +247,10 @@ func TestParallelChatWithHistory(t *testing.T) {
defer cleanup()
initialTimeout := 120 * time.Second
streamTimeout := 20 * time.Second
prepareParallelHistoryModel(ctx, t, client, modelName)
// Get the server running (if applicable) warm the model up with a single initial empty request
slog.Info("loading", "model", modelName)
err := client.Generate(ctx,
&api.GenerateRequest{Model: modelName, KeepAlive: &api.Duration{Duration: 10 * time.Second}},
func(response api.GenerateResponse) error { return nil },
)
if err != nil {
t.Fatalf("failed to load model %s: %s", modelName, err)
}
preloadGenerateModel(ctx, t, client, api.GenerateRequest{Model: modelName, KeepAlive: &api.Duration{Duration: 10 * time.Second}})
gpuPercent := getGPUPercent(ctx, t, client, modelName)
if gpuPercent < 80 && gpuPercent > 50 {
slog.Warn("Low GPU percentage - increasing timeouts", "percent", gpuPercent)
@@ -302,8 +291,16 @@ func TestParallelChatWithHistory(t *testing.T) {
wg.Wait()
}
func prepareParallelHistoryModel(ctx context.Context, t *testing.T, client *api.Client, modelName string) {
t.Helper()
skipRegisteredMinVRAM(t, modelName)
requireCapability(ctx, t, client, modelName, "completion")
skipIfTargetArchitecture(ctx, t, client, modelName)
skipIfModelTooLargeForSweepVRAM(ctx, t, client, modelName)
}
// Send generate requests with prior context and ensure the response is coherant and expected
func TestChatWithHistory(t *testing.T) {
func runChatWithHistory(t *testing.T) {
req := api.ChatRequest{
Model: smol,
Stream: &stream,
@@ -324,16 +321,10 @@ func TestChatWithHistory(t *testing.T) {
defer cancel()
client, _, cleanup := InitServerConnection(ctx, t)
defer cleanup()
pullOrSkip(ctx, t, client, req.Model)
// Get the server running (if applicable) warm the model up with a single initial request
slog.Info("loading", "model", req.Model)
err := client.Generate(ctx,
&api.GenerateRequest{Model: req.Model, KeepAlive: &api.Duration{Duration: 10 * time.Second}, Options: req.Options},
func(response api.GenerateResponse) error { return nil },
)
if err != nil {
t.Fatalf("failed to load model %s: %s", req.Model, err)
}
preloadGenerateModel(ctx, t, client, api.GenerateRequest{Model: req.Model, KeepAlive: &api.Duration{Duration: 10 * time.Second}, Options: req.Options})
assistant := DoChat(ctx, t, client, req, rainbowExpected, 30*time.Second, 20*time.Second)

View File

@@ -136,7 +136,7 @@ func runOllamaCreate(ctx context.Context, t *testing.T, args ...string) {
}
}
func TestCreateSafetensorsLLM(t *testing.T) {
func runCreateSafetensorsLLM(t *testing.T) {
if testModel != "" {
t.Skip("exercises create pipeline with a fixed source model, not applicable with model override")
}
@@ -214,7 +214,7 @@ func TestCreateSafetensorsLLM(t *testing.T) {
}
}
func TestCreateGGUF(t *testing.T) {
func runCreateGGUF(t *testing.T) {
if testModel != "" {
t.Skip("exercises create pipeline with a fixed source model, not applicable with model override")
}

View File

@@ -0,0 +1,204 @@
//go:build integration
package integration
import (
"context"
"encoding/json"
"fmt"
"os"
"path/filepath"
"strings"
"testing"
"time"
"github.com/ollama/ollama/api"
)
func registerEmbeddingCases(models []string) {
registerEmbeddingCasesWithFallback(models, false)
}
func registerLibraryEmbeddingCases(models []string) {
registerEmbeddingCasesWithFallback(models, true)
}
func registerEmbeddingCasesWithFallback(models []string, smokeMissing bool) {
testCases, err := loadEmbeddingTestCases()
if err != nil {
registerIntegrationCases(integrationCase{
Key: "embed/testdata",
Case: "embed",
Model: "testdata",
Run: func(t *testing.T) {
t.Fatalf("failed to load embedding test data: %s", err)
},
})
return
}
if testModel != "" {
models = []string{testModel}
}
cases := make([]integrationCase, 0, len(models))
for _, model := range models {
model := model
expected, ok := embeddingExpected(testCases, model)
if !ok {
if smokeMissing || testModel != "" {
cases = append(cases, embeddingSmokeCase(model))
continue
}
cases = append(cases, integrationCase{
Key: "embed/" + model,
Case: "embed",
Model: model,
Run: func(t *testing.T) {
t.Skipf("no embedding expectation for model %s", model)
},
})
continue
}
cases = append(cases, embeddingCase(model, expected))
}
registerIntegrationCases(cases...)
}
func embeddingSmokeCase(model string) integrationCase {
return integrationCase{
Key: "embed/" + model,
Case: "embed",
Model: model,
Run: func(t *testing.T) {
runEmbeddingSmokeModel(t, model)
},
}
}
func embeddingCase(model string, expected []float64) integrationCase {
return integrationCase{
Key: "embed/" + model,
Case: "embed",
Model: model,
Run: func(t *testing.T) {
runEmbeddingModel(t, model, expected)
},
}
}
func loadEmbeddingTestCases() (map[string][]float64, error) {
data, err := os.ReadFile(filepath.Join("testdata", "embed.json"))
if err != nil {
return nil, err
}
testCases := map[string][]float64{}
if err := json.Unmarshal(data, &testCases); err != nil {
return nil, err
}
return testCases, nil
}
func embeddingExpected(testCases map[string][]float64, model string) ([]float64, bool) {
if expected, ok := testCases[model]; ok {
return expected, true
}
if !strings.Contains(model, ":") {
expected, ok := testCases[model+":latest"]
return expected, ok
}
return nil, false
}
func runEmbeddingModel(t *testing.T, model string, expected []float64) {
t.Helper()
softTimeout, hardTimeout := getTimeouts(t)
ctx, cancel := context.WithTimeout(context.Background(), hardTimeout)
defer cancel()
client, _, cleanup := InitServerConnection(ctx, t)
defer cleanup()
if time.Since(started) > softTimeout {
t.Skip("skipping remaining tests to avoid excessive runtime")
}
pullOrSkip(ctx, t, client, model)
skipIfModelTooLargeForSweepVRAM(ctx, t, client, model)
req := api.EmbeddingRequest{
Model: model,
Prompt: "why is the sky blue?",
KeepAlive: &api.Duration{Duration: 10 * time.Second},
Options: map[string]any{
"temperature": 0,
"seed": 123,
},
}
resp, err := client.Embeddings(ctx, &req)
if err != nil {
t.Fatalf("embeddings call failed %s", err)
}
defer func() {
client.Generate(ctx, &api.GenerateRequest{Model: req.Model, KeepAlive: &api.Duration{Duration: 0}}, func(rsp api.GenerateResponse) error { return nil })
}()
if len(resp.Embedding) == 0 {
t.Errorf("zero length embedding response")
}
if len(expected) != len(resp.Embedding) {
expStr := make([]string, len(resp.Embedding))
for i, v := range resp.Embedding {
expStr[i] = fmt.Sprintf("%0.6f", v)
}
// When adding new models, use this output to populate the testdata/embed.json
fmt.Printf("expected\n%s\n", strings.Join(expStr, ", "))
t.Fatalf("expected %d, got %d", len(expected), len(resp.Embedding))
}
sim := cosineSimilarity(resp.Embedding, expected)
if sim < 0.99 {
t.Fatalf("expected %v, got %v (similarity: %f)", expected[0:5], resp.Embedding[0:5], sim)
}
}
func runEmbeddingSmokeModel(t *testing.T, model string) {
t.Helper()
softTimeout, hardTimeout := getTimeouts(t)
ctx, cancel := context.WithTimeout(context.Background(), hardTimeout)
defer cancel()
client, _, cleanup := InitServerConnection(ctx, t)
defer cleanup()
if time.Since(started) > softTimeout {
t.Skip("skipping remaining tests to avoid excessive runtime")
}
requireCapability(ctx, t, client, model, "embedding")
skipIfModelTooLargeForSweepVRAM(ctx, t, client, model)
req := api.EmbedRequest{
Model: model,
Input: []string{"cat", "kitten", "dog"},
KeepAlive: &api.Duration{Duration: 10 * time.Second},
}
resp, err := embedTestHelper(ctx, client, t, req)
if err != nil {
t.Fatal(err)
}
defer func() {
client.Generate(ctx, &api.GenerateRequest{Model: req.Model, KeepAlive: &api.Duration{Duration: 0}}, func(rsp api.GenerateResponse) error { return nil })
}()
if len(resp.Embeddings) != 3 {
t.Fatalf("expected 3 embeddings, got %d", len(resp.Embeddings))
}
for i, embedding := range resp.Embeddings {
if len(embedding) == 0 {
t.Fatalf("embedding %d was empty", i)
}
}
cosRelated := cosineSimilarity(resp.Embeddings[0], resp.Embeddings[1])
cosUnrelated := cosineSimilarity(resp.Embeddings[0], resp.Embeddings[2])
if cosRelated <= cosUnrelated {
t.Fatalf("expected related terms to be closer than unrelated terms: cat/kitten=%f cat/dog=%f", cosRelated, cosUnrelated)
}
}

View File

@@ -10,7 +10,6 @@ import (
"testing"
"time"
"github.com/google/go-cmp/cmp"
"github.com/ollama/ollama/api"
)
@@ -61,6 +60,19 @@ func requireEmbedErrorContainsAny(t *testing.T, err error, substrings ...string)
t.Fatalf("expected error containing one of %q, got: %v", substrings, err)
}
func requireSimilarEmbedding(t *testing.T, want, got []float32) {
t.Helper()
if len(got) != len(want) {
t.Fatalf("expected %d embedding floats, got %d", len(want), len(got))
}
sim := cosineSimilarity(got, want)
if sim < 0.999 {
t.Fatalf("expected embedding similar to %v, got %v (similarity: %f)", want[0:5], got[0:5], sim)
}
}
func euclideanDistance[V float32 | float64](v1, v2 []V) V {
if len(v1) != len(v2) {
return V(math.Inf(1))
@@ -88,13 +100,13 @@ func manhattanDistance[V float32 | float64](v1, v2 []V) V {
return sum
}
func TestEmbedCosineDistanceCorrelation(t *testing.T) {
func runEmbedCosineDistanceCorrelation(t *testing.T, models []string) {
ctx, cancel := context.WithTimeout(context.Background(), 4*time.Minute)
defer cancel()
client, _, cleanup := InitServerConnection(ctx, t)
defer cleanup()
for _, model := range testModels(libraryEmbedModels) {
for _, model := range testModels(models) {
t.Run(model, func(t *testing.T) {
if testModel != "" {
requireCapability(ctx, t, client, model, "embedding")
@@ -163,7 +175,7 @@ func TestEmbedCosineDistanceCorrelation(t *testing.T) {
}
}
func TestAllMiniLMEmbeddings(t *testing.T) {
func runAllMiniLMEmbeddings(t *testing.T) {
if testModel != "" {
t.Skip("uses hardcoded model, not applicable with model override")
}
@@ -196,7 +208,7 @@ func TestAllMiniLMEmbeddings(t *testing.T) {
}
}
func TestAllMiniLMEmbed(t *testing.T) {
func runAllMiniLMEmbed(t *testing.T) {
if testModel != "" {
t.Skip("uses hardcoded model, not applicable with model override")
}
@@ -236,7 +248,7 @@ func TestAllMiniLMEmbed(t *testing.T) {
}
}
func TestAllMiniLMBatchEmbed(t *testing.T) {
func runAllMiniLMBatchEmbed(t *testing.T) {
if testModel != "" {
t.Skip("uses hardcoded model, not applicable with model override")
}
@@ -286,7 +298,7 @@ func TestAllMiniLMBatchEmbed(t *testing.T) {
}
}
func TestAllMiniLMEmbedTruncate(t *testing.T) {
func runAllMiniLMEmbedTruncate(t *testing.T) {
if testModel != "" {
t.Skip("uses hardcoded model, not applicable with model override")
}
@@ -321,9 +333,7 @@ func TestAllMiniLMEmbedTruncate(t *testing.T) {
t.Fatal(err)
}
if diff := cmp.Diff(want.Embeddings[0], got.Embeddings[0]); diff != "" {
t.Errorf("embedding mismatch (-want +got):\n%s", diff)
}
requireSimilarEmbedding(t, want.Embeddings[0], got.Embeddings[0])
},
},
{
@@ -338,9 +348,7 @@ func TestAllMiniLMEmbedTruncate(t *testing.T) {
t.Fatal(err)
}
t.Logf("PromptEvalCount: want=%d got=%d", want.PromptEvalCount, got.PromptEvalCount)
if diff := cmp.Diff(want.Embeddings[0], got.Embeddings[0]); diff != "" {
t.Errorf("embedding mismatch (-want +got):\n%s", diff)
}
requireSimilarEmbedding(t, want.Embeddings[0], got.Embeddings[0])
},
},
{
@@ -356,9 +364,7 @@ func TestAllMiniLMEmbedTruncate(t *testing.T) {
t.Fatal(err)
}
t.Logf("PromptEvalCount: want=%d got=%d", want.PromptEvalCount, got.PromptEvalCount)
if diff := cmp.Diff(want.Embeddings[0], got.Embeddings[0]); diff != "" {
t.Errorf("embedding mismatch (-want +got):\n%s", diff)
}
requireSimilarEmbedding(t, want.Embeddings[0], got.Embeddings[0])
},
},
{
@@ -432,7 +438,7 @@ func embedTestHelper(ctx context.Context, client *api.Client, t *testing.T, req
return client.Embed(ctx, &req)
}
func TestEmbedTruncation(t *testing.T) {
func runEmbedTruncation(t *testing.T, models []string) {
// Use test deadline if set, otherwise default to 2 minutes
timeout := 2 * time.Minute
if deadline, ok := t.Deadline(); ok {
@@ -443,7 +449,7 @@ func TestEmbedTruncation(t *testing.T) {
client, _, cleanup := InitServerConnection(ctx, t)
defer cleanup()
for _, model := range testModels(libraryEmbedModels) {
for _, model := range testModels(models) {
model := model
t.Run(model, func(t *testing.T) {
if testModel != "" {
@@ -454,6 +460,9 @@ func TestEmbedTruncation(t *testing.T) {
t.Skip("skipping remaining tests to avoid timeout")
}
pullOrSkip(ctx, t, client, model)
skipIfModelTooLargeForSweepVRAM(ctx, t, client, model)
// Give each model its own budget to account for first-time pulls/loads
mctx, mcancel := context.WithTimeout(ctx, 3*time.Minute)
defer mcancel()
@@ -507,19 +516,22 @@ func TestEmbedTruncation(t *testing.T) {
}
}
// TestEmbedLargeInput tests that embedding models can handle large inputs that would exceed typical batch sizes.
func TestEmbedLargeInput(t *testing.T) {
// runEmbedLargeInput tests that embedding models can handle large inputs that would exceed typical batch sizes.
func runEmbedLargeInput(t *testing.T, models []string) {
ctx, cancel := context.WithTimeout(context.Background(), 3*time.Minute)
defer cancel()
client, _, cleanup := InitServerConnection(ctx, t)
defer cleanup()
for _, model := range testModels(libraryEmbedModels) {
for _, model := range testModels(models) {
model := model
t.Run(model, func(t *testing.T) {
if testModel != "" {
requireCapability(ctx, t, client, model, "embedding")
}
pullOrSkip(ctx, t, client, model)
skipIfModelTooLargeForSweepVRAM(ctx, t, client, model)
mctx, mcancel := context.WithTimeout(ctx, 2*time.Minute)
defer mcancel()
@@ -567,11 +579,11 @@ func TestEmbedLargeInput(t *testing.T) {
}
}
// TestEmbedStatusCode tests that errors from the embedding endpoint
// runEmbedStatusCode tests that errors from the embedding endpoint
// properly preserve their HTTP status codes when returned to the client.
// This test specifically checks the error handling path in EmbedHandler
// where api.StatusError errors should maintain their original status code.
func TestEmbedStatusCode(t *testing.T) {
func runEmbedStatusCode(t *testing.T, models []string) {
// Use test deadline if set, otherwise default to 2 minutes
timeout := 2 * time.Minute
if deadline, ok := t.Deadline(); ok {
@@ -582,7 +594,7 @@ func TestEmbedStatusCode(t *testing.T) {
client, _, cleanup := InitServerConnection(ctx, t)
defer cleanup()
for _, model := range testModels(libraryEmbedModels) {
for _, model := range testModels(models) {
model := model
t.Run(model, func(t *testing.T) {
if testModel != "" {
@@ -598,6 +610,7 @@ func TestEmbedStatusCode(t *testing.T) {
// Pull the model if needed
pullOrSkip(mctx, t, client, model)
skipIfModelTooLargeForSweepVRAM(mctx, t, client, model)
t.Run("truncation error status code", func(t *testing.T) {
truncFalse := false

View File

@@ -13,7 +13,7 @@ import (
"github.com/ollama/ollama/api"
)
func TestImageGeneration(t *testing.T) {
func runImageGeneration(t *testing.T) {
if testModel != "" {
t.Skip("uses hardcoded models, not applicable with model override")
}
@@ -51,6 +51,7 @@ func TestImageGeneration(t *testing.T) {
t.Logf("Generating image with prompt: %s", tc.prompt)
imageBase64, err := generateImage(ctx, client, tc.imageGenModel, tc.prompt)
if err != nil {
skipIfMLXUnsupported(t, err)
if strings.Contains(err.Error(), "image generation not available") {
t.Skip("Target system does not support image generation")
} else if strings.Contains(err.Error(), "executable file not found in") { // Windows pattern, not yet supported
@@ -79,10 +80,7 @@ func TestImageGeneration(t *testing.T) {
t.Logf("Generated image: %d bytes", len(imageData))
// Preload vision model and check GPU loading
err = client.Generate(ctx, &api.GenerateRequest{Model: tc.visionModel}, func(response api.GenerateResponse) error { return nil })
if err != nil {
t.Fatalf("failed to load vision model: %v", err)
}
preloadGenerateModel(ctx, t, client, api.GenerateRequest{Model: tc.visionModel})
// Use vision model to describe the image
chatReq := api.ChatRequest{

View File

@@ -1,72 +0,0 @@
//go:build integration && library
package integration
import (
"context"
"fmt"
"log/slog"
"os"
"testing"
"time"
"github.com/ollama/ollama/api"
)
// First run of this scenario on a target system will take a long time to download
// ~1.5TB of models. Set a sufficiently large -timeout for your network speed
func TestLibraryModelsChat(t *testing.T) {
softTimeout, hardTimeout := getTimeouts(t)
slog.Info("Setting timeouts", "soft", softTimeout, "hard", hardTimeout)
ctx, cancel := context.WithTimeout(context.Background(), hardTimeout)
defer cancel()
client, _, cleanup := InitServerConnection(ctx, t)
defer cleanup()
targetArch := os.Getenv("OLLAMA_TEST_ARCHITECTURE")
for _, model := range testModels(libraryChatModels) {
t.Run(model, func(t *testing.T) {
if time.Now().Sub(started) > softTimeout {
t.Skip("skipping remaining tests to avoid excessive runtime")
}
pullOrSkip(ctx, t, client, model)
if targetArch != "" {
resp, err := client.Show(ctx, &api.ShowRequest{Name: model})
if err != nil {
t.Fatalf("unable to show model: %s", err)
}
arch := resp.ModelInfo["general.architecture"].(string)
if arch != targetArch {
t.Skip(fmt.Sprintf("Skipping %s architecture %s != %s", model, arch, targetArch))
}
}
req := api.ChatRequest{
Model: model,
Messages: []api.Message{
{
Role: "user",
Content: blueSkyPrompt,
},
},
KeepAlive: &api.Duration{Duration: 10 * time.Second},
Options: map[string]interface{}{
"temperature": 0.1,
"seed": 123,
},
}
anyResp := blueSkyExpected
// Special cases
if model == "duckdb-nsql" {
anyResp = []string{"select", "from"}
} else if model == "granite3-guardian" || model == "shieldgemma" || model == "llama-guard3" || model == "bespoke-minicheck" {
anyResp = []string{"yes", "no", "safe", "unsafe"}
} else if model == "openthinker" {
anyResp = []string{"plugin", "im_sep", "components", "function call"}
} else if model == "starcoder" || model == "starcoder2" || model == "magicoder" || model == "deepseek-coder" {
req.Messages[0].Content = "def fibonacci():"
anyResp = []string{"f(n)", "sequence", "n-1", "main()", "__main__", "while"}
}
DoChat(ctx, t, client, req, anyResp, 120*time.Second, 30*time.Second)
})
}
}

View File

@@ -11,69 +11,53 @@ import (
"github.com/ollama/ollama/api"
)
func TestVisionModels(t *testing.T) {
skipUnderMinVRAM(t, 6)
defaultVisionModels := []string{
"gemma4",
"qwen2.5vl",
// "llama3.2-vision", // TODO: re-enable when llama.cpp supports mllama.
"gemma3",
"qwen3-vl:8b",
"qwen3-vl:30b",
"ministral-3",
}
skipIfNoVisionOverride(t)
for _, model := range testModels(defaultVisionModels) {
t.Run(model, func(t *testing.T) {
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Minute)
defer cancel()
client, _, cleanup := InitServerConnection(ctx, t)
defer cleanup()
requireCapability(ctx, t, client, model, "vision")
pullOrSkip(ctx, t, client, model)
image, err := base64.StdEncoding.DecodeString(imageEncoding)
if err != nil {
t.Fatal(err)
}
req := api.ChatRequest{
Model: model,
Messages: []api.Message{
{
Role: "user",
Content: "what does the text in this image say?",
Images: []api.ImageData{
image,
},
},
},
Stream: &stream,
Options: map[string]any{
"seed": 42,
"temperature": 0.0,
},
KeepAlive: &api.Duration{Duration: 10 * time.Second},
}
// Preload to skip if we're less than 80% on GPU to avoid extremely slow tests
err = client.Generate(ctx, &api.GenerateRequest{Model: req.Model}, func(response api.GenerateResponse) error { return nil })
if err != nil {
t.Fatalf("failed to load model %s: %s", req.Model, err)
}
skipIfNotGPULoaded(ctx, t, client, req.Model, 80)
// Note: sometimes it returns "the ollamas" sometimes "the ollams"
// llava models on CPU can be quite slow to start
DoChat(ctx, t, client, req, []string{"the ollam"}, 240*time.Second, 30*time.Second)
})
}
func registerVisionTextCases(models []string) {
registerModelIntegrationCases("vision-text", models, runVisionTextModel)
}
func TestIntegrationSplitBatch(t *testing.T) {
func runVisionTextModel(t *testing.T, model string) {
t.Helper()
skipUnderMinVRAM(t, 6)
skipKnownIntegrationFlake(t, "vision-text", model)
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Minute)
defer cancel()
client, _, cleanup := InitServerConnection(ctx, t)
defer cleanup()
requireCapability(ctx, t, client, model, "vision")
pullOrSkip(ctx, t, client, model)
image, err := base64.StdEncoding.DecodeString(imageEncoding)
if err != nil {
t.Fatal(err)
}
req := api.ChatRequest{
Model: model,
Messages: []api.Message{
{
Role: "user",
Content: "what does the text in this image say?",
Images: []api.ImageData{
image,
},
},
},
Stream: &stream,
Options: map[string]any{
"seed": 42,
"temperature": 0.0,
},
KeepAlive: &api.Duration{Duration: 10 * time.Second},
}
// Preload to skip if we're less than 80% on GPU to avoid extremely slow tests
preloadGenerateModel(ctx, t, client, api.GenerateRequest{Model: req.Model})
skipIfNotGPULoaded(ctx, t, client, req.Model, 80)
DoChat(ctx, t, client, req, []string{"the ollam", "ollamas"}, 240*time.Second, 30*time.Second)
}
func runIntegrationSplitBatch(t *testing.T, model string) {
if testModel != "" {
t.Skip("uses hardcoded model, not applicable with model override")
}
@@ -83,7 +67,7 @@ func TestIntegrationSplitBatch(t *testing.T) {
t.Fatal(err)
}
req := api.GenerateRequest{
Model: "gemma3:4b",
Model: model,
// Fill up a chunk of the batch so the image will partially spill over into the next one
System: "Lorem ipsum dolor sit amet, consectetur adipiscing elit. Sed aliquet, justo in malesuada lobortis, odio ligula volutpat quam, quis faucibus ipsum magna quis sapien. Aliquam in venenatis diam, eu viverra magna. Phasellus imperdiet hendrerit volutpat. Vivamus sem ex, facilisis placerat felis non, dictum elementum est. Phasellus aliquam imperdiet lacus, eget placerat ligula sodales vel. Pellentesque nec auctor mi. Curabitur arcu nisi, faucibus eget nunc id, viverra interdum mi. Curabitur ornare ipsum ex, ac euismod ex aliquam in. Vestibulum id magna at purus accumsan fermentum. Proin scelerisque posuere nunc quis interdum. Maecenas sed mollis nisl. Etiam vitae ipsum interdum, placerat est quis, tincidunt velit. Nullam tempor nibh non lorem volutpat efficitur. Cras laoreet diam imperdiet ipsum auctor bibendum. Suspendisse ultrices urna sed metus sagittis suscipit. Quisque ullamcorper aliquam nibh ut mollis. Aenean dapibus mauris pharetra, venenatis elit ac, hendrerit odio. Cras vestibulum erat tempor, lobortis justo eu, lobortis ipsum. Nam laoreet dapibus sem. Proin vel diam ultrices, elementum ante et, ornare lectus. Proin eu accumsan nisl. Praesent ac ex vitae ipsum vulputate tristique facilisis sit amet lacus. Nullam faucibus magna a pellentesque pretium. Nunc lacinia ullamcorper sollicitudin. Donec vitae accumsan turpis, sed porttitor est. Donec porttitor mi vitae augue faucibus, vel mollis diam tincidunt.",
Prompt: "what does the text in this image say?",

View File

@@ -16,7 +16,7 @@ import (
"github.com/ollama/ollama/api"
)
func TestMaxQueue(t *testing.T) {
func runMaxQueue(t *testing.T) {
t.Skip("this test needs to be re-evaluated to use a proper embedding model")
if os.Getenv("OLLAMA_TEST_EXISTING") != "" {

View File

@@ -1,187 +0,0 @@
//go:build integration && models
package integration
import (
"context"
"encoding/json"
"fmt"
"io/ioutil"
"log/slog"
"os"
"path/filepath"
"strconv"
"strings"
"testing"
"time"
"github.com/ollama/ollama/api"
"github.com/ollama/ollama/format"
)
func TestModelsChat(t *testing.T) {
softTimeout, hardTimeout := getTimeouts(t)
slog.Info("Setting timeouts", "soft", softTimeout, "hard", hardTimeout)
ctx, cancel := context.WithTimeout(context.Background(), hardTimeout)
defer cancel()
client, _, cleanup := InitServerConnection(ctx, t)
defer cleanup()
// TODO use info API eventually
var maxVram uint64
var err error
if s := os.Getenv("OLLAMA_MAX_VRAM"); s != "" {
maxVram, err = strconv.ParseUint(s, 10, 64)
if err != nil {
t.Fatalf("invalid OLLAMA_MAX_VRAM %v", err)
}
} else {
slog.Warn("No VRAM info available, testing all models, so larger ones might timeout...")
}
chatModels := append(ollamaEngineChatModels, llamaRunnerChatModels...)
chatModels = append(chatModels, mlxEngineChatModels...)
for _, model := range testModels(chatModels) {
t.Run(model, func(t *testing.T) {
if time.Now().Sub(started) > softTimeout {
t.Skip("skipping remaining tests to avoid excessive runtime")
}
pullOrSkip(ctx, t, client, model)
if maxVram > 0 {
resp, err := client.List(ctx)
if err != nil {
t.Fatalf("list models failed %v", err)
}
for _, m := range resp.Models {
if m.Name == model && float32(m.Size)*1.2 > float32(maxVram) {
t.Skipf("model %s is too large for available VRAM: %s > %s", model, format.HumanBytes(m.Size), format.HumanBytes(int64(maxVram)))
}
}
}
initialTimeout := 120 * time.Second
streamTimeout := 30 * time.Second
slog.Info("loading", "model", model)
err := client.Generate(ctx,
&api.GenerateRequest{Model: model, KeepAlive: &api.Duration{Duration: 10 * time.Second}},
func(response api.GenerateResponse) error { return nil },
)
if err != nil {
skipIfMLXUnsupported(t, err)
t.Fatalf("failed to load model %s: %s", model, err)
}
gpuPercent := getGPUPercent(ctx, t, client, model)
if gpuPercent < 80 {
slog.Warn("Low GPU percentage - increasing timeouts", "percent", gpuPercent)
initialTimeout = 240 * time.Second
streamTimeout = 40 * time.Second
}
// TODO - fiddle with context size
req := api.ChatRequest{
Model: model,
Messages: []api.Message{
{
Role: "user",
Content: blueSkyPrompt,
},
},
KeepAlive: &api.Duration{Duration: 10 * time.Second},
Options: map[string]interface{}{
"temperature": 0,
"seed": 123,
},
}
DoChat(ctx, t, client, req, blueSkyExpected, initialTimeout, streamTimeout)
// best effort unload once we're done with the model
client.Generate(ctx, &api.GenerateRequest{Model: req.Model, KeepAlive: &api.Duration{Duration: 0}}, func(rsp api.GenerateResponse) error { return nil })
})
}
}
func TestModelsEmbed(t *testing.T) {
softTimeout, hardTimeout := getTimeouts(t)
ctx, cancel := context.WithTimeout(context.Background(), hardTimeout)
defer cancel()
client, _, cleanup := InitServerConnection(ctx, t)
defer cleanup()
// TODO use info API eventually
var maxVram uint64
var err error
if s := os.Getenv("OLLAMA_MAX_VRAM"); s != "" {
maxVram, err = strconv.ParseUint(s, 10, 64)
if err != nil {
t.Fatalf("invalid OLLAMA_MAX_VRAM %v", err)
}
} else {
slog.Warn("No VRAM info available, testing all models, so larger ones might timeout...")
}
data, err := ioutil.ReadFile(filepath.Join("testdata", "embed.json"))
if err != nil {
t.Fatalf("failed to open test data file: %s", err)
}
testCase := map[string][]float64{}
err = json.Unmarshal(data, &testCase)
if err != nil {
t.Fatalf("failed to load test data: %s", err)
}
for model, expected := range testCase {
if testModel != "" && model != testModel {
continue
}
t.Run(model, func(t *testing.T) {
if time.Now().Sub(started) > softTimeout {
t.Skip("skipping remaining tests to avoid excessive runtime")
}
pullOrSkip(ctx, t, client, model)
if maxVram > 0 {
resp, err := client.List(ctx)
if err != nil {
t.Fatalf("list models failed %v", err)
}
for _, m := range resp.Models {
if m.Name == model && float32(m.Size)*1.2 > float32(maxVram) {
t.Skipf("model %s is too large for available VRAM: %s > %s", model, format.HumanBytes(m.Size), format.HumanBytes(int64(maxVram)))
}
}
}
req := api.EmbeddingRequest{
Model: model,
Prompt: "why is the sky blue?",
KeepAlive: &api.Duration{Duration: 10 * time.Second},
Options: map[string]interface{}{
"temperature": 0,
"seed": 123,
},
}
resp, err := client.Embeddings(ctx, &req)
if err != nil {
t.Fatalf("embeddings call failed %s", err)
}
defer func() {
// best effort unload once we're done with the model
client.Generate(ctx, &api.GenerateRequest{Model: req.Model, KeepAlive: &api.Duration{Duration: 0}}, func(rsp api.GenerateResponse) error { return nil })
}()
if len(resp.Embedding) == 0 {
t.Errorf("zero length embedding response")
}
if len(expected) != len(resp.Embedding) {
expStr := make([]string, len(resp.Embedding))
for i, v := range resp.Embedding {
expStr[i] = fmt.Sprintf("%0.6f", v)
}
// When adding new models, use this output to populate the testdata/embed.json
fmt.Printf("expected\n%s\n", strings.Join(expStr, ", "))
t.Fatalf("expected %d, got %d", len(expected), len(resp.Embedding))
}
sim := cosineSimilarity(resp.Embedding, expected)
if sim < 0.99 {
t.Fatalf("expected %v, got %v (similarity: %f)", expected[0:5], resp.Embedding[0:5], sim)
}
})
}
}

View File

@@ -1,278 +0,0 @@
//go:build integration && perf
package integration
import (
"context"
"fmt"
"io/ioutil"
"log/slog"
"math"
"os"
"path/filepath"
"strconv"
"strings"
"testing"
"time"
"github.com/ollama/ollama/api"
"github.com/ollama/ollama/format"
)
var (
// Models that don't work reliably with the large context prompt in this test case
longContextFlakes = []string{
"granite-code:latest",
"nemotron-mini:latest",
"falcon:latest", // 2k model
"falcon2:latest", // 2k model
"minicpm-v:latest",
"qwen:latest",
}
)
// Note: this test case can take a long time to run, particularly on models with
// large contexts. Run with -timeout set to a large value to get reasonable coverage
// Example usage:
//
// go test --tags=integration,perf -count 1 ./integration -v -timeout 90m -run TestModelsPerf 2>&1 | tee int.log
// cat int.log | grep MODEL_PERF_HEADER | head -1| cut -f2- -d: > perf.csv
// cat int.log | grep MODEL_PERF_DATA | cut -f2- -d: >> perf.csv
func TestModelsPerf(t *testing.T) {
doModelPerfTest(t, append(ollamaEngineChatModels, llamaRunnerChatModels...))
}
func TestLibraryModelsPerf(t *testing.T) {
doModelPerfTest(t, libraryChatModels)
}
func doModelPerfTest(t *testing.T, chatModels []string) {
softTimeout, hardTimeout := getTimeouts(t)
slog.Info("Setting timeouts", "soft", softTimeout, "hard", hardTimeout)
ctx, cancel := context.WithTimeout(context.Background(), hardTimeout)
defer cancel()
client, _, cleanup := InitServerConnection(ctx, t)
defer cleanup()
// TODO use info API eventually
var maxVram uint64
var err error
if s := os.Getenv("OLLAMA_MAX_VRAM"); s != "" {
maxVram, err = strconv.ParseUint(s, 10, 64)
if err != nil {
t.Fatalf("invalid OLLAMA_MAX_VRAM %v", err)
}
} else {
slog.Warn("No VRAM info available, testing all models, so larger ones might timeout...")
}
data, err := ioutil.ReadFile(filepath.Join("testdata", "shakespeare.txt"))
if err != nil {
t.Fatalf("failed to open test data file: %s", err)
}
longPrompt := "summarize the following: " + string(data)
targetArch := os.Getenv("OLLAMA_TEST_ARCHITECTURE")
for _, model := range chatModels {
if !strings.Contains(model, ":") {
model = model + ":latest"
}
t.Run(model, func(t *testing.T) {
if time.Now().Sub(started) > softTimeout {
t.Skip("skipping remaining tests to avoid excessive runtime")
}
pullOrSkip(ctx, t, client, model)
var maxContext int
resp, err := client.Show(ctx, &api.ShowRequest{Model: model})
if err != nil {
t.Fatalf("show failed: %s", err)
}
arch := resp.ModelInfo["general.architecture"].(string)
maxContext = int(resp.ModelInfo[fmt.Sprintf("%s.context_length", arch)].(float64))
if targetArch != "" && arch != targetArch {
t.Skip(fmt.Sprintf("Skipping %s architecture %s != %s", model, arch, targetArch))
}
if maxVram > 0 {
resp, err := client.List(ctx)
if err != nil {
t.Fatalf("list models failed %v", err)
}
for _, m := range resp.Models {
// For these tests we want to exercise a some amount of overflow on the CPU
if m.Name == model && float32(m.Size)*0.75 > float32(maxVram) {
t.Skipf("model %s is too large %s for available VRAM %s", model, format.HumanBytes(m.Size), format.HumanBytes(int64(maxVram)))
}
}
}
slog.Info("scneario", "model", model, "max_context", maxContext)
loaded := false
defer func() {
// best effort unload once we're done with the model
if loaded {
client.Generate(ctx, &api.GenerateRequest{Model: model, KeepAlive: &api.Duration{Duration: 0}}, func(rsp api.GenerateResponse) error { return nil })
}
}()
// Some models don't handle the long context data well so skip them to avoid flaky test results
longContextFlake := false
for _, flake := range longContextFlakes {
if model == flake {
longContextFlake = true
break
}
}
// iterate through a few context sizes for coverage without excessive runtime
var contexts []int
keepGoing := true
if maxContext > 16384 {
contexts = []int{4096, 8192, 16384, maxContext}
} else if maxContext > 8192 {
contexts = []int{4096, 8192, maxContext}
} else if maxContext > 4096 {
contexts = []int{4096, maxContext}
} else if maxContext > 0 {
contexts = []int{maxContext}
} else {
t.Fatal("unknown max context size")
}
for _, numCtx := range contexts {
if !keepGoing && numCtx > 8192 { // Always try up to 8k before bailing out
break
}
skipLongPrompt := false
// Workaround bug 11172 temporarily...
maxPrompt := longPrompt
// If we fill the context too full with the prompt, many models
// quickly hit context shifting and go bad.
if len(maxPrompt) > numCtx*2 { // typically yields ~1/2 full context
maxPrompt = maxPrompt[:numCtx*2]
}
testCases := []struct {
prompt string
anyResp []string
}{
{blueSkyPrompt, blueSkyExpected},
{maxPrompt, []string{"shakespeare", "oppression", "sorrows", "gutenberg", "child", "license", "sonnet", "melancholy", "love", "sorrow", "beauty"}},
}
var gpuPercent int
for _, tc := range testCases {
if len(tc.prompt) > 100 && (longContextFlake || skipLongPrompt) {
slog.Info("skipping long prompt", "model", model, "num_ctx", numCtx, "gpu_percent", gpuPercent)
continue
}
req := api.ChatRequest{
Model: model,
Messages: []api.Message{
{
Role: "user",
Content: tc.prompt,
},
},
KeepAlive: &api.Duration{Duration: 20 * time.Second}, // long enough to ensure a ps returns
Options: map[string]interface{}{
"temperature": 0,
"seed": 123,
"num_ctx": numCtx,
},
}
atLeastOne := false
var resp api.ChatResponse
stream := false
req.Stream = &stream
// Avoid potentially getting stuck indefinitely
limit := 5 * time.Minute
genCtx, cancel := context.WithDeadlineCause(
ctx,
time.Now().Add(limit),
fmt.Errorf("generate on model %s with ctx %d took longer than %v", model, numCtx, limit),
)
defer cancel()
err = client.Chat(genCtx, &req, func(rsp api.ChatResponse) error {
resp = rsp
return nil
})
if err != nil {
// Avoid excessive test runs, but don't consider a failure with massive context
if numCtx > 16384 && strings.Contains(err.Error(), "took longer") {
slog.Warn("max context was taking too long, skipping", "error", err)
keepGoing = false
skipLongPrompt = true
continue
}
t.Fatalf("generate error: ctx:%d err:%s", numCtx, err)
}
loaded = true
for _, expResp := range tc.anyResp {
if strings.Contains(strings.ToLower(resp.Message.Content), expResp) {
atLeastOne = true
break
}
}
if !atLeastOne {
t.Fatalf("response didn't contain expected values: ctx:%d expected:%v response:%s ", numCtx, tc.anyResp, resp.Message.Content)
}
models, err := client.ListRunning(ctx)
if err != nil {
slog.Warn("failed to list running models", "error", err)
continue
}
if len(models.Models) > 1 {
slog.Warn("multiple models loaded, may impact performance results", "loaded", models.Models)
}
for _, m := range models.Models {
if m.Name == model {
if m.SizeVRAM == 0 {
slog.Info("Model fully loaded into CPU")
gpuPercent = 0
keepGoing = false
skipLongPrompt = true
} else if m.SizeVRAM == m.Size {
slog.Info("Model fully loaded into GPU")
gpuPercent = 100
} else {
sizeCPU := m.Size - m.SizeVRAM
cpuPercent := math.Round(float64(sizeCPU) / float64(m.Size) * 100)
gpuPercent = int(100 - cpuPercent)
slog.Info("Model split between CPU/GPU", "CPU", cpuPercent, "GPU", gpuPercent)
keepGoing = false
// Heuristic to avoid excessive test run time
if gpuPercent < 90 {
skipLongPrompt = true
}
}
}
}
// Round the logged prompt count for comparisons across versions/configurations which can vary slightly
fmt.Fprintf(os.Stderr, "MODEL_PERF_HEADER:%s,%s,%s,%s,%s,%s,%s\n",
"MODEL",
"CONTEXT",
"GPU PERCENT",
"APPROX PROMPT COUNT",
"LOAD TIME",
"PROMPT EVAL TPS",
"EVAL TPS",
)
fmt.Fprintf(os.Stderr, "MODEL_PERF_DATA:%s,%d,%d,%d,%0.2f,%0.2f,%0.2f\n",
model,
numCtx,
gpuPercent,
(resp.PromptEvalCount/10)*10,
float64(resp.LoadDuration)/1000000000.0,
float64(resp.PromptEvalCount)/(float64(resp.PromptEvalDuration)/1000000000.0),
float64(resp.EvalCount)/(float64(resp.EvalDuration)/1000000000.0),
)
}
}
})
}
}

View File

@@ -1,4 +1,4 @@
//go:build integration && models
//go:build integration && release
package integration
@@ -14,7 +14,11 @@ import (
"github.com/ollama/ollama/api"
)
func TestQuantization(t *testing.T) {
func runQuantization(t *testing.T) {
if testModel != "" {
t.Skip("exercises quantization with a fixed source model, not applicable with model override")
}
sourceModels := []string{
"qwen2.5:0.5b-instruct-fp16",
}

View File

@@ -0,0 +1,48 @@
//go:build integration && fast
package integration
var (
fastNumPredictModel = "llama3.2:1b"
fastChatModels = []integrationModel{
{Name: "gemma4", MinVRAMGB: 8},
{Name: "gemma4:12b", MinVRAMGB: 16},
{Name: "qwen3.5:2b-nvfp4", MinVRAMGB: 4},
}
fastEmbedModels = []string{"qwen3-embedding"}
fastVisionTextModels = []string{"gemma4"}
fastToolsModels = []string{"qwen3.5:2b"}
fastToolsStressModels = []string{"lfm2.5"}
fastAudioModels = []string{"gemma4:e2b"}
)
func init() {
// API/basic/context/concurrency smoke cases
registerIntegrationCases(
integrationTestCase("api-generate", smol, runAPIGenerate),
integrationTestCase("api-chat", smol, runAPIChat),
integrationTestCase("api-list-models", "", runAPIListModels),
integrationTestCase("api-show-model", "llama3.2", runAPIShowModel),
integrationTestCase("generate-logprobs", smol, runAPIGenerateLogprobs),
integrationTestCase("chat-logprobs", smol, runAPIChatLogprobs),
integrationTestCase("blue-sky", smol, runBlueSky),
integrationTestCase("thinking-enabled", smol, runThinkingEnabled),
integrationTestCase("thinking-suppressed", smol, runThinkingSuppressed),
integrationModelTestCase("num-predict", fastNumPredictModel, runNumPredict),
integrationTestCase("embedding-api", "all-minilm", runAllMiniLMEmbeddings),
integrationTestCase("embed-api-truncate", "all-minilm", runAllMiniLMEmbedTruncate),
integrationTestCase("context-long-input", smol, runLongInputContext),
integrationTestCase("context-exhaustion", smol, runContextExhaustion),
integrationTestCase("generate-history", smol, runGenerateWithHistory),
integrationTestCase("concurrent-chat", smol, runConcurrentChat),
)
// Model-parametric cases
registerModelMinVRAM(fastChatModels)
registerChatCases(testModels(modelNames(fastChatModels)))
registerEmbeddingCases(testModels(fastEmbedModels))
registerVisionTextCases(testModels(fastVisionTextModels))
registerToolCases(testModels(fastToolsModels))
registerToolStressCases(testModels(fastToolsStressModels))
registerAudioTranscriptionCases(testModels(fastAudioModels))
}

View File

@@ -0,0 +1,143 @@
//go:build integration && (fast || release || library)
package integration
import (
"strings"
"testing"
)
func runIntegrationGroup(t *testing.T, cases ...string) {
t.Helper()
selected := map[string]struct{}{}
for _, c := range cases {
selected[c] = struct{}{}
}
var ran bool
for _, c := range integrationCases {
if _, ok := selected[c.Case]; !ok {
continue
}
ran = true
c := c
name := c.Case
if c.Model != "" {
name += "/" + testName(c.Model)
}
t.Run(name, c.Run)
}
if !ran {
t.Skip("no integration cases selected")
}
}
func TestAPI(t *testing.T) {
runIntegrationGroup(t,
"api-generate",
"api-chat",
"api-list-models",
"api-show-model",
"generate-logprobs",
"chat-logprobs",
)
}
func TestBasic(t *testing.T) {
runIntegrationGroup(t,
"blue-sky",
"unicode-input",
"unicode-output",
"unicode-model-dir",
"num-predict",
"thinking-enabled",
"thinking-suppressed",
)
}
func TestChat(t *testing.T) {
runIntegrationGroup(t,
"chat",
"chat-history",
)
}
func TestEmbedding(t *testing.T) {
runIntegrationGroup(t,
"embed",
"embed-correlation",
"embedding-api",
"embed-api",
"embed-api-batch",
"embed-api-truncate",
"embed-truncation",
"embed-large-input",
"embed-status-code",
)
}
func TestVision(t *testing.T) {
runIntegrationGroup(t,
"vision-multiturn",
"vision-count",
"vision-scene",
"vision-spatial",
"vision-detail",
"vision-multi-image",
"vision-description",
"vision-split-batch",
"vision-text",
)
}
func TestAudio(t *testing.T) {
runIntegrationGroup(t,
"audio-transcription",
"audio-response",
"openai-audio-transcription",
"openai-chat-audio",
)
}
func TestContext(t *testing.T) {
runIntegrationGroup(t,
"context-long-input",
"context-exhaustion",
"generate-history",
"parallel-generate-history",
"parallel-chat-history",
)
}
func TestConcurrency(t *testing.T) {
runIntegrationGroup(t,
"concurrent-chat",
"scheduler-multimodel",
"scheduler-max-queue",
)
}
func TestTools(t *testing.T) {
runIntegrationGroup(t,
"tools",
"tools-stress",
)
}
func TestCreate(t *testing.T) {
runIntegrationGroup(t,
"create-safetensors",
"create-gguf",
)
}
func TestQuantization(t *testing.T) {
runIntegrationGroup(t, "quantization")
}
func TestImageGeneration(t *testing.T) {
runIntegrationGroup(t, "image-generation")
}
func testName(s string) string {
return strings.NewReplacer("/", "~", " ", "_").Replace(s)
}

View File

@@ -0,0 +1,231 @@
//go:build integration && library
package integration
// Broad public library inventory, roughly newest to oldest from
// https://ollama.com/library?sort=newest. Cloud-only models are omitted;
// models with only very large local tags are kept commented in place.
var libraryModels = []string{
"lfm2.5",
"mistral-medium-3.5",
"granite4.1",
"nemotron3",
"laguna-xs.2",
"qwen3.6",
"medgemma1.5",
"medgemma",
"nemotron-cascade-2",
"gemma4",
"lfm2",
"nemotron-3-super",
"qwen3.5",
"qwen3-coder-next",
"glm-ocr",
"lfm2.5-thinking",
"glm-4.7-flash",
"translategemma",
"nemotron-3-nano",
"functiongemma",
"olmo-3.1",
"olmo-3",
"nomic-embed-text-v2-moe",
"devstral-small-2",
"rnj-1",
"devstral-2",
"qwen3-next",
"ministral-3",
"deepseek-ocr",
// "cogito-2.1", // only local tags are very large (404.0GB minimum)
"gpt-oss-safeguard",
"qwen3-vl",
"granite4",
"qwen3-embedding",
"embeddinggemma",
// "deepseek-v3.1", // only local tags are very large (404.0GB minimum)
"gpt-oss",
"qwen3-coder",
"mistral-small3.2",
"gemma3n",
"magistral",
"devstral",
"qwen2.5vl",
"phi4-reasoning",
"phi4-mini-reasoning",
"qwen3",
"granite3.3",
"deepcoder",
"mistral-small3.1",
"cogito",
"llama4",
"exaone-deep",
"command-a",
"gemma3",
"command-r7b-arabic",
"granite3.2-vision",
"phi4-mini",
"granite3.2",
"r1-1776",
"deepscaler",
"openthinker",
"deepseek-r1",
"olmo2",
"command-r7b",
// "deepseek-v3", // only local tags are very large (404.0GB minimum)
"phi4",
"dolphin3",
"smallthinker",
"granite3.1-dense",
"granite3.1-moe",
"falcon3",
"granite-embedding",
"exaone3.5",
"llama3.3",
"snowflake-arctic-embed2",
"sailor2",
"qwq",
"marco-o1",
"tulu3",
"athene-v2",
"opencoder",
"llama3.2-vision",
"smollm2",
"granite3-guardian",
"aya-expanse",
"granite3-dense",
"granite3-moe",
"nemotron",
"shieldgemma",
"llama-guard3",
"llama3.2",
"qwen2.5-coder",
"solar-pro",
"nemotron-mini",
"qwen2.5",
"bespoke-minicheck",
"mistral-small",
"reader-lm",
"minicpm-v",
// "deepseek-v2.5", // only local tags are very large (133.0GB minimum)
"reflection",
"yi-coder",
"qwen2-math",
"hermes3",
"phi3.5",
"smollm",
"bge-large",
"paraphrase-multilingual",
"bge-m3",
"mistral-large",
"llama3.1",
"nuextract",
"mistral-nemo",
"firefunction-v2",
"llama3-groq-tool-use",
"mathstral",
"codegeex4",
"glm4",
"internlm2",
"gemma2",
"deepseek-coder-v2",
"qwen2",
"deepseek-v2",
"codestral",
"granite-code",
"aya",
"falcon2",
"llama3-chatqa",
"llava-phi3",
"llava-llama3",
"llama3-gradient",
"moondream",
"phi3",
"dolphin-llama3",
"llama3",
"codeqwen",
"snowflake-arctic-embed",
"dbrx",
"command-r-plus",
"wizardlm2",
"codegemma",
"command-r",
"mxbai-embed-large",
"dolphincoder",
"starcoder2",
"all-minilm",
"nomic-embed-text",
"gemma",
"stablelm2",
"duckdb-nsql",
"qwen",
"tinydolphin",
"stable-code",
"nous-hermes2-mixtral",
"megadolphin",
"llama-pro",
"tinyllama",
"openhermes",
"notux",
"notus",
"dolphin-mistral",
"nous-hermes2",
"dolphin-phi",
"phi",
"solar",
"dolphin-mixtral",
"mixtral",
"bakllava",
"llava",
"stablelm-zephyr",
"magicoder",
"deepseek-llm",
"meditron",
"starling-lm",
"orca2",
"deepseek-coder",
"alfred",
"goliath",
"neural-chat",
"openchat",
"yi",
"yarn-mistral",
"yarn-llama2",
"xwinlm",
"mistrallite",
"codebooga",
"mistral-openorca",
"zephyr",
"nexusraven",
"samantha-mistral",
"starcoder",
"sqlcoder",
"mistral",
"falcon",
"wizardcoder",
"phind-codellama",
"codellama",
"wizardlm",
"wizardlm-uncensored",
"wizard-vicuna",
"wizard-vicuna-uncensored",
"wizard-math",
"vicuna",
"stable-beluga",
"orca-mini",
"open-orca-platypus2",
"nous-hermes",
"medllama2",
"llama2",
"llama2-uncensored",
"llama2-chinese",
"everythinglm",
"codeup",
}
func init() {
// Broad model sweeps. Each case skips models that do not expose the
// capability it is testing.
registerChatCases(testModels(libraryModels))
registerLibraryEmbeddingCases(testModels(libraryModels))
registerToolCases(testModels(libraryModels))
registerVisionTextCases(testModels(libraryModels))
}

View File

@@ -0,0 +1,132 @@
//go:build integration && release
package integration
var (
releaseUnicodeInputModel = integrationModel{Name: "deepseek-coder-v2:16b-lite-instruct-q2_K", MinVRAMGB: 12}
releaseUnicodeOutputModel = "gemma2:2b"
releaseNumPredictModel = "llama3.2:1b"
releaseParallelHistoryModel = integrationModel{Name: "gpt-oss:20b", MinVRAMGB: 16}
releaseChatModels = []integrationModel{
{Name: "gemma4", MinVRAMGB: 8},
{Name: "gemma4:12b", MinVRAMGB: 16},
{Name: "lfm2.5", MinVRAMGB: 6},
{Name: "granite4.1:8b", MinVRAMGB: 6},
{Name: "gpt-oss:20b", MinVRAMGB: 16},
{Name: "qwen3.6:27b", MinVRAMGB: 20},
{Name: "qwen3.5:2b", MinVRAMGB: 4},
{Name: "qwen3.5:2b-nvfp4", MinVRAMGB: 4},
{Name: "deepseek-r1:8b", MinVRAMGB: 6},
{Name: "mistral-small3.2:latest", MinVRAMGB: 16},
{Name: "llama3.2:latest"},
{Name: "gemma4:e2b-nvfp4", MinVRAMGB: 8},
}
releaseEmbedModels = []string{
"embeddinggemma",
"nomic-embed-text",
"all-minilm",
"bge-large",
"bge-m3",
"granite-embedding",
"mxbai-embed-large",
"paraphrase-multilingual",
"snowflake-arctic-embed",
"snowflake-arctic-embed2",
"qwen3-embedding",
}
releaseVisionModels = []string{
"nemotron3:33b",
"gemma4",
"qwen3.6:27b",
// "llama3.2-vision", // TODO: re-enable when llama.cpp supports mllama.
}
releaseVisionTextModels = []string{
"gemma4",
"qwen3.6:27b",
"qwen3.5:2b",
// "llama3.2-vision", // TODO: re-enable when llama.cpp supports mllama.
"ministral-3:3b",
}
releaseToolsModels = []string{
"lfm2.5",
"nemotron3:33b",
"gemma4",
"gpt-oss:20b",
"qwen3.6:27b",
}
releaseAudioModels = []string{
"nemotron3:33b",
"gemma4:e2b",
"gemma4:e4b",
}
)
const releaseSplitBatchVisionModel = "qwen3.5:2b"
func init() {
// Fixed release regression cases
registerIntegrationCases(
integrationTestCase("api-generate", smol, runAPIGenerate),
integrationTestCase("api-chat", smol, runAPIChat),
integrationTestCase("api-list-models", "", runAPIListModels),
integrationTestCase("api-show-model", "llama3.2", runAPIShowModel),
integrationTestCase("generate-logprobs", smol, runAPIGenerateLogprobs),
integrationTestCase("chat-logprobs", smol, runAPIChatLogprobs),
integrationTestCase("blue-sky", smol, runBlueSky),
integrationModelTestCase("unicode-input", releaseUnicodeInputModel.Name, runUnicode),
integrationModelTestCase("unicode-output", releaseUnicodeOutputModel, runExtendedUnicodeOutput),
integrationTestCase("unicode-model-dir", smol, runUnicodeModelDir),
integrationModelTestCase("num-predict", releaseNumPredictModel, runNumPredict),
integrationModelsTestCase("embed-correlation", releaseEmbedModels, runEmbedCosineDistanceCorrelation),
integrationTestCase("embedding-api", "all-minilm", runAllMiniLMEmbeddings),
integrationTestCase("embed-api", "all-minilm", runAllMiniLMEmbed),
integrationTestCase("embed-api-batch", "all-minilm", runAllMiniLMBatchEmbed),
integrationTestCase("embed-api-truncate", "all-minilm", runAllMiniLMEmbedTruncate),
integrationModelsTestCase("embed-truncation", releaseEmbedModels, runEmbedTruncation),
integrationModelsTestCase("embed-large-input", releaseEmbedModels, runEmbedLargeInput),
integrationModelsTestCase("embed-status-code", releaseEmbedModels, runEmbedStatusCode),
integrationModelsTestCase("vision-multiturn", releaseVisionModels, runVisionMultiTurn),
integrationModelsTestCase("vision-count", releaseVisionModels, runVisionObjectCounting),
integrationModelsTestCase("vision-scene", releaseVisionModels, runVisionSceneUnderstanding),
integrationModelsTestCase("vision-spatial", releaseVisionModels, runVisionSpatialReasoning),
integrationModelsTestCase("vision-detail", releaseVisionModels, runVisionDetailRecognition),
integrationModelsTestCase("vision-multi-image", releaseVisionModels, runVisionMultiImage),
integrationModelsTestCase("vision-description", releaseVisionModels, runVisionImageDescription),
integrationModelTestCase("vision-split-batch", releaseSplitBatchVisionModel, runIntegrationSplitBatch),
integrationModelsTestCase("audio-response", releaseAudioModels, runAudioResponse),
integrationModelsTestCase("openai-audio-transcription", releaseAudioModels, runOpenAIAudioTranscription),
integrationModelsTestCase("openai-chat-audio", releaseAudioModels, runOpenAIChatWithAudio),
integrationTestCase("context-long-input", smol, runLongInputContext),
integrationTestCase("context-exhaustion", smol, runContextExhaustion),
integrationModelTestCase("parallel-generate-history", releaseParallelHistoryModel.Name, runParallelGenerateWithHistory),
integrationTestCase("generate-history", smol, runGenerateWithHistory),
integrationModelTestCase("parallel-chat-history", defaultTestModel(releaseParallelHistoryModel.Name), runParallelChatWithHistory),
integrationTestCase("chat-history", smol, runChatWithHistory),
integrationTestCase("concurrent-chat", smol, runConcurrentChat),
integrationTestCase("scheduler-multimodel", "", runMultiModelStress),
integrationTestCase("scheduler-max-queue", smol, runMaxQueue),
integrationTestCase("thinking-enabled", smol, runThinkingEnabled),
integrationTestCase("thinking-suppressed", smol, runThinkingSuppressed),
integrationTestCase("create-safetensors", "", runCreateSafetensorsLLM),
integrationTestCase("create-gguf", "", runCreateGGUF),
integrationTestCase("quantization", "qwen2.5:0.5b-instruct-fp16", runQuantization),
integrationTestCase("image-generation", "", runImageGeneration),
)
// Model-parametric cases
registerModelMinVRAM([]integrationModel{releaseUnicodeInputModel, releaseParallelHistoryModel})
registerModelMinVRAM(releaseChatModels)
registerChatCases(testModels(modelNames(releaseChatModels)))
registerEmbeddingCases(testModels(releaseEmbedModels))
registerVisionTextCases(testModels(releaseVisionTextModels))
registerToolCases(testModels(releaseToolsModels))
registerToolStressCases(testModels(releaseToolsModels))
registerAudioTranscriptionCases(testModels(releaseAudioModels))
}

View File

@@ -0,0 +1,9 @@
//go:build integration && !fast && !release && !library && !imagegen
package integration
import "testing"
func TestIntegrationRequiresScope(t *testing.T) {
t.Fatal("integration tests require one of the fast, release, or library tags")
}

140
integration/reg_test.go Normal file
View File

@@ -0,0 +1,140 @@
//go:build integration
package integration
import "testing"
type integrationCase struct {
Key string
Case string
Model string
Run func(t *testing.T)
}
type integrationModel struct {
Name string
MinVRAMGB uint64
}
var (
integrationCases []integrationCase
integrationCaseKeys = map[string]struct{}{}
modelMinVRAMGB = map[string]uint64{}
)
func registerIntegrationCases(cases ...integrationCase) {
for _, c := range cases {
if _, ok := integrationCaseKeys[c.Key]; ok {
continue
}
integrationCaseKeys[c.Key] = struct{}{}
integrationCases = append(integrationCases, c)
}
}
func integrationTestCase(name, model string, run func(t *testing.T)) integrationCase {
key := name
if model != "" {
key += "/" + model
}
return integrationCase{
Key: key,
Case: name,
Model: model,
Run: run,
}
}
func integrationModelTestCase(name, model string, run func(*testing.T, string)) integrationCase {
return integrationTestCase(name, model, func(t *testing.T) {
run(t, model)
})
}
func integrationModelsTestCase(name string, models []string, run func(*testing.T, []string)) integrationCase {
return integrationTestCase(name, "", func(t *testing.T) {
run(t, models)
})
}
func registerModelIntegrationCases(name string, models []string, run func(*testing.T, string)) {
cases := make([]integrationCase, 0, len(models))
for _, model := range models {
model := model
cases = append(cases, integrationCase{
Key: name + "/" + model,
Case: name,
Model: model,
Run: func(t *testing.T) {
run(t, model)
},
})
}
registerIntegrationCases(cases...)
}
func modelNames(models []integrationModel) []string {
names := make([]string, 0, len(models))
for _, model := range models {
names = append(names, model.Name)
}
return names
}
func registerModelMinVRAM(models []integrationModel) {
for _, model := range models {
if model.MinVRAMGB > 0 {
modelMinVRAMGB[model.Name] = model.MinVRAMGB
}
}
}
func skipRegisteredMinVRAM(t *testing.T, model string) {
t.Helper()
if v, ok := modelMinVRAMGB[model]; ok {
skipUnderMinVRAM(t, v)
}
}
type knownIntegrationFlake struct {
Scenario string
Model string
Reason string
}
var knownIntegrationFlakes = []knownIntegrationFlake{
{
Scenario: "tools-stress/multi_turn",
Model: "gemma4",
Reason: "returns an empty response on the agent-style multi-turn tool prompt",
},
{
Scenario: "tools-stress/multi_turn",
Model: "qwen3.5:2b",
Reason: "returns an empty response after the tool result in the agent-style multi-turn prompt",
},
{
Scenario: "vision-text",
Model: "qwen3.5:2b",
Reason: "times out instead of returning OCR text for the Ollamas image",
},
{
Scenario: "vision-multiturn",
Model: "gemma4",
Reason: "counts five animals in the Ollamas image instead of four",
},
{
Scenario: "vision-count",
Model: "gemma4",
Reason: "counts five animals in the docs image instead of four",
},
}
func skipKnownIntegrationFlake(t *testing.T, scenario, model string) {
t.Helper()
for _, flake := range knownIntegrationFlakes {
if flake.Scenario == scenario && flake.Model == model {
t.Skipf("known model/scenario flake: %s", flake.Reason)
}
}
}

View File

@@ -11,9 +11,30 @@ import (
"github.com/ollama/ollama/api"
)
// TestThinkingEnabled verifies that when thinking is requested, the model
// produces both thinking and content output without leaking raw channel tags.
func TestThinkingEnabled(t *testing.T) {
var rawThinkingProtocolTags = []string{
"<|channel>",
"<channel|>",
"<think>",
"</think>",
"<assistant>",
"</assistant>",
"<tool_call>",
"</tool_call>",
}
func rejectRawThinkingProtocolTags(t *testing.T, field, value string) {
t.Helper()
for _, tag := range rawThinkingProtocolTags {
if strings.Contains(value, tag) {
t.Errorf("%s contains raw protocol tag %q: %s", field, tag, value)
}
}
}
// runThinkingEnabled verifies that thinking-capable models honor an explicit
// thinking request, complete a reasoning trace, and return the final answer
// without leaking raw channel tags.
func runThinkingEnabled(t *testing.T) {
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Minute)
defer cancel()
@@ -33,12 +54,15 @@ func TestThinkingEnabled(t *testing.T) {
Stream: &stream,
Think: &think,
Messages: []api.Message{
{Role: "user", Content: "What is 12 * 15? Think step by step."},
{Role: "user", Content: "What is 12 multiplied by 15? Give the final answer."},
},
Options: map[string]any{
"temperature": 0,
"seed": 42,
"num_predict": 512,
// Deep-thinking models can use several thousand tokens on
// simple problems before producing their final answer. Keep
// this high enough to verify a natural stop, not truncation.
"num_predict": 8192,
},
}
@@ -61,22 +85,17 @@ func TestThinkingEnabled(t *testing.T) {
if thinking == "" {
t.Error("expected non-empty thinking output when thinking is enabled")
}
// The answer (180) should appear in thinking, content, or both.
// Some models put everything in thinking and leave content empty
// if they hit the token limit while still thinking.
combined := thinking + " " + content
if !strings.Contains(combined, "180") {
t.Errorf("expected '180' in thinking or content, got thinking=%q content=%q", thinking, content)
if content == "" {
t.Error("expected non-empty final content after thinking")
} else if !strings.Contains(content, "180") {
t.Errorf("expected final answer 180, got content=%q", content)
}
if response.DoneReason != "stop" {
t.Errorf("expected completed response, got done reason %q", response.DoneReason)
}
// Neither thinking nor content should contain raw channel tags
if strings.Contains(content, "<|channel>") || strings.Contains(content, "<channel|>") {
t.Errorf("content contains raw channel tags: %s", content)
}
if strings.Contains(thinking, "<|channel>") || strings.Contains(thinking, "<channel|>") {
t.Errorf("thinking contains raw channel tags: %s", thinking)
}
rejectRawThinkingProtocolTags(t, "content", content)
rejectRawThinkingProtocolTags(t, "thinking", thinking)
t.Logf("thinking (%d chars): %.100s...", len(thinking), thinking)
t.Logf("content (%d chars): %s", len(content), content)
@@ -84,9 +103,9 @@ func TestThinkingEnabled(t *testing.T) {
}
}
// TestThinkingSuppressed verifies that when thinking is NOT requested,
// runThinkingSuppressed verifies that when thinking is explicitly disabled,
// the model does not leak thinking/channel content into the response.
func TestThinkingSuppressed(t *testing.T) {
func runThinkingSuppressed(t *testing.T) {
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Minute)
defer cancel()
@@ -100,10 +119,11 @@ func TestThinkingSuppressed(t *testing.T) {
pullOrSkip(ctx, t, client, modelName)
stream := false
think := api.ThinkValue{Value: false}
req := api.ChatRequest{
Model: modelName,
Stream: &stream,
// Think is nil — thinking not requested
Think: &think,
Messages: []api.Message{
{Role: "user", Content: "What is the capital of Japan? Answer in one word."},
},
@@ -129,24 +149,16 @@ func TestThinkingSuppressed(t *testing.T) {
content := response.Message.Content
thinking := response.Message.Thinking
// The answer should appear in content or thinking
combined := content + " " + thinking
if !strings.Contains(combined, "Tokyo") {
t.Errorf("expected 'Tokyo' in content or thinking, got content=%q thinking=%q", content, thinking)
// With thinking disabled, the answer must be returned as content.
if !strings.Contains(content, "Tokyo") {
t.Errorf("expected 'Tokyo' in content, got content=%q thinking=%q", content, thinking)
}
// Content must NOT contain channel/thinking tags
if strings.Contains(content, "<|channel>") || strings.Contains(content, "<channel|>") {
t.Errorf("content contains leaked channel tags when thinking not requested: %s", content)
}
if strings.Contains(content, "thought") && strings.Contains(content, "<channel|>") {
t.Errorf("content contains leaked thinking block: %s", content)
}
rejectRawThinkingProtocolTags(t, "content", content)
rejectRawThinkingProtocolTags(t, "thinking", thinking)
// Thinking field should ideally be empty when not requested.
// Some small models may still produce thinking output; log but don't fail.
if thinking != "" {
t.Logf("WARNING: model produced thinking output when not requested (%d chars): %.100s...", len(thinking), thinking)
t.Errorf("expected empty thinking when thinking is disabled, got %q", thinking)
}
t.Logf("content: %s", content)

View File

@@ -15,117 +15,90 @@ import (
"github.com/ollama/ollama/api"
)
// TestAPIToolCallingStress tests tool calling with complex, agent-style prompts
// that include large system messages, multiple tools, and multi-turn conversations.
// This catches cache corruption and parser bugs that simple tool tests miss.
func TestAPIToolCallingStress(t *testing.T) {
func registerToolStressCases(models []string) {
registerModelIntegrationCases("tools-stress", models, runAPIToolCallingStressModel)
}
var toolStressSkipModels = map[string]string{
"lfm2.5-thinking": "returns text instead of tool calls with complex system prompts",
"qwen3.5:2b": "2B model too small for reliable multi-tool agent prompts",
"qwen3-vl": "vision model, extremely slow with complex tool prompts",
"llama3.2": "3B model too small for reliable multi-tool agent prompts",
"mistral": "7B v0.3 returns text instead of tool calls with complex prompts",
"mixtral:8x22b": "returns text instead of tool calls with complex prompts",
"qwen2": "returns text instead of tool calls with complex prompts",
"granite3.3": "returns text instead of tool calls with complex prompts",
}
func runAPIToolCallingStressModel(t *testing.T, model string) {
initialTimeout := 120 * time.Second
streamTimeout := 120 * time.Second
softTimeout, _ := getTimeouts(t)
if time.Since(started) > softTimeout {
t.Skip("skipping remaining tests to avoid excessive runtime")
}
ctx, cancel := context.WithTimeout(context.Background(), 15*time.Minute)
defer cancel()
client, _, cleanup := InitServerConnection(ctx, t)
defer cleanup()
minVRAM := map[string]uint64{
"qwen3-vl": 16,
"gpt-oss:20b": 16,
"gpt-oss:120b": 70,
"qwen3": 6,
"llama3.1": 8,
"llama3.2": 4,
"mistral": 6,
"qwen2.5": 6,
"qwen2": 6,
"ministral-3": 20,
"mistral-nemo": 9,
"mistral-small": 16,
"mixtral:8x22b": 80,
"qwq": 20,
"granite3.3": 7,
runAPIToolCallingStressModelWithClient(t, ctx, client, model, initialTimeout, streamTimeout, toolsMinVRAM, toolStressSkipModels)
}
func runAPIToolCallingStressModelWithClient(t *testing.T, ctx context.Context, client *api.Client, model string, initialTimeout, streamTimeout time.Duration, minVRAM map[string]uint64, skipModels map[string]string) {
t.Helper()
// Skip known-bad models unless explicitly requested via env var
if reason, ok := skipModels[model]; ok && testModel == "" {
t.Skipf("skipping: %s", reason)
}
// Models that don't reliably produce tool calls with complex/multi-tool prompts.
// The stress test uses a large system prompt with many tools, simulating coding agents.
// Some models are too small, too slow, or not designed for this use case.
skipModels := map[string]string{
"lfm2.5-thinking": "returns text instead of tool calls with complex system prompts",
"qwen3-vl": "vision model, extremely slow with complex tool prompts",
"llama3.2": "3B model too small for reliable multi-tool agent prompts",
"mistral": "7B v0.3 returns text instead of tool calls with complex prompts",
"mixtral:8x22b": "returns text instead of tool calls with complex prompts",
"qwen2": "returns text instead of tool calls with complex prompts",
"granite3.3": "returns text instead of tool calls with complex prompts",
if v, ok := minVRAM[model]; ok {
skipUnderMinVRAM(t, v)
}
requireCapability(ctx, t, client, model, "tools")
models := testModels(libraryToolsModels)
// Preload and skip if not sufficiently GPU-loaded to avoid timeouts
preloadGenerateModel(ctx, t, client, api.GenerateRequest{Model: model})
skipIfNotGPULoaded(ctx, t, client, model, 80)
softTimeout, _ := getTimeouts(t)
tools := stressTestTools()
for _, model := range models {
t.Run(model, func(t *testing.T) {
if time.Since(started) > softTimeout {
t.Skip("skipping remaining tests to avoid excessive runtime")
return
}
// Skip known-bad models unless explicitly requested via env var
if reason, ok := skipModels[model]; ok && testModel == "" {
t.Skipf("skipping: %s", reason)
}
if testModel != "" {
requireCapability(ctx, t, client, model, "tools")
}
if v, ok := minVRAM[model]; ok {
skipUnderMinVRAM(t, v)
}
// Large system prompt that mimics real coding agents (opencode, Claude Code, etc.)
// This is intentionally very long (~5000+ tokens) to match the prompt sizes that
// real coding agents send. The combination of a large system prompt, many tools,
// and thinking mode is what triggers failures in some models.
systemPrompt := stressTestSystemPrompt()
pullOrSkip(ctx, t, client, model)
// Test 1: First request (fresh prompt processing)
// Use a direct prompt that tells the model exactly what tool to use,
// reducing the chance it asks for clarification instead.
t.Run("first_request", func(t *testing.T) {
testToolCall(t, ctx, client, model, systemPrompt, tools,
"Run git diff main to review the code changes on the current branch.",
initialTimeout, streamTimeout)
})
// Preload and skip if not sufficiently GPU-loaded to avoid timeouts
err := client.Generate(ctx, &api.GenerateRequest{Model: model}, func(response api.GenerateResponse) error { return nil })
if err != nil {
t.Fatalf("failed to load model %s: %s", model, err)
}
skipIfNotGPULoaded(ctx, t, client, model, 80)
// Test 2: Repeat with same prompt (tests cache reuse)
t.Run("cached_request", func(t *testing.T) {
testToolCall(t, ctx, client, model, systemPrompt, tools,
"Run git diff main to review the code changes on the current branch.",
initialTimeout, streamTimeout)
})
tools := stressTestTools()
// Test 3: Different user message (partial cache hit)
t.Run("different_user_message", func(t *testing.T) {
testToolCall(t, ctx, client, model, systemPrompt, tools,
"Read the file at ./go.mod and tell me what dependencies we have.",
initialTimeout, streamTimeout)
})
// Large system prompt that mimics real coding agents (opencode, Claude Code, etc.)
// This is intentionally very long (~5000+ tokens) to match the prompt sizes that
// real coding agents send. The combination of a large system prompt, many tools,
// and thinking mode is what triggers failures in some models.
systemPrompt := stressTestSystemPrompt()
// Test 1: First request (fresh prompt processing)
// Use a direct prompt that tells the model exactly what tool to use,
// reducing the chance it asks for clarification instead.
t.Run("first_request", func(t *testing.T) {
testToolCall(t, ctx, client, model, systemPrompt, tools,
"Run git diff main to review the code changes on the current branch.",
initialTimeout, streamTimeout)
})
// Test 2: Repeat with same prompt (tests cache reuse)
t.Run("cached_request", func(t *testing.T) {
testToolCall(t, ctx, client, model, systemPrompt, tools,
"Run git diff main to review the code changes on the current branch.",
initialTimeout, streamTimeout)
})
// Test 3: Different user message (partial cache hit)
t.Run("different_user_message", func(t *testing.T) {
testToolCall(t, ctx, client, model, systemPrompt, tools,
"Read the file at ./go.mod and tell me what dependencies we have.",
initialTimeout, streamTimeout)
})
// Test 4: Multi-turn with tool response
t.Run("multi_turn", func(t *testing.T) {
testToolCallMultiTurn(t, ctx, client, model, systemPrompt, tools,
initialTimeout, streamTimeout)
})
})
}
// Test 4: Multi-turn with tool response
t.Run("multi_turn", func(t *testing.T) {
skipKnownIntegrationFlake(t, "tools-stress/multi_turn", model)
testToolCallMultiTurn(t, ctx, client, model, systemPrompt, tools,
initialTimeout, streamTimeout)
})
}
func newTool(name, description string, required []string, props map[string]api.ToolProperty) api.Tool {

View File

@@ -20,134 +20,138 @@ func testPropsMap(m map[string]api.ToolProperty) *api.ToolPropertiesMap {
return props
}
func TestAPIToolCalling(t *testing.T) {
func registerToolCases(models []string) {
registerModelIntegrationCases("tools", models, runAPIToolCallingModel)
}
var toolsMinVRAM = map[string]uint64{
"gemma4": 8,
"lfm2.5": 6,
"granite4.1:3b": 4,
"granite4.1:8b": 6,
"nemotron3:33b": 32,
"qwen3.5:2b": 4,
"qwen3.6:27b": 20,
"qwen3-vl": 16,
"gpt-oss:20b": 16,
"gpt-oss:120b": 70,
"qwen3": 6,
"llama3.1": 8,
"llama3.2": 4,
"mistral": 6,
"qwen2.5": 6,
"qwen2": 6,
"ministral-3": 20,
"mistral-nemo": 9,
"mistral-small": 16,
"mixtral:8x22b": 80,
"qwq": 20,
"granite3.3": 7,
}
func runAPIToolCallingModel(t *testing.T, model string) {
initialTimeout := 60 * time.Second
streamTimeout := 60 * time.Second
softTimeout, hardTimeout := getTimeouts(t)
if time.Since(started) > softTimeout {
t.Skip("skipping remaining tests to avoid excessive runtime")
}
ctx, cancel := context.WithTimeout(context.Background(), hardTimeout)
defer cancel()
client, _, cleanup := InitServerConnection(ctx, t)
defer cleanup()
minVRAM := map[string]uint64{
"gemma4": 8,
"qwen3-vl": 16,
"gpt-oss:20b": 16,
"gpt-oss:120b": 70,
"qwen3": 6,
"llama3.1": 8,
"llama3.2": 4,
"mistral": 6,
"qwen2.5": 6,
"qwen2": 6,
"ministral-3": 20,
"mistral-nemo": 9,
"mistral-small": 16,
"mixtral:8x22b": 80,
"qwq": 20,
"granite3.3": 7,
runAPIToolCallingModelWithClient(t, ctx, client, model, initialTimeout, streamTimeout, toolsMinVRAM)
}
func runAPIToolCallingModelWithClient(t *testing.T, ctx context.Context, client *api.Client, model string, initialTimeout, streamTimeout time.Duration, minVRAM map[string]uint64) {
t.Helper()
if v, ok := minVRAM[model]; ok {
skipUnderMinVRAM(t, v)
}
requireCapability(ctx, t, client, model, "tools")
tools := []api.Tool{
{
Type: "function",
Function: api.ToolFunction{
Name: "get_weather",
Description: "Get the current weather for a location",
Parameters: api.ToolFunctionParameters{
Type: "object",
Required: []string{"location"},
Properties: testPropsMap(map[string]api.ToolProperty{
"location": {
Type: api.PropertyType{"string"},
Description: "The city and state, e.g. San Francisco, CA",
},
}),
},
},
},
}
models := testModels(libraryToolsModels)
req := api.ChatRequest{
Model: model,
Messages: []api.Message{
{
Role: "user",
Content: "Call get_weather with location set to San Francisco.",
},
},
Tools: tools,
Options: map[string]any{
"temperature": 0,
},
KeepAlive: &api.Duration{Duration: 10 * time.Second},
}
for _, model := range models {
t.Run(model, func(t *testing.T) {
if time.Now().Sub(started) > softTimeout {
t.Skip("skipping remaining tests to avoid excessive runtime")
return
}
stallTimer := time.NewTimer(initialTimeout)
var gotToolCall bool
var lastToolCall api.ToolCall
if testModel != "" {
requireCapability(ctx, t, client, model, "tools")
}
if v, ok := minVRAM[model]; ok {
skipUnderMinVRAM(t, v)
}
fn := func(response api.ChatResponse) error {
if len(response.Message.ToolCalls) > 0 {
gotToolCall = true
lastToolCall = response.Message.ToolCalls[len(response.Message.ToolCalls)-1]
}
if !stallTimer.Reset(streamTimeout) {
return fmt.Errorf("stall was detected while streaming response, aborting")
}
return nil
}
pullOrSkip(ctx, t, client, model)
stream := true
req.Stream = &stream
done := make(chan int)
var genErr error
go func() {
genErr = client.Chat(ctx, &req, fn)
done <- 0
}()
tools := []api.Tool{
{
Type: "function",
Function: api.ToolFunction{
Name: "get_weather",
Description: "Get the current weather in a given location",
Parameters: api.ToolFunctionParameters{
Type: "object",
Required: []string{"location"},
Properties: testPropsMap(map[string]api.ToolProperty{
"location": {
Type: api.PropertyType{"string"},
Description: "The city and state, e.g. San Francisco, CA",
},
}),
},
},
},
}
select {
case <-stallTimer.C:
t.Errorf("tool-calling chat never started. Timed out after: %s", initialTimeout.String())
case <-done:
if genErr != nil {
t.Fatalf("chat failed: %v", genErr)
}
req := api.ChatRequest{
Model: model,
Messages: []api.Message{
{
Role: "user",
Content: "Call get_weather with location set to San Francisco.",
},
},
Tools: tools,
Options: map[string]any{
"temperature": 0,
},
KeepAlive: &api.Duration{Duration: 10 * time.Second},
}
if !gotToolCall {
t.Fatalf("expected at least one tool call, got none")
}
stallTimer := time.NewTimer(initialTimeout)
var gotToolCall bool
var lastToolCall api.ToolCall
if lastToolCall.Function.Name != "get_weather" {
t.Errorf("unexpected tool called: got %q want %q", lastToolCall.Function.Name, "get_weather")
}
fn := func(response api.ChatResponse) error {
if len(response.Message.ToolCalls) > 0 {
gotToolCall = true
lastToolCall = response.Message.ToolCalls[len(response.Message.ToolCalls)-1]
}
if !stallTimer.Reset(streamTimeout) {
return fmt.Errorf("stall was detected while streaming response, aborting")
}
return nil
}
stream := true
req.Stream = &stream
done := make(chan int)
var genErr error
go func() {
genErr = client.Chat(ctx, &req, fn)
done <- 0
}()
select {
case <-stallTimer.C:
t.Errorf("tool-calling chat never started. Timed out after: %s", initialTimeout.String())
case <-done:
if genErr != nil {
t.Fatalf("chat failed: %v", genErr)
}
if !gotToolCall {
t.Fatalf("expected at least one tool call, got none")
}
if lastToolCall.Function.Name != "get_weather" {
t.Errorf("unexpected tool called: got %q want %q", lastToolCall.Function.Name, "get_weather")
}
if _, ok := lastToolCall.Function.Arguments.Get("location"); !ok {
t.Errorf("expected tool arguments to include 'location', got: %s", lastToolCall.Function.Arguments.String())
}
case <-ctx.Done():
t.Error("outer test context done while waiting for tool-calling chat")
}
})
if _, ok := lastToolCall.Function.Arguments.Get("location"); !ok {
t.Errorf("expected tool arguments to include 'location', got: %s", lastToolCall.Function.Arguments.String())
}
case <-ctx.Done():
t.Error("outer test context done while waiting for tool-calling chat")
}
}

View File

@@ -31,270 +31,18 @@ import (
)
var (
smol = "llama3.2:1b"
stream = false
// testModel is set via OLLAMA_TEST_MODEL env var. When set, all tests
// that loop over model lists will test only this model, and smol is
// also overridden to use it.
testModel string
testModel = os.Getenv("OLLAMA_TEST_MODEL")
smol = defaultTestModel("llama3.2:1b")
stream = false
)
var (
started = time.Now()
// Note: add newer models at the top of the list to test them first
ollamaEngineChatModels = []string{
"nemotron3:33b",
// "laguna-xs.2:q4_K_M", // TODO: re-enable when llama.cpp supports laguna.
"gemma4",
"lfm2.5-thinking",
"ministral-3",
"qwen3-coder:30b",
"gpt-oss:20b",
"gemma3n:e2b",
"mistral-small3.2:latest",
"deepseek-r1:1.5b",
// "llama3.2-vision:latest", // TODO: re-enable when llama.cpp supports mllama.
"qwen2.5-coder:latest",
"qwen2.5vl:3b",
"qwen3:0.6b", // dense
"qwen3:1.7b", // dense
"qwen3:30b", // MOE
"gemma3:1b",
"llama3.1:latest",
"llama3.2:latest",
"gemma2:latest",
"minicpm-v:latest", // arch=qwen2
"granite-code:latest", // arch=llama
}
// MLX-backed safetensors tags. These exercise the mlxrunner subprocess
// on platforms where MLX is available (today: macOS; Linux/Windows CUDA
// coming). On other platforms, skipIfMLXUnsupported turns the load
// failure into a test skip.
mlxEngineChatModels = []string{
"laguna-xs.2:nvfp4",
"qwen3.5:2b-nvfp4", // ~2.5GB, Qwen3_5 arch
"gemma4:e2b-nvfp4", // ~7.1GB, Gemma4 arch (skipped under low VRAM)
}
llamaRunnerChatModels = []string{
"mistral:latest",
"falcon3:latest",
"granite3-moe:latest",
"command-r:latest",
"nemotron-mini:latest",
"phi3.5:latest",
"internlm2:latest",
"codellama:latest", // arch=llama
"phi3:latest",
}
// Some library models are quite large - ensure large VRAM and sufficient disk space
// before running scenarios based on this set
libraryChatModels = []string{
"alfred",
"athene-v2",
"aya-expanse",
"aya",
"bakllava",
"bespoke-minicheck",
"codebooga",
"codegeex4",
"codegemma",
"codellama",
"codeqwen",
"codestral",
"codeup",
"cogito",
"command-a",
"command-r-plus",
"command-r",
"command-r7b-arabic",
"command-r7b",
"dbrx",
"deepcoder",
"deepscaler",
"deepseek-coder-v2",
"deepseek-coder",
"deepseek-llm",
"deepseek-r1",
// "deepseek-v2.5", // requires 155 GB VRAM
"deepseek-v2",
// "deepseek-v3", // requires 482 GB VRAM
"devstral",
"dolphin-llama3",
"dolphin-mistral",
"dolphin-mixtral",
"dolphin-phi",
"dolphin3",
"dolphincoder",
"duckdb-nsql",
"everythinglm",
"exaone-deep",
"exaone3.5",
"falcon",
"falcon2",
"falcon3",
"firefunction-v2",
"gemma",
"gemma2",
"gemma3",
"gemma3n",
"gemma4",
"glm4",
"goliath",
"gpt-oss:20b",
"granite-code",
"granite3-dense",
"granite3-guardian",
"granite3-moe",
"granite3.1-dense",
"granite3.1-moe",
"granite3.2-vision",
"granite3.2",
"granite3.3",
"hermes3",
"internlm2",
"lfm2.5-thinking",
"llama-guard3",
"llama-pro",
"llama2-chinese",
"llama2-uncensored",
"llama2",
"llama3-chatqa",
"llama3-gradient",
"llama3-groq-tool-use",
"llama3.1",
// "llama3.2-vision", // TODO: re-enable when llama.cpp supports mllama.
"llama3.2",
"llama3.3",
"llama3",
"llama4",
"llava-llama3",
"llava-phi3",
"llava",
"magicoder",
"magistral",
"marco-o1",
"mathstral",
"meditron",
"medllama2",
"megadolphin",
"minicpm-v",
"ministral-3",
"mistral-large",
"mistral-nemo",
"mistral-openorca",
"mistral-small",
"mistral-small3.1",
"mistral-small3.2",
"mistral",
"mistrallite",
"mixtral",
"moondream",
"nemotron-mini",
"nemotron",
"neural-chat",
"nexusraven",
"notus",
"nous-hermes",
"nous-hermes2-mixtral",
"nous-hermes2",
"nuextract",
"olmo2",
"open-orca-platypus2",
"openchat",
"opencoder",
"openhermes",
"openthinker",
"orca-mini",
"orca2",
// "phi", // unreliable
"phi3.5",
"phi3",
"phi4-mini-reasoning",
"phi4-mini",
"phi4-reasoning",
"phi4",
"phind-codellama",
"qwen",
"qwen2-math",
"qwen2.5-coder",
"qwen2.5",
"qwen2.5vl",
"qwen2",
"qwen3:0.6b", // dense
"qwen3:30b", // MOE
"qwq",
"r1-1776",
"reader-lm",
"reflection",
"sailor2",
"samantha-mistral",
"shieldgemma",
"smallthinker",
"smollm",
"smollm2",
"solar",
"sqlcoder",
"stable-beluga",
"stable-code",
"stablelm-zephyr",
"stablelm2",
"starcoder",
"starcoder2",
"starling-lm",
"tinydolphin",
"tinyllama",
"tulu3",
"vicuna",
"wizard-math",
"wizard-vicuna-uncensored",
"wizard-vicuna",
"wizardcoder",
"wizardlm-uncensored",
"wizardlm2",
"xwinlm",
"yarn-llama2",
"yarn-mistral",
"yi-coder",
"yi",
"zephyr",
}
libraryEmbedModels = []string{
"embeddinggemma",
"nomic-embed-text",
"all-minilm",
"bge-large",
"bge-m3",
"granite-embedding",
"mxbai-embed-large",
"paraphrase-multilingual",
"snowflake-arctic-embed",
"snowflake-arctic-embed2",
"qwen3-embedding",
}
libraryToolsModels = []string{
"nemotron3:33b",
// "laguna-xs.2", // TODO: re-enable when llama.cpp supports laguna.
"gemma4",
"lfm2.5-thinking",
"qwen3-vl",
"gpt-oss:20b",
"gpt-oss:120b",
"qwen3",
"llama3.1",
"llama3.2",
"mistral",
"qwen2.5",
"ministral-3",
"mistral-nemo",
"mistral-small",
"mixtral:8x22b",
"qwq",
"granite3.3",
}
blueSkyPrompt = "why is the sky blue? Be brief but factual in your reply"
blueSkyExpected = []string{"rayleigh", "scatter", "atmosphere", "nitrogen", "oxygen", "wavelength", "interact"}
@@ -314,13 +62,18 @@ func init() {
logger := slog.New(slog.NewTextHandler(os.Stdout, &slog.HandlerOptions{Level: slog.LevelDebug}))
slog.SetDefault(logger)
testModel = os.Getenv("OLLAMA_TEST_MODEL")
if testModel != "" {
slog.Info("test model override", "model", testModel)
smol = testModel
}
}
func defaultTestModel(model string) string {
if testModel != "" {
return testModel
}
return model
}
// testModels returns the override model as a single-element slice when
// OLLAMA_TEST_MODEL is set, otherwise returns the provided default list.
func testModels(defaults []string) []string {
@@ -500,7 +253,7 @@ func PullIfMissing(ctx context.Context, client *api.Client, modelName string) er
}
slog.Info("model missing", "model", modelName)
stallDuration := 60 * time.Second // This includes checksum verification, which can take a while on larger models, and slower systems
stallDuration := 2 * time.Minute // Includes checksum verification, which can take a while on larger models and slower systems.
stallTimer := time.NewTimer(stallDuration)
fn := func(resp api.ProgressResponse) error {
// fmt.Print(".")
@@ -632,14 +385,7 @@ func DoGenerate(ctx context.Context, t *testing.T, client *api.Client, genReq ap
verify := func() {
// Verify the response contains the expected data
response = buf.String()
atLeastOne := false
for _, resp := range anyResp {
if strings.Contains(strings.ToLower(response), resp) {
atLeastOne = true
break
}
}
if !atLeastOne {
if !containsExpectedResponse(response, anyResp) {
t.Fatalf("%s: none of %v found in %s", genReq.Model, anyResp, response)
}
}
@@ -777,14 +523,7 @@ func DoChat(ctx context.Context, t *testing.T, client *api.Client, req api.ChatR
verify := func() {
// Verify the response contains the expected data
response = buf.String()
atLeastOne := false
for _, resp := range anyResp {
if strings.Contains(strings.ToLower(response), resp) {
atLeastOne = true
break
}
}
if !atLeastOne {
if !containsExpectedResponse(response, anyResp) {
t.Fatalf("%s: none of %v found in \"%s\" -- request was:%s", req.Model, anyResp, response, summarizeMessages(req.Messages))
}
}
@@ -816,6 +555,24 @@ func DoChat(ctx context.Context, t *testing.T, client *api.Client, req api.ChatR
return &api.Message{Role: role, Content: buf.String()}
}
func containsExpectedResponse(response string, anyResp []string) bool {
lowerResponse := strings.ToLower(response)
normalizedResponse := normalizeResponseText(response)
for _, resp := range anyResp {
if strings.Contains(lowerResponse, strings.ToLower(resp)) {
return true
}
if strings.Contains(normalizedResponse, normalizeResponseText(resp)) {
return true
}
}
return false
}
func normalizeResponseText(s string) string {
return strings.Join(strings.Fields(strings.ToLower(s)), " ")
}
func ChatRequests() ([]api.ChatRequest, [][]string) {
genReqs, results := GenerateRequests()
reqs := make([]api.ChatRequest, len(genReqs))
@@ -835,6 +592,16 @@ func ChatRequests() ([]api.ChatRequest, [][]string) {
return reqs, results
}
func preloadGenerateModel(ctx context.Context, t *testing.T, client *api.Client, req api.GenerateRequest) {
t.Helper()
slog.Info("loading", "model", req.Model)
err := client.Generate(ctx, &req, func(response api.GenerateResponse) error { return nil })
if err != nil {
skipIfMLXUnsupported(t, err)
t.Fatalf("failed to load model %s: %s", req.Model, err)
}
}
// skipIfMLXUnsupported converts an MLX runner startup error into a test skip
// when the fingerprint matches "the MLX stack is not wired up on this host",
// and only on platforms where MLX is not yet expected to work. On Apple
@@ -851,7 +618,8 @@ func skipIfMLXUnsupported(t *testing.T, err error) {
if err == nil {
return
}
if runtime.GOOS == "darwin" && runtime.GOARCH == "arm64" {
targetGOOS, targetGOARCH := targetPlatform()
if targetGOOS == "darwin" && targetGOARCH == "arm64" {
return
}
msg := err.Error()
@@ -863,17 +631,52 @@ func skipIfMLXUnsupported(t *testing.T, err error) {
"image generation is not supported on",
} {
if strings.Contains(msg, s) {
t.Skipf("MLX not available on %s/%s: %v", runtime.GOOS, runtime.GOARCH, err)
t.Skipf("MLX not available on target %s/%s (runner %s/%s): %v", targetGOOS, targetGOARCH, runtime.GOOS, runtime.GOARCH, err)
}
}
}
func targetPlatform() (goos, goarch string) {
goos = normalizeTargetGOOS(os.Getenv("OLLAMA_TEST_HOST_OS"))
goarch = normalizeTargetGOARCH(os.Getenv("OLLAMA_TEST_HOST_ARCH"))
if goos == "" {
goos = runtime.GOOS
}
if goarch == "" {
goarch = runtime.GOARCH
}
return goos, goarch
}
func normalizeTargetGOOS(goos string) string {
switch strings.ToLower(goos) {
case "darwin":
return "darwin"
case "linux":
return "linux"
case "windows", "win32nt":
return "windows"
default:
return strings.ToLower(goos)
}
}
func normalizeTargetGOARCH(goarch string) string {
switch strings.ToLower(goarch) {
case "aarch64", "arm64":
return "arm64"
case "x86_64", "amd64":
return "amd64"
default:
return strings.ToLower(goarch)
}
}
// skipIfModelTooLargeForVRAM skips the test when the model's on-disk size
// is larger than OLLAMA_MAX_VRAM by enough that even partial GPU offload
// won't help. Uses the same 0.75x gate as TestPerfModels (model_perf_test.go)
// so vision/audio tests stay runnable on systems where the model is slightly
// over VRAM and a portion legitimately spills to CPU. No-op when
// OLLAMA_MAX_VRAM is unset.
// won't help. The 0.75x gate keeps vision/audio tests runnable on systems
// where the model is slightly over VRAM and a portion legitimately spills to
// CPU. No-op when OLLAMA_MAX_VRAM is unset.
func skipIfModelTooLargeForVRAM(ctx context.Context, t *testing.T, client *api.Client, modelName string) {
t.Helper()
s := os.Getenv("OLLAMA_MAX_VRAM")

View File

@@ -13,17 +13,6 @@ import (
"github.com/ollama/ollama/types/model"
)
// Default set of vision models to test. When OLLAMA_TEST_MODEL is set,
// only that model is tested (with a capability check for vision).
var defaultVisionModels = []string{
"nemotron3:33b",
"gemma4",
"gemma3",
// "llama3.2-vision", // TODO: re-enable when llama.cpp supports mllama.
"qwen2.5vl",
"qwen3-vl:8b",
}
// decodeTestImages returns the test images.
func decodeTestImages(t *testing.T) (abbeyRoad, docs, ollamaHome api.ImageData) {
t.Helper()
@@ -68,22 +57,17 @@ func skipIfNoVisionOverride(t *testing.T) {
// setupVisionModel pulls the model, preloads it, and skips if not GPU-loaded.
func setupVisionModel(ctx context.Context, t *testing.T, client *api.Client, model string) {
t.Helper()
if testModel == "" {
pullOrSkip(ctx, t, client, model)
}
pullOrSkip(ctx, t, client, model)
skipIfModelTooLargeForVRAM(ctx, t, client, model)
requireCapability(ctx, t, client, model, "vision")
err := client.Generate(ctx, &api.GenerateRequest{Model: model}, func(response api.GenerateResponse) error { return nil })
if err != nil {
t.Fatalf("failed to load model %s: %s", model, err)
}
preloadGenerateModel(ctx, t, client, api.GenerateRequest{Model: model})
skipIfNotGPULoaded(ctx, t, client, model, 80)
}
// TestVisionMultiTurn sends an image, gets a response, then asks follow-up
// runVisionMultiTurn sends an image, gets a response, then asks follow-up
// questions about the same image. This verifies that the KV cache correctly
// handles cached image tokens across turns.
func TestVisionMultiTurn(t *testing.T) {
func runVisionMultiTurn(t *testing.T, models []string) {
skipUnderMinVRAM(t, 16)
skipIfNoVisionOverride(t)
@@ -93,8 +77,9 @@ func TestVisionMultiTurn(t *testing.T) {
"llama3.2-vision": "miscounts animals (says 3 instead of 4) on turn 2",
}
for _, model := range testModels(defaultVisionModels) {
for _, model := range testModels(models) {
t.Run(model, func(t *testing.T) {
skipKnownIntegrationFlake(t, "vision-multiturn", model)
if reason, ok := skipModels[model]; ok && testModel == "" {
t.Skipf("skipping: %s", reason)
}
@@ -151,8 +136,8 @@ func TestVisionMultiTurn(t *testing.T) {
}
}
// TestVisionObjectCounting asks the model to count objects in an image.
func TestVisionObjectCounting(t *testing.T) {
// runVisionObjectCounting asks the model to count objects in an image.
func runVisionObjectCounting(t *testing.T, models []string) {
skipUnderMinVRAM(t, 16)
skipIfNoVisionOverride(t)
@@ -160,8 +145,9 @@ func TestVisionObjectCounting(t *testing.T) {
"llama3.2-vision": "consistently miscounts (says 3 instead of 4)",
}
for _, model := range testModels(defaultVisionModels) {
for _, model := range testModels(models) {
t.Run(model, func(t *testing.T) {
skipKnownIntegrationFlake(t, "vision-count", model)
if reason, ok := skipModels[model]; ok && testModel == "" {
t.Skipf("skipping: %s", reason)
}
@@ -191,9 +177,9 @@ func TestVisionObjectCounting(t *testing.T) {
}
}
// TestVisionSceneUnderstanding tests whether the model can identify
// runVisionSceneUnderstanding tests whether the model can identify
// cultural references and scene context from an image.
func TestVisionSceneUnderstanding(t *testing.T) {
func runVisionSceneUnderstanding(t *testing.T, models []string) {
skipUnderMinVRAM(t, 16)
skipIfNoVisionOverride(t)
@@ -203,7 +189,7 @@ func TestVisionSceneUnderstanding(t *testing.T) {
"minicpm-v": "too small for cultural reference detection",
}
for _, model := range testModels(defaultVisionModels) {
for _, model := range testModels(models) {
t.Run(model, func(t *testing.T) {
if reason, ok := skipModels[model]; ok && testModel == "" {
t.Skipf("skipping: %s", reason)
@@ -236,13 +222,13 @@ func TestVisionSceneUnderstanding(t *testing.T) {
}
}
// TestVisionSpatialReasoning tests the model's ability to identify
// runVisionSpatialReasoning tests the model's ability to identify
// objects based on their spatial position in the image.
func TestVisionSpatialReasoning(t *testing.T) {
func runVisionSpatialReasoning(t *testing.T, models []string) {
skipUnderMinVRAM(t, 16)
skipIfNoVisionOverride(t)
for _, model := range testModels(defaultVisionModels) {
for _, model := range testModels(models) {
t.Run(model, func(t *testing.T) {
ctx, cancel := context.WithTimeout(context.Background(), 3*time.Minute)
defer cancel()
@@ -274,13 +260,13 @@ func TestVisionSpatialReasoning(t *testing.T) {
}
}
// TestVisionDetailRecognition tests whether the model can identify
// runVisionDetailRecognition tests whether the model can identify
// small details like accessories in an image.
func TestVisionDetailRecognition(t *testing.T) {
func runVisionDetailRecognition(t *testing.T, models []string) {
skipUnderMinVRAM(t, 16)
skipIfNoVisionOverride(t)
for _, model := range testModels(defaultVisionModels) {
for _, model := range testModels(models) {
t.Run(model, func(t *testing.T) {
ctx, cancel := context.WithTimeout(context.Background(), 3*time.Minute)
defer cancel()
@@ -310,10 +296,10 @@ func TestVisionDetailRecognition(t *testing.T) {
}
}
// TestVisionMultiImage sends two images in a single message and asks
// runVisionMultiImage sends two images in a single message and asks
// the model to compare and contrast them. This exercises multi-image
// encoding and cross-image reasoning.
func TestVisionMultiImage(t *testing.T) {
func runVisionMultiImage(t *testing.T, models []string) {
skipUnderMinVRAM(t, 16)
skipIfNoVisionOverride(t)
@@ -322,7 +308,7 @@ func TestVisionMultiImage(t *testing.T) {
"llama3.2-vision": "does not support multi-image input",
}
for _, model := range testModels(defaultVisionModels) {
for _, model := range testModels(models) {
t.Run(model, func(t *testing.T) {
if reason, ok := skipModels[model]; ok && testModel == "" {
t.Skipf("skipping: %s", reason)
@@ -357,14 +343,14 @@ func TestVisionMultiImage(t *testing.T) {
}
}
// TestVisionImageDescription verifies that the model can describe the contents
// runVisionImageDescription verifies that the model can describe the contents
// of the ollama homepage image (a cartoon llama with "Start building with
// open models" text). Basic sanity check that the vision pipeline works.
func TestVisionImageDescription(t *testing.T) {
func runVisionImageDescription(t *testing.T, models []string) {
skipUnderMinVRAM(t, 16)
skipIfNoVisionOverride(t)
for _, model := range testModels(defaultVisionModels) {
for _, model := range testModels(models) {
t.Run(model, func(t *testing.T) {
ctx, cancel := context.WithTimeout(context.Background(), 3*time.Minute)
defer cancel()

View File

@@ -88,6 +88,7 @@ This table tracks the dispatch surface. Keep it brief; the handler comments in
| `deepseekocr` | Maps to `deepseek2-ocr`, injects missing OCR/MoE metadata, and hides embedded SAM/vision/projector tensors. | DeepSeek OCR projector translation. |
| `glmocr` | Maps GLM OCR metadata/tensors to the llama.cpp-compatible view. | GLM OCR projector translation. |
| `glm4moelite` | Maps GLM-4.7 Flash MLA metadata to the `deepseek2` path and fixes special-token metadata. | n/a |
| `laguna` | Renames legacy attention-gate tensors and SWA RoPE metadata to current llama.cpp names. | n/a |
| `nemotron_h_moe` | Fixes latent-FFN variants and hides MTP tensors. | n/a |
| `nemotron_h_omni` | Selects the Nemotron text loader and hides audio/vision/projector tensors from the text loader. | Nemotron V2 VL projector translation; audio remains disabled. |
| `llama` with Llama 3 markers | Fixes Llama 3 tokenizer metadata. | n/a |

View File

@@ -118,6 +118,43 @@ void fix_glm4moelite_eog_token_ids(gguf_context * meta) {
}
}
// =========================================================================
// laguna (text side)
// =========================================================================
bool detect_ollama_laguna(const gguf_context * meta, const ggml_context * ctx) {
const int64_t arch_kid = gguf_find_key(meta, "general.architecture");
if (arch_kid < 0 || std::strcmp(gguf_get_val_str(meta, arch_kid), "laguna") != 0) return false;
return has_key(meta, "laguna.rope.swa.dimension_count")
|| has_key(meta, "laguna.rope.swa.freq_base")
|| ggml_get_tensor(const_cast<ggml_context *>(ctx), "blk.0.attn_g.weight") != nullptr;
}
void handle_laguna(gguf_context * meta, ggml_context * ctx) {
if (!detect_ollama_laguna(meta, ctx)) return;
OLLAMA_COMPAT_LOG_INFO("%s: detected Ollama-format laguna GGUF; applying compatibility fixes\n", __func__);
copy_u32_kv(meta, "laguna.rope.swa.dimension_count", "laguna.rope.dimension_count_swa");
copy_f32_kv(meta, "laguna.rope.swa.freq_base", "laguna.rope.freq_base_swa");
std::vector<std::pair<std::string, std::string>> renames;
const int64_t n = gguf_get_n_tensors(meta);
constexpr const char * old_suffix = ".attn_g.weight";
constexpr const char * new_suffix = ".attn_gate.weight";
for (int64_t i = 0; i < n; ++i) {
const std::string name(gguf_get_tensor_name(meta, i));
const size_t pos = name.rfind(old_suffix);
if (pos != std::string::npos && pos + std::strlen(old_suffix) == name.size()) {
renames.emplace_back(name, name.substr(0, pos) + new_suffix);
}
}
for (const auto & [from, to] : renames) {
rename_tensor(meta, ctx, from.c_str(), to.c_str());
}
}
bool get_u32_kv(const gguf_context * meta, const char * key, uint32_t & out) {
const int64_t kid = gguf_find_key(meta, key);
if (kid < 0) return false;
@@ -3221,6 +3258,7 @@ bool translate_metadata(const llama_model_loader * ml,
if (arch_name == "qwen35moe") handle_qwen35moe(ml, meta, ctx);
if (arch_name == "qwen35") handle_qwen35 (ml, meta, ctx);
if (arch_name == "qwen3next") handle_qwen3next(meta, ctx);
if (arch_name == "laguna") handle_laguna (meta, ctx);
if (arch_name == "gptoss") handle_gptoss (ml, meta, ctx, arch_name);
if (arch_name == "lfm2") handle_lfm2 (ml, meta, ctx);
if (arch_name == "olmo3") handle_olmo3 (meta, arch_name);

View File

@@ -0,0 +1,42 @@
diff --git a/src/models/laguna.cpp b/src/models/laguna.cpp
index 1d705440e..7e708b144 100644
--- a/src/models/laguna.cpp
+++ b/src/models/laguna.cpp
@@ -275,16 +275,35 @@ llama_model_laguna::graph::graph(const llama_model & model, const llm_graph_para
if ((uint32_t)il >= hparams.n_layer_dense_lead) {
// MoE: sigmoid routing + score-correction bias + sum-norm +
// routed_scaling_factor (all handled by build_moe_ffn).
+ ggml_tensor * up_scale = nullptr;
+ float expert_weights_scale = hparams.expert_weights_scale;
+
+#if defined(GGML_USE_METAL)
+ const ggml_type down_type = model.layers[il].ffn_down_exps->type;
+ if (n_tokens >= 32 && (ggml_is_quantized(down_type) || down_type == GGML_TYPE_F16)) {
+ // At 32 prompt tokens, Metal switches MUL_MAT_ID from its
+ // range-safe matrix-vector kernel to FP16 matrix tiles.
+ // Laguna's routed SwiGLU activations can overflow those tiles.
+ // Generation uses one token and does not need this workaround.
+ constexpr float down_input_scale = 1.0f / 256.0f;
+ up_scale = ggml_fill(ctx0,
+ model.layers[il].ffn_exp_probs_b, down_input_scale);
+ expert_weights_scale /= down_input_scale;
+ }
+#endif
+
ggml_tensor * moe_out = build_moe_ffn(cur,
model.layers[il].ffn_gate_inp,
model.layers[il].ffn_up_exps,
model.layers[il].ffn_gate_exps,
model.layers[il].ffn_down_exps,
model.layers[il].ffn_exp_probs_b,
n_expert, n_expert_used,
LLM_FFN_SILU,
hparams.expert_weights_norm,
- hparams.expert_weights_scale,
+ expert_weights_scale,
(llama_expert_gating_func_type) hparams.expert_gating_func,
- il);
+ il,
+ nullptr, nullptr,
+ up_scale);
cb(moe_out, "ffn_moe_out", il);

View File

@@ -1,93 +0,0 @@
diff --git a/src/llama-arch.cpp b/src/llama-arch.cpp
--- a/src/llama-arch.cpp
+++ b/src/llama-arch.cpp
@@ -136,2 +136,3 @@ static const std::map<llm_arch, const char *> LLM_ARCH_NAMES = {
{ LLM_ARCH_MELLUM, "mellum" },
+ { LLM_ARCH_LAGUNA, "laguna" },
{ LLM_ARCH_UNKNOWN, "(unknown)" },
@@ -398,2 +399,3 @@ static const std::map<llm_tensor, const char *> LLM_TENSOR_NAMES = {
{ LLM_TENSOR_ATTN_GATE, "blk.%d.attn_gate" },
+ { LLM_TENSOR_ATTN_GATE_LAGUNA, "blk.%d.attn_g" },
{ LLM_TENSOR_FFN_POST_NORM, "blk.%d.post_ffw_norm" },
@@ -596,2 +598,3 @@ static const std::map<llm_tensor, llm_tensor_info> LLM_TENSOR_INFOS = {
{LLM_TENSOR_ATTN_GATE, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT}},
+ {LLM_TENSOR_ATTN_GATE_LAGUNA, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT}},
{LLM_TENSOR_FFN_GATE, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT}},
diff --git a/src/llama-arch.h b/src/llama-arch.h
--- a/src/llama-arch.h
+++ b/src/llama-arch.h
@@ -145,2 +145,3 @@ enum llm_arch {
LLM_ARCH_DFLASH,
+ LLM_ARCH_LAGUNA,
LLM_ARCH_UNKNOWN,
@@ -382,2 +383,3 @@ enum llm_tensor {
LLM_TENSOR_ATTN_GATE,
+ LLM_TENSOR_ATTN_GATE_LAGUNA,
LLM_TENSOR_FFN_GATE_INP,
diff --git a/src/llama-model.cpp b/src/llama-model.cpp
--- a/src/llama-model.cpp
+++ b/src/llama-model.cpp
@@ -48,4 +48,6 @@ static llama_model * llama_model_mapping(llm_arch arch, const llama_model_params
case LLM_ARCH_TALKIE:
return new llama_model_talkie(params);
+ case LLM_ARCH_LAGUNA:
+ return new llama_model_laguna(params);
case LLM_ARCH_DECI:
return new llama_model_deci(params);
@@ -2525,2 +2527,3 @@ llama_rope_type llama_model_rope_type(const llama_model * model) {
case LLM_ARCH_MELLUM:
+ case LLM_ARCH_LAGUNA:
case LLM_ARCH_DFLASH:
diff --git a/src/llama-model-loader.cpp b/src/llama-model-loader.cpp
--- a/src/llama-model-loader.cpp
+++ b/src/llama-model-loader.cpp
@@ -505,2 +505,3 @@ namespace GGUFMeta {
template bool llama_model_loader::get_key_or_arr<std::array<uint32_t, 512>>(enum llm_kv kid, std::array<uint32_t, 512> & result, uint32_t n, bool required);
+ template bool llama_model_loader::get_key_or_arr<uint32_t, 512>(const std::string & key, std::array<uint32_t, 512> & result, uint32_t n, bool required);
template bool llama_model_loader::get_key_or_arr<std::array<float, 512>>(enum llm_kv kid, std::array<float, 512> & result, uint32_t n, bool required);
diff --git a/src/llama-vocab.cpp b/src/llama-vocab.cpp
--- a/src/llama-vocab.cpp
+++ b/src/llama-vocab.cpp
@@ -359,2 +359,8 @@ struct llm_tokenizer_bpe : llm_tokenizer {
break;
+ case LLAMA_VOCAB_PRE_TYPE_LAGUNA:
+ regex_exprs = {
+ "(?:\\r?\\n)+(?!\\r?\\n)",
+ "(?:'[sS]|'[tT]|'[rR][eE]|'[vV][eE]|'[mM]|'[lL][lL]|'[dD])|[^\\r\\n\\p{L}\\p{N}]?\\p{L}+|\\p{N}| ?[^\\s\\p{L}\\p{N}]+[\\r\\n]*|\\s*[\\r\\n]+|\\s+(?!\\S)|\\s+",
+ };
+ break;
case LLAMA_VOCAB_PRE_TYPE_GPT2:
@@ -2100,2 +2106,4 @@ void llama_vocab::impl::load(llama_model_loader & ml, const LLM_KV & kv) {
ignore_merges = true;
+ } else if (tokenizer_pre == "laguna") {
+ pre_type = LLAMA_VOCAB_PRE_TYPE_LAGUNA;
} else if (
@@ -2773,2 +2781,3 @@ void llama_vocab::impl::load(llama_model_loader & ml, const LLM_KV & kv) {
|| t.first == "<end▁of▁sentence>" // deepseek-ocr
+ || t.first == "</assistant>" // poolside Laguna (eos_token_ids)
) {
diff --git a/src/llama-vocab.h b/src/llama-vocab.h
--- a/src/llama-vocab.h
+++ b/src/llama-vocab.h
@@ -65,2 +65,3 @@ enum llama_vocab_pre_type {
LLAMA_VOCAB_PRE_TYPE_MELLUM2 = 55,
+ LLAMA_VOCAB_PRE_TYPE_LAGUNA = 56,
};
diff --git a/src/models/models.h b/src/models/models.h
--- a/src/models/models.h
+++ b/src/models/models.h
@@ -1657,1 +1657,14 @@ struct llama_model_arcee : public llama_model_base {
+struct llama_model_laguna : public llama_model_base {
+ llama_model_laguna(const struct llama_model_params & params) : llama_model_base(params) {}
+ void load_arch_hparams(llama_model_loader & ml) override;
+ void load_arch_tensors(llama_model_loader & ml) override;
+
+ struct graph : public llm_graph_context {
+ graph(const llama_model & model, const llm_graph_params & params);
+ };
+
+ std::unique_ptr<llm_graph_context> build_arch_graph(const llm_graph_params & params) const override;
+};
+
+
struct llama_model_arcee : public llama_model_base {

View File

@@ -1,232 +0,0 @@
#include "models/models.h"
void llama_model_laguna::load_arch_hparams(llama_model_loader & ml) {
ml.get_key(LLM_KV_ATTENTION_LAYERNORM_RMS_EPS, hparams.f_norm_rms_eps);
// MoE
ml.get_key(LLM_KV_LEADING_DENSE_BLOCK_COUNT, hparams.n_layer_dense_lead, false);
ml.get_key(LLM_KV_EXPERT_FEED_FORWARD_LENGTH, hparams.n_ff_exp);
ml.get_key(LLM_KV_EXPERT_SHARED_FEED_FORWARD_LENGTH, hparams.n_ff_shexp, false);
ml.get_key(LLM_KV_EXPERT_SHARED_COUNT, hparams.n_expert_shared, false);
ml.get_key(LLM_KV_EXPERT_WEIGHTS_SCALE, hparams.expert_weights_scale, false);
ml.get_key(LLM_KV_EXPERT_WEIGHTS_NORM, hparams.expert_weights_norm, false);
ml.get_key(LLM_KV_EXPERT_GATING_FUNC, hparams.expert_gating_func, false);
ml.get_key(LLM_KV_ATTENTION_SLIDING_WINDOW, hparams.n_swa, false);
hparams.swa_type = LLAMA_SWA_TYPE_STANDARD;
ml.get_key_or_arr("laguna.attention.layer_types", hparams.is_swa_impl, hparams.n_layer(), false);
ml.get_key("laguna.rope.swa.dimension_count", hparams.n_rot_swa, false);
ml.get_key("laguna.rope.swa.freq_base", hparams.rope_freq_base_train_swa, false);
ml.get_key("laguna.rope.scaling.beta_fast", hparams.yarn_beta_fast, false);
ml.get_key("laguna.rope.scaling.beta_slow", hparams.yarn_beta_slow, false);
type = LLM_TYPE_UNKNOWN;
}
void llama_model_laguna::load_arch_tensors(llama_model_loader &) {
LLAMA_LOAD_LOCALS;
const int64_t n_ff_exp = hparams.n_ff_exp;
const int64_t n_ff_shexp = hparams.n_ff_shexp;
tok_embd = create_tensor(tn(LLM_TENSOR_TOKEN_EMBD, "weight"), {n_embd, n_vocab}, 0);
output_norm = create_tensor(tn(LLM_TENSOR_OUTPUT_NORM, "weight"), {n_embd}, 0);
output = create_tensor(tn(LLM_TENSOR_OUTPUT, "weight"), {n_embd, n_vocab}, TENSOR_NOT_REQUIRED);
if (output == NULL) {
output = create_tensor(tn(LLM_TENSOR_TOKEN_EMBD, "weight"), {n_embd, n_vocab}, TENSOR_DUPLICATED);
}
for (int i = 0; i < n_layer; ++i) {
auto & layer = layers[i];
const int64_t n_head_i = hparams.n_head(i);
const int64_t n_head_kv_i = hparams.n_head_kv(i);
const int64_t n_embd_q = n_embd_head_k * n_head_i;
const int64_t n_embd_kv = n_embd_head_k * n_head_kv_i;
layer.attn_norm = create_tensor(tn(LLM_TENSOR_ATTN_NORM, "weight", i), {n_embd}, 0);
layer.wq = create_tensor(tn(LLM_TENSOR_ATTN_Q, "weight", i), {n_embd, n_embd_q}, 0);
layer.wk = create_tensor(tn(LLM_TENSOR_ATTN_K, "weight", i), {n_embd, n_embd_kv}, 0);
layer.wv = create_tensor(tn(LLM_TENSOR_ATTN_V, "weight", i), {n_embd, n_embd_kv}, 0);
layer.wo = create_tensor(tn(LLM_TENSOR_ATTN_OUT, "weight", i), {n_embd_q, n_embd}, 0);
layer.wqkv_gate = create_tensor(tn(LLM_TENSOR_ATTN_GATE_LAGUNA, "weight", i), {n_embd, n_head_i}, 0);
layer.attn_q_norm = create_tensor(tn(LLM_TENSOR_ATTN_Q_NORM, "weight", i), {n_embd_head_k}, 0);
layer.attn_k_norm = create_tensor(tn(LLM_TENSOR_ATTN_K_NORM, "weight", i), {n_embd_head_k}, 0);
layer.ffn_norm = create_tensor(tn(LLM_TENSOR_FFN_NORM, "weight", i), {n_embd}, 0);
if (i < (int) hparams.n_layer_dense_lead) {
layer.ffn_gate = create_tensor(tn(LLM_TENSOR_FFN_GATE, "weight", i), {n_embd, n_ff}, 0);
layer.ffn_down = create_tensor(tn(LLM_TENSOR_FFN_DOWN, "weight", i), { n_ff, n_embd}, 0);
layer.ffn_up = create_tensor(tn(LLM_TENSOR_FFN_UP, "weight", i), {n_embd, n_ff}, 0);
} else {
layer.ffn_gate_inp = create_tensor(tn(LLM_TENSOR_FFN_GATE_INP, "weight", i), {n_embd, n_expert}, 0);
layer.ffn_exp_probs_b = create_tensor(tn(LLM_TENSOR_FFN_EXP_PROBS_B, "bias", i), {n_expert}, TENSOR_NOT_REQUIRED);
layer.ffn_gate_exps = create_tensor(tn(LLM_TENSOR_FFN_GATE_EXPS, "weight", i), { n_embd, n_ff_exp, n_expert}, 0);
layer.ffn_down_exps = create_tensor(tn(LLM_TENSOR_FFN_DOWN_EXPS, "weight", i), {n_ff_exp, n_embd, n_expert}, 0);
layer.ffn_up_exps = create_tensor(tn(LLM_TENSOR_FFN_UP_EXPS, "weight", i), { n_embd, n_ff_exp, n_expert}, 0);
layer.ffn_gate_shexp = create_tensor(tn(LLM_TENSOR_FFN_GATE_SHEXP, "weight", i), {n_embd, n_ff_shexp}, 0);
layer.ffn_down_shexp = create_tensor(tn(LLM_TENSOR_FFN_DOWN_SHEXP, "weight", i), {n_ff_shexp, n_embd}, 0);
layer.ffn_up_shexp = create_tensor(tn(LLM_TENSOR_FFN_UP_SHEXP, "weight", i), {n_embd, n_ff_shexp}, 0);
}
}
}
std::unique_ptr<llm_graph_context> llama_model_laguna::build_arch_graph(const llm_graph_params & params) const {
return std::make_unique<graph>(*this, params);
}
llama_model_laguna::graph::graph(const llama_model & model, const llm_graph_params & params) :
llm_graph_context(params) {
const int64_t n_embd_head = hparams.n_embd_head_v();
GGML_ASSERT(n_embd_head == hparams.n_embd_head_k());
const float kq_scale = 1.0f / sqrtf(float(n_embd_head));
ggml_tensor * cur;
ggml_tensor * inpL;
inpL = build_inp_embd(model.tok_embd);
ggml_tensor * inp_pos = build_inp_pos();
auto * inp_attn = build_attn_inp_kv_iswa();
ggml_tensor * inp_out_ids = build_inp_out_ids();
for (int il = 0; il < n_layer; ++il) {
ggml_tensor * inpSA = inpL;
const int64_t n_head_il = hparams.n_head(il);
const int64_t n_head_kv_il = hparams.n_head_kv(il);
const bool is_swa = hparams.is_swa(il);
const int rope_n_dims = hparams.n_rot(il);
const float rope_base = is_swa ? hparams.rope_freq_base_train_swa : hparams.rope_freq_base_train;
const float rope_scale = is_swa ? hparams.rope_freq_scale_train_swa : hparams.rope_freq_scale_train;
const float rope_ext = is_swa ? 0.0f : 1.0f;
const float rope_bfast = hparams.yarn_beta_fast;
const float rope_bslow = hparams.yarn_beta_slow;
cur = build_norm(inpL, model.layers[il].attn_norm, NULL, LLM_NORM_RMS, il);
cb(cur, "attn_norm", il);
// self-attention
{
ggml_tensor * Qcur = build_lora_mm(model.layers[il].wq, cur);
cb(Qcur, "Qcur", il);
ggml_tensor * Kcur = build_lora_mm(model.layers[il].wk, cur);
cb(Kcur, "Kcur", il);
ggml_tensor * Vcur = build_lora_mm(model.layers[il].wv, cur);
cb(Vcur, "Vcur", il);
ggml_tensor * gate = build_lora_mm(model.layers[il].wqkv_gate, cur);
cb(gate, "gate", il);
Qcur = ggml_reshape_3d(ctx0, Qcur, n_embd_head, n_head_il, n_tokens);
Kcur = ggml_reshape_3d(ctx0, Kcur, n_embd_head, n_head_kv_il, n_tokens);
Vcur = ggml_reshape_3d(ctx0, Vcur, n_embd_head, n_head_kv_il, n_tokens);
Qcur = build_norm(Qcur, model.layers[il].attn_q_norm, NULL, LLM_NORM_RMS, il);
cb(Qcur, "Qcur_normed", il);
Kcur = build_norm(Kcur, model.layers[il].attn_k_norm, NULL, LLM_NORM_RMS, il);
cb(Kcur, "Kcur_normed", il);
Qcur = ggml_rope_ext(ctx0, Qcur, inp_pos, nullptr,
rope_n_dims, rope_type, hparams.n_ctx_orig_yarn, rope_base, rope_scale,
rope_ext, hparams.rope_attn_factor, rope_bfast, rope_bslow);
Kcur = ggml_rope_ext(ctx0, Kcur, inp_pos, nullptr,
rope_n_dims, rope_type, hparams.n_ctx_orig_yarn, rope_base, rope_scale,
rope_ext, hparams.rope_attn_factor, rope_bfast, rope_bslow);
cb(Qcur, "Qcur", il);
cb(Kcur, "Kcur", il);
cb(Vcur, "Vcur", il);
cur = build_attn(inp_attn,
nullptr, nullptr, nullptr,
Qcur, Kcur, Vcur, nullptr, nullptr, nullptr, kq_scale, il);
cb(cur, "attn_pregate", il);
gate = ggml_softplus(ctx0, gate);
cur = ggml_reshape_3d(ctx0, cur, n_embd_head, n_head_il, n_tokens);
gate = ggml_reshape_3d(ctx0, gate, 1, n_head_il, n_tokens);
cur = ggml_mul(ctx0, cur, gate);
cur = ggml_reshape_2d(ctx0, cur, n_embd_head * n_head_il, n_tokens);
cb(cur, "attn_gated", il);
cur = build_lora_mm(model.layers[il].wo, cur, model.layers[il].wo_s);
cb(cur, "attn_out", il);
}
if (il == n_layer - 1 && inp_out_ids) {
cur = ggml_get_rows(ctx0, cur, inp_out_ids);
inpSA = ggml_get_rows(ctx0, inpSA, inp_out_ids);
}
ggml_tensor * ffn_inp = ggml_add(ctx0, cur, inpSA);
cb(ffn_inp, "ffn_inp", il);
// feed-forward
cur = build_norm(ffn_inp, model.layers[il].ffn_norm, NULL, LLM_NORM_RMS, il);
cb(cur, "ffn_norm", il);
if ((uint32_t) il < hparams.n_layer_dense_lead) {
cur = build_ffn(cur,
model.layers[il].ffn_up, NULL, NULL,
model.layers[il].ffn_gate, NULL, NULL,
model.layers[il].ffn_down, NULL, NULL,
NULL, LLM_FFN_SILU, LLM_FFN_PAR, il);
cb(cur, "ffn_out", il);
} else {
ggml_tensor * moe_out = build_moe_ffn(cur,
model.layers[il].ffn_gate_inp,
model.layers[il].ffn_up_exps,
model.layers[il].ffn_gate_exps,
model.layers[il].ffn_down_exps,
model.layers[il].ffn_exp_probs_b,
n_expert, n_expert_used,
LLM_FFN_SILU, hparams.expert_weights_norm,
hparams.expert_weights_scale,
(llama_expert_gating_func_type) hparams.expert_gating_func,
il);
cb(moe_out, "ffn_moe_out", il);
ggml_tensor * ffn_shexp = build_ffn(cur,
model.layers[il].ffn_up_shexp, NULL, NULL,
model.layers[il].ffn_gate_shexp, NULL, NULL,
model.layers[il].ffn_down_shexp, NULL, NULL,
NULL, LLM_FFN_SILU, LLM_FFN_PAR, il);
cb(ffn_shexp, "ffn_shexp", il);
cur = ggml_add(ctx0, moe_out, ffn_shexp);
cb(cur, "ffn_out", il);
}
cur = ggml_add(ctx0, cur, ffn_inp);
cur = build_cvec(cur, il);
cb(cur, "l_out", il);
inpL = cur;
}
cur = inpL;
cur = build_norm(cur, model.output_norm, NULL, LLM_NORM_RMS, -1);
cb(cur, "result_norm", -1);
res->t_embd = cur;
cur = build_lora_mm(model.output, cur, model.output_s);
cb(cur, "result_output", -1);
res->t_logits = cur;
ggml_build_forward_expand(gf, cur);
}

View File

@@ -96,6 +96,14 @@ func (p *GLM46Parser) Add(s string, done bool) (content string, thinking string,
p.buffer.WriteString(s)
events := p.parseEvents()
if done && (p.state == glm46ParserState_ToolStartedEatingWhitespace || p.state == glm46ParserState_CollectingToolContent) {
event, err := p.finalizeToolCall()
if err != nil {
return "", "", nil, fmt.Errorf("incomplete GLM tool call: %v", err)
}
events = append(events, event)
}
var toolCalls []api.ToolCall
var contentSb strings.Builder
var thinkingSb strings.Builder
@@ -123,6 +131,66 @@ func (p *GLM46Parser) Add(s string, done bool) (content string, thinking string,
return contentSb.String(), thinkingSb.String(), toolCalls, nil
}
func (p *GLM46Parser) finalizeToolCall() (glm46EventRawToolCall, error) {
raw := p.buffer.String()
if overlapLen := overlap(raw, glm46ToolCloseTag); overlapLen > 0 {
raw = strings.TrimRightFunc(raw[:len(raw)-overlapLen], unicode.IsSpace)
}
escaped := escapeGLM46Content(raw)
var parsed GLMToolCallXML
if err := xml.Unmarshal([]byte("<tool_call>"+escaped+"</tool_call>"), &parsed); err != nil {
return glm46EventRawToolCall{}, err
}
if err := validateFinalGLM46ToolCall(parsed, p.tools); err != nil {
return glm46EventRawToolCall{}, err
}
p.buffer.Reset()
p.state = glm46ParserState_CollectingContent
return glm46EventRawToolCall{raw: raw}, nil
}
// validateFinalGLM46ToolCall is intentionally stricter than normal GLM parsing.
// At end-of-stream only the outer closing tag may be missing; repairing a
// truncated argument could turn partial model output into a mutating tool call.
func validateFinalGLM46ToolCall(parsed GLMToolCallXML, tools []api.Tool) error {
functionName := strings.TrimSpace(parsed.Content)
if functionName == "" {
return fmt.Errorf("empty function name")
}
if len(parsed.Keys) != len(parsed.Values) {
return fmt.Errorf("mismatched arg_key and arg_value counts: %d keys, %d values", len(parsed.Keys), len(parsed.Values))
}
var declaredTool *api.Tool
for i := range tools {
if tools[i].Function.Name == functionName {
declaredTool = &tools[i]
break
}
}
if declaredTool == nil {
return fmt.Errorf("tool %q is not declared", functionName)
}
seen := make(map[string]struct{}, len(parsed.Keys))
for _, rawKey := range parsed.Keys {
key := strings.TrimSpace(rawKey)
if key == "" {
return fmt.Errorf("empty argument name")
}
seen[key] = struct{}{}
}
for _, required := range declaredTool.Function.Parameters.Required {
if _, ok := seen[required]; !ok {
return fmt.Errorf("required argument %q is missing for tool %q", required, functionName)
}
}
return nil
}
func (p *GLM46Parser) parseEvents() []glm46Event {
var all []glm46Event

View File

@@ -3,6 +3,7 @@ package parsers
import (
"encoding/xml"
"reflect"
"strings"
"testing"
"github.com/ollama/ollama/api"
@@ -445,6 +446,140 @@ func TestGLM46ParserStreaming(t *testing.T) {
}
}
func TestGLM46ParserFinalizesCompleteToolCallOnDone(t *testing.T) {
type chunk struct {
content string
done bool
}
toolBody := `grep
<arg_key>pattern</arg_key>
<arg_value>needle</arg_value>
<arg_key>path</arg_key>
<arg_value>.</arg_value>`
tests := []struct {
name string
chunks []chunk
}{
{
name: "empty final chunk",
chunks: []chunk{
{content: "<tool_call>" + toolBody},
{done: true},
},
},
{
name: "tool body in final chunk",
chunks: []chunk{
{content: "<tool_call>" + toolBody, done: true},
},
},
{
name: "partial outer close in final chunk",
chunks: []chunk{
{content: "<tool_call>" + toolBody + "</tool_", done: true},
},
},
}
tools := glm46FinalizationTestTools()
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
parser := GLM46Parser{}
parser.Init(tools, nil, nil)
var calls []api.ToolCall
for _, chunk := range tt.chunks {
content, thinking, got, err := parser.Add(chunk.content, chunk.done)
if err != nil {
t.Fatal(err)
}
if content != "" || thinking != "" {
t.Fatalf("content=%q thinking=%q, want empty", content, thinking)
}
calls = append(calls, got...)
}
if len(calls) != 1 {
t.Fatalf("got %d tool calls, want 1", len(calls))
}
if calls[0].Function.Name != "grep" {
t.Fatalf("tool name=%q, want grep", calls[0].Function.Name)
}
if pattern, ok := calls[0].Function.Arguments.Get("pattern"); !ok || pattern != "needle" {
t.Fatalf("pattern=%#v, %v; want needle", pattern, ok)
}
if path, ok := calls[0].Function.Arguments.Get("path"); !ok || path != "." {
t.Fatalf("path=%#v, %v; want .", path, ok)
}
})
}
}
func TestGLM46ParserRejectsIncompleteToolCallOnDone(t *testing.T) {
tests := []struct {
name string
input string
}{
{name: "empty body", input: "<tool_call>"},
{name: "partial tool name", input: "<tool_call>gr"},
{
name: "incomplete argument value",
input: `<tool_call>write
<arg_key>path</arg_key>
<arg_value>out.txt</arg_value>
<arg_key>content</arg_key>
<arg_value>partial`,
},
{
name: "undeclared tool",
input: `<tool_call>shell
<arg_key>command</arg_key>
<arg_value>echo unsafe</arg_value>`,
},
{
name: "missing required argument",
input: `<tool_call>write
<arg_key>path</arg_key>
<arg_value>out.txt</arg_value>`,
},
}
tools := glm46FinalizationTestTools()
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
parser := GLM46Parser{}
parser.Init(tools, nil, nil)
content, thinking, calls, err := parser.Add(tt.input, true)
if err == nil || !strings.Contains(err.Error(), "incomplete GLM tool call") {
t.Fatalf("error=%v, want incomplete GLM tool call", err)
}
if content != "" || thinking != "" || len(calls) != 0 {
t.Fatalf("content=%q thinking=%q calls=%v, want no output", content, thinking, calls)
}
})
}
}
func glm46FinalizationTestTools() []api.Tool {
grep := tool("grep", map[string]api.ToolProperty{
"pattern": {Type: api.PropertyType{"string"}},
"path": {Type: api.PropertyType{"string"}},
})
grep.Function.Parameters.Required = []string{"pattern", "path"}
write := tool("write", map[string]api.ToolProperty{
"path": {Type: api.PropertyType{"string"}},
"content": {Type: api.PropertyType{"string"}},
})
write.Function.Parameters.Required = []string{"path", "content"}
return []api.Tool{grep, write}
}
// TestGLMToolCallXMLOrderPreservation verifies that xml.Unmarshal preserves
// document order when collecting multiple elements with the same tag name into slices.
// This is a critical assumption for the GLM-4.6 parser's struct-based approach.

View File

@@ -94,6 +94,17 @@ func (p *LagunaParser) Init(tools []api.Tool, lastMessage *api.Message, thinkVal
return tools
}
// LagunaV8Parser matches the v8 renderer, which closes any assistant history
// turn and emits a fresh assistant generation prompt instead of continuing the
// final assistant message in place.
type LagunaV8Parser struct {
LagunaParser
}
func (p *LagunaV8Parser) Init(tools []api.Tool, _ *api.Message, thinkValue *api.ThinkValue) []api.Tool {
return p.LagunaParser.Init(tools, nil, thinkValue)
}
func (p *LagunaParser) Add(s string, done bool) (content string, thinking string, calls []api.ToolCall, err error) {
p.buffer.WriteString(s)
var contentSB, thinkingSB strings.Builder

View File

@@ -20,6 +20,23 @@ func lagunaTestTools() []api.Tool {
}}
}
func lagunaParseChunks(t *testing.T, parser Parser, chunks ...string) (string, string, []api.ToolCall) {
t.Helper()
var content, thinking string
var calls []api.ToolCall
for i, chunk := range chunks {
chunkContent, chunkThinking, chunkCalls, err := parser.Add(chunk, i == len(chunks)-1)
if err != nil {
t.Fatalf("Add(%q, done=%t): %v", chunk, i == len(chunks)-1, err)
}
content += chunkContent
thinking += chunkThinking
calls = append(calls, chunkCalls...)
}
return content, thinking, calls
}
func TestLagunaParserToolCall(t *testing.T) {
parser := ParserForName("laguna")
if parser == nil {
@@ -515,6 +532,27 @@ func TestLagunaParserNonAssistantLastMessageStillPrimesThinking(t *testing.T) {
}
}
func TestLagunaV8ParserAssistantHistoryStillPrimesThinking(t *testing.T) {
// Laguna v8 closes assistant history and emits a fresh generation prompt,
// so an assistant tail message must not switch the parser into prefill mode.
parser := ParserForName("poolside-v1")
if parser == nil {
t.Fatal("expected poolside-v1 parser")
}
if !parser.HasToolSupport() || !parser.HasThinkingSupport() {
t.Fatal("poolside-v1 parser should advertise tools and thinking")
}
parser.Init(nil, &api.Message{Role: "assistant", Content: "Previous."}, &api.ThinkValue{Value: true})
content, thinking, calls, err := parser.Add("Reasoning.</think>Answer.", true)
if err != nil {
t.Fatal(err)
}
if content != "Answer." || thinking != "Reasoning." || len(calls) != 0 {
t.Fatalf("content=%q thinking=%q calls=%d", content, thinking, len(calls))
}
}
func TestLagunaParserStripsLeadingContentWhitespace(t *testing.T) {
// No-think prompts prime </think>, so the model emits a leading newline
// before content; the parser drops it.
@@ -575,3 +613,89 @@ func TestLagunaParserSplitToolTag(t *testing.T) {
t.Fatalf("second chunk content=%q thinking=%q calls=%d", content, thinking, len(calls))
}
}
func TestLagunaParserPartialToolCallFakeoutInContent(t *testing.T) {
parser := ParserForName("laguna")
parser.Init(lagunaTestTools(), nil, nil)
content, thinking, calls := lagunaParseChunks(t, parser, "Document literal <tool_call", " fakeout")
if content != "Document literal <tool_call fakeout" || thinking != "" || len(calls) != 0 {
t.Fatalf("content=%q thinking=%q calls=%d", content, thinking, len(calls))
}
}
func TestLagunaParserPartialToolCallFakeoutInThinking(t *testing.T) {
parser := ParserForName("laguna")
parser.Init(lagunaTestTools(), nil, &api.ThinkValue{Value: true})
content, thinking, calls := lagunaParseChunks(t, parser, "<think>Document literal <tool_c", " fakeout")
if content != "" || thinking != "Document literal <tool_c fakeout" || len(calls) != 0 {
t.Fatalf("content=%q thinking=%q calls=%d", content, thinking, len(calls))
}
}
func TestLagunaParserPartialThinkOpenFakeoutInContent(t *testing.T) {
tests := []struct {
name string
thinkValue *api.ThinkValue
last *api.Message
}{
{
name: "default off",
},
{
name: "explicit off",
thinkValue: &api.ThinkValue{Value: false},
},
{
name: "enabled assistant prefill content",
thinkValue: &api.ThinkValue{Value: true},
last: &api.Message{Role: "assistant", Content: "prefill"},
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
parser := ParserForName("laguna")
parser.Init(nil, tt.last, tt.thinkValue)
content, thinking, calls := lagunaParseChunks(t, parser, "Document literal <think", " fakeout")
if content != "Document literal <think fakeout" || thinking != "" || len(calls) != 0 {
t.Fatalf("content=%q thinking=%q calls=%d", content, thinking, len(calls))
}
})
}
}
func TestLagunaParserPartialThinkCloseFakeoutAtContentStart(t *testing.T) {
tests := []struct {
name string
thinkValue *api.ThinkValue
last *api.Message
}{
{
name: "default off",
},
{
name: "explicit off",
thinkValue: &api.ThinkValue{Value: false},
},
{
name: "enabled assistant prefill content",
thinkValue: &api.ThinkValue{Value: true},
last: &api.Message{Role: "assistant", Content: "prefill"},
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
parser := ParserForName("laguna")
parser.Init(nil, tt.last, tt.thinkValue)
content, thinking, calls := lagunaParseChunks(t, parser, "</think", " fakeout")
if content != "</think fakeout" || thinking != "" || len(calls) != 0 {
t.Fatalf("content=%q thinking=%q calls=%d", content, thinking, len(calls))
}
})
}
}

View File

@@ -94,6 +94,8 @@ func ParserForName(name string) Parser {
return &LFM2Parser{hasThinkingSupport: true}
case "laguna":
return &LagunaParser{}
case "poolside-v1":
return &LagunaV8Parser{}
case "cohere":
return &CohereParser{}
default:

View File

@@ -82,15 +82,16 @@ func (r *LagunaRenderer) Render(messages []api.Message, tools []api.Tool, think
sb.WriteString(content)
sb.WriteString("\n</user>\n")
case "assistant":
content, reasoning := lagunaV2AssistantContent(message.Content, message.Thinking)
lastMessage := i == len(messages)-1
prefill := lastMessage && (strings.TrimSpace(content) != "" || strings.TrimSpace(message.Thinking) != "" || len(message.ToolCalls) > 0)
prefill := lastMessage && (strings.TrimSpace(content) != "" || strings.TrimSpace(reasoning) != "" || len(message.ToolCalls) > 0)
sb.WriteString("<assistant>\n")
// Every assistant turn opens with the reasoning block: a full
// <think>…</think> when there is reasoning, otherwise a bare
// </think> marking the turn as direct.
if reasoning := strings.TrimSpace(message.Thinking); reasoning != "" {
if reasoning := strings.TrimSpace(reasoning); reasoning != "" {
sb.WriteString("<think>\n")
sb.WriteString(reasoning)
sb.WriteString("\n</think>\n")
@@ -112,7 +113,7 @@ func (r *LagunaRenderer) Render(messages []api.Message, tools []api.Tool, think
sb.WriteString(name)
sb.WriteString("</arg_key>\n")
sb.WriteString("<arg_value>")
sb.WriteString(formatToolCallArgument(value))
sb.WriteString(formatLagunaToolCallArgument(value))
sb.WriteString("</arg_value>\n")
}
sb.WriteString("</tool_call>\n")
@@ -146,3 +147,136 @@ func (r *LagunaRenderer) Render(messages []api.Message, tools []api.Tool, think
return sb.String(), nil
}
func lagunaV2AssistantContent(content, reasoning string) (string, string) {
parts := strings.Split(content, lagunaThoughtClose)
if len(parts) == 1 {
return content, reasoning
}
if reasoning == "" {
before := strings.TrimRight(parts[0], "\n")
if i := strings.LastIndex(before, lagunaThoughtOpen); i >= 0 {
before = before[i+len(lagunaThoughtOpen):]
}
reasoning = strings.TrimLeft(before, "\n")
}
content = strings.TrimLeft(parts[len(parts)-1], "\n")
return content, reasoning
}
type LagunaV8Renderer struct{}
func (r *LagunaV8Renderer) LeadingBOS() string {
return lagunaBOS
}
func (r *LagunaV8Renderer) Render(messages []api.Message, tools []api.Tool, think *api.ThinkValue) (string, error) {
var sb strings.Builder
sb.WriteString(lagunaBOS)
thinkingEnabled := think != nil && think.Bool()
systemMessage := lagunaDefaultSystem
firstMessageIsSystem := len(messages) > 0 && messages[0].Role == "system"
if firstMessageIsSystem {
systemMessage = messages[0].Content
}
hasSystem := strings.TrimSpace(systemMessage) != ""
if hasSystem || len(tools) > 0 || thinkingEnabled {
sb.WriteString("<system>")
if hasSystem {
sb.WriteString(strings.TrimRightFunc(systemMessage, unicode.IsSpace))
if len(tools) > 0 {
sb.WriteString("\n\n")
}
}
if len(tools) > 0 {
sb.WriteString("### Tools\n\n")
sb.WriteString("You may call functions to assist with the user query.\n")
sb.WriteString("All available function signatures are listed below:\n")
sb.WriteString("<available_tools>\n")
for _, tool := range tools {
if b, err := marshalWithSpaces(tool); err == nil {
sb.Write(b)
sb.WriteByte('\n')
}
}
sb.WriteString("</available_tools>")
}
sb.WriteString("</system>\n")
}
for i, message := range messages {
if i == 0 && firstMessageIsSystem {
continue
}
content := message.Content
switch message.Role {
case "user":
sb.WriteString("<user>")
sb.WriteString(content)
sb.WriteString("</user>\n")
case "assistant":
sb.WriteString("<assistant>")
if thinkingEnabled {
sb.WriteString(lagunaThoughtOpen)
sb.WriteString(message.Thinking)
sb.WriteString(lagunaThoughtClose)
} else {
sb.WriteString(lagunaThoughtClose)
}
if content != "" {
sb.WriteString(content)
}
for _, toolCall := range message.ToolCalls {
sb.WriteString("<tool_call>")
sb.WriteString(toolCall.Function.Name)
for name, value := range toolCall.Function.Arguments.All() {
sb.WriteString("<arg_key>")
sb.WriteString(name)
sb.WriteString("</arg_key>")
sb.WriteString("<arg_value>")
sb.WriteString(formatLagunaToolCallArgument(value))
sb.WriteString("</arg_value>")
}
sb.WriteString("</tool_call>")
}
sb.WriteString("</assistant>\n")
case "tool":
sb.WriteString("<tool_response>")
sb.WriteString(content)
sb.WriteString("</tool_response>\n")
case "system":
sb.WriteString("<system>")
sb.WriteString(content)
sb.WriteString("</system>\n")
}
}
sb.WriteString("<assistant>")
if thinkingEnabled {
sb.WriteString(lagunaThoughtOpen)
} else {
sb.WriteString(lagunaThoughtClose)
}
return sb.String(), nil
}
func formatLagunaToolCallArgument(value any) string {
switch v := value.(type) {
case string:
return v
case []byte:
return string(v)
}
if b, err := marshalWithSpaces(value); err == nil {
return string(b)
}
return formatToolCallArgument(value)
}

View File

@@ -4,6 +4,7 @@ import (
"encoding/json"
"os"
"os/exec"
"path/filepath"
"strings"
"testing"
@@ -11,17 +12,23 @@ import (
"github.com/ollama/ollama/api"
)
const (
lagunaV2Template = "testdata/laguna_v2_chat_template.jinja2"
lagunaV8Template = "testdata/laguna_v8_chat_template.jinja2"
)
// lagunaToolJSON is the get_weather tool as serialized into <available_tools>,
// matching lagunaWeatherTool().
const lagunaToolJSON = `{"type": "function", "function": {"name": "get_weather", "description": "Get weather", "parameters": {"type": "object", "required": ["location"], "properties": {"location": {"type": "string", "description": "City"}}}}}`
const lagunaMathToolJSON = `{"type": "function", "function": {"name": "add", "description": "Add numbers", "parameters": {"type": "object", "required": ["a", "b"], "properties": {"a": {"type": "number", "description": "First number"}, "b": {"type": "number", "description": "Second number"}}}}}`
// TestLagunaRendererReferenceFlowCoverage checks the renderer against the Laguna
// chat template. Each want is byte-for-byte template output (verified by
// rendering chat_template.jinja), except that history tool-calls use the clean
// form — the template leaks Jinja indentation there.
// TestLagunaRendererReferenceFlowCoverage checks the renderer against byte-for-byte
// expected output from the Laguna v2 chat template. VERIFY_JINJA2=1 also verifies
// these expected values against the checked-in Jinja fixture.
func TestLagunaRendererReferenceFlowCoverage(t *testing.T) {
weather := lagunaWeatherTool()
think := func(v bool) *api.ThinkValue { return &api.ThinkValue{Value: v} }
verifyJinja2 := lagunaVerifyJinja2(t)
// system header is always emitted; with no system message the default is used
defaultHeader := "〈|EOS|〉<system>\n\n" + lagunaDefaultSystem + "\n</system>\n"
@@ -33,6 +40,10 @@ func TestLagunaRendererReferenceFlowCoverage(t *testing.T) {
think *api.ThinkValue
want string
}{
{
name: "empty_messages",
want: defaultHeader + "<assistant>\n</think>",
},
{
name: "user_only_default",
messages: []api.Message{{Role: "user", Content: "Hello"}},
@@ -59,6 +70,14 @@ func TestLagunaRendererReferenceFlowCoverage(t *testing.T) {
want: "〈|EOS|〉<system>\n\nStay concise.\n</system>\n" +
"<user>\nHi\n</user>\n<assistant>\n</think>",
},
{
name: "empty_first_system_opts_out_of_header",
messages: []api.Message{
{Role: "system", Content: ""},
{Role: "user", Content: "Hi"},
},
want: "〈|EOS|〉<user>\nHi\n</user>\n<assistant>\n</think>",
},
{
name: "additional_system",
messages: []api.Message{
@@ -71,6 +90,22 @@ func TestLagunaRendererReferenceFlowCoverage(t *testing.T) {
"<system>\nSecondary.\n</system>\n" +
"<assistant>\n</think>",
},
{
name: "empty_first_system_with_tools",
messages: []api.Message{
{Role: "system", Content: ""},
{Role: "user", Content: "Weather?"},
},
tools: weather,
want: "〈|EOS|〉<system>\n\n\n### Tools\n\n" +
"You may call functions to assist with the user query.\n" +
"All available function signatures are listed below:\n" +
"<available_tools>\n" + lagunaToolJSON + "\n</available_tools>\n\n" +
"For each function call, return an unescaped XML-like object with function name and arguments within '<tool_call>' and '</tool_call>' tags, like here:\n" +
"<tool_call>function-name\n<arg_key>argument-key</arg_key>\n<arg_value>value-of-argument-key</arg_value>\n</tool_call>" +
"\n</system>\n" +
"<user>\nWeather?\n</user>\n<assistant>\n</think>",
},
{
name: "tools_in_header",
messages: []api.Message{
@@ -102,6 +137,19 @@ func TestLagunaRendererReferenceFlowCoverage(t *testing.T) {
"\n</system>\n" +
"<user>\nWeather?\n</user>\n<assistant>\n</think>",
},
{
name: "multiple_tools_in_header",
messages: []api.Message{{Role: "user", Content: "Add then report weather"}},
tools: append(weather, lagunaMathTool()...),
want: "〈|EOS|〉<system>\n\n" + lagunaDefaultSystem + "\n\n### Tools\n\n" +
"You may call functions to assist with the user query.\n" +
"All available function signatures are listed below:\n" +
"<available_tools>\n" + lagunaToolJSON + "\n" + lagunaMathToolJSON + "\n</available_tools>\n\n" +
"For each function call, return an unescaped XML-like object with function name and arguments within '<tool_call>' and '</tool_call>' tags, like here:\n" +
"<tool_call>function-name\n<arg_key>argument-key</arg_key>\n<arg_value>value-of-argument-key</arg_value>\n</tool_call>" +
"\n</system>\n" +
"<user>\nAdd then report weather\n</user>\n<assistant>\n</think>",
},
{
name: "assistant_history",
messages: []api.Message{
@@ -138,12 +186,80 @@ func TestLagunaRendererReferenceFlowCoverage(t *testing.T) {
"<user>\nThanks\n</user>\n<assistant>\n<think>",
},
{
name: "final_assistant_prefill",
name: "assistant_extracts_thinking_from_content",
messages: []api.Message{
{Role: "user", Content: "Complete this"},
{Role: "assistant", Content: "Partial"},
{Role: "user", Content: "Explain"},
{Role: "assistant", Content: "<think>\nPlan\n</think>\nAnswer\n\n"},
{Role: "user", Content: "Next"},
},
want: defaultHeader + "<user>\nComplete this\n</user>\n<assistant>\n</think>\nPartial\n",
think: think(true),
want: defaultHeader +
"<user>\nExplain\n</user>\n" +
"<assistant>\n<think>\nPlan\n</think>\nAnswer\n</assistant>\n" +
"<user>\nNext\n</user>\n<assistant>\n<think>",
},
{
name: "assistant_thinking_metadata_overrides_content_tags",
messages: []api.Message{
{Role: "user", Content: "Explain"},
{Role: "assistant", Thinking: "Use metadata.", Content: "<think>Ignore this</think>\nAnswer"},
{Role: "user", Content: "Next"},
},
want: defaultHeader +
"<user>\nExplain\n</user>\n" +
"<assistant>\n<think>\nUse metadata.\n</think>\nAnswer\n</assistant>\n" +
"<user>\nNext\n</user>\n<assistant>\n</think>",
},
{
name: "assistant_whitespace_content_only",
messages: []api.Message{
{Role: "user", Content: "Continue"},
{Role: "assistant", Content: " \n\t "},
{Role: "user", Content: "Next"},
},
want: defaultHeader +
"<user>\nContinue\n</user>\n" +
"<assistant>\n</think>\n</assistant>\n" +
"<user>\nNext\n</user>\n<assistant>\n</think>",
},
{
name: "assistant_multiple_tool_calls_mixed_args",
messages: []api.Message{
{Role: "user", Content: "Do calls"},
{
Role: "assistant",
ToolCalls: []api.ToolCall{
{Function: api.ToolCallFunction{
Name: "echo",
Arguments: testArgsOrdered([]orderedArg{
{Key: "text", Value: "hello"},
{Key: "count", Value: 2},
}),
}},
{Function: api.ToolCallFunction{
Name: "configure",
Arguments: testArgsOrdered([]orderedArg{
{Key: "flag", Value: true},
{Key: "options", Value: map[string]any{"mode": "fast"}},
}),
}},
},
},
{Role: "user", Content: "Done?"},
},
want: defaultHeader +
"<user>\nDo calls\n</user>\n" +
"<assistant>\n</think>\n" +
"<tool_call>echo\n" +
"<arg_key>text</arg_key>\n<arg_value>hello</arg_value>\n" +
"<arg_key>count</arg_key>\n<arg_value>2</arg_value>\n" +
"</tool_call>\n" +
"<tool_call>configure\n" +
"<arg_key>flag</arg_key>\n<arg_value>true</arg_value>\n" +
"<arg_key>options</arg_key>\n<arg_value>{\"mode\": \"fast\"}</arg_value>\n" +
"</tool_call>\n" +
"</assistant>\n" +
"<user>\nDone?\n</user>\n<assistant>\n</think>",
},
}
@@ -157,22 +273,346 @@ func TestLagunaRendererReferenceFlowCoverage(t *testing.T) {
if diff := cmp.Diff(tt.want, got); diff != "" {
t.Fatalf("renderer output mismatch vs template (-want +got):\n%s", diff)
}
if verifyJinja2 {
jinja := renderLagunaJinja2Template(t, lagunaV2Template, tt.messages, tt.tools, tt.think)
if diff := cmp.Diff(jinja, tt.want); diff != "" {
t.Fatalf("hardcoded expected mismatch vs Jinja2 template (-jinja +want):\n%s", diff)
}
if diff := cmp.Diff(jinja, got); diff != "" {
t.Fatalf("renderer output mismatch vs Jinja2 template (-jinja +got):\n%s", diff)
}
}
})
}
}
func TestLagunaRendererMatchesLocalJinjaControlFlow(t *testing.T) {
if os.Getenv("VERIFY_LAGUNA_JINJA2") == "" {
t.Skip("set VERIFY_LAGUNA_JINJA2=1 to compare against the local Laguna chat_template.jinja")
func TestLagunaRendererAssistantPrefill(t *testing.T) {
got, err := (&LagunaRenderer{}).Render([]api.Message{
{Role: "user", Content: "Complete this"},
{Role: "assistant", Content: "Partial"},
}, nil, nil)
if err != nil {
t.Fatal(err)
}
python := "/Users/daniel/.codex/worktrees/7038/ollama/.venv/bin/python3"
if _, err := os.Stat(python); err != nil {
t.Fatalf("VERIFY_LAGUNA_JINJA2 requires %s with jinja2 installed", python)
want := "〈|EOS|〉<system>\n\n" + lagunaDefaultSystem + "\n</system>\n" +
"<user>\nComplete this\n</user>\n<assistant>\n</think>\nPartial\n"
if diff := cmp.Diff(want, got); diff != "" {
t.Fatalf("renderer prefill mismatch (-want +got):\n%s", diff)
}
}
func TestLagunaRendererKnownJinja2Differences(t *testing.T) {
if !lagunaVerifyJinja2(t) {
t.Skip("set VERIFY_JINJA2=1 to run Jinja2 difference checks")
}
messages := []api.Message{
{Role: "user", Content: "Complete this"},
{Role: "assistant", Content: "Partial"},
}
got, err := (&LagunaRenderer{}).Render(messages, nil, nil)
if err != nil {
t.Fatal(err)
}
jinja := renderLagunaJinja2Template(t, lagunaV2Template, messages, nil, nil)
if got == jinja {
t.Fatal("v2 assistant prefill no longer differs from Jinja2 output")
}
wantJinja := "〈|EOS|〉<system>\n\n" + lagunaDefaultSystem + "\n</system>\n" +
"<user>\nComplete this\n</user>\n<assistant>\n</think>\nPartial\n</assistant>\n<assistant>\n</think>"
if diff := cmp.Diff(wantJinja, jinja); diff != "" {
t.Fatalf("v2 assistant prefill Jinja2 reference mismatch (-want +jinja):\n%s", diff)
}
}
func TestLagunaV8RendererReferenceFlowCoverage(t *testing.T) {
weather := lagunaWeatherTool()
think := func(v bool) *api.ThinkValue { return &api.ThinkValue{Value: v} }
verifyJinja2 := lagunaVerifyJinja2(t)
defaultHeader := "〈|EOS|〉<system>" + lagunaDefaultSystem + "</system>\n"
tests := []struct {
name string
messages []api.Message
tools []api.Tool
think *api.ThinkValue
want string
}{
{
name: "empty_messages",
want: defaultHeader + "<assistant></think>",
},
{
name: "user_only_default",
messages: []api.Message{{Role: "user", Content: "Hello"}},
want: defaultHeader + "<user>Hello</user>\n<assistant></think>",
},
{
name: "user_only_think",
messages: []api.Message{{Role: "user", Content: "Hello"}},
think: think(true),
want: defaultHeader + "<user>Hello</user>\n<assistant><think>",
},
{
name: "user_only_nothink",
messages: []api.Message{{Role: "user", Content: "Hello"}},
think: think(false),
want: defaultHeader + "<user>Hello</user>\n<assistant></think>",
},
{
name: "first_system_is_header",
messages: []api.Message{
{Role: "system", Content: "Stay concise.\n\n"},
{Role: "user", Content: "Hi"},
},
want: "〈|EOS|〉<system>Stay concise.</system>\n" +
"<user>Hi</user>\n<assistant></think>",
},
{
name: "empty_first_system_opts_out_of_header",
messages: []api.Message{
{Role: "system", Content: ""},
{Role: "user", Content: "Hi"},
},
want: "〈|EOS|〉<user>Hi</user>\n<assistant></think>",
},
{
name: "empty_first_system_with_tools",
messages: []api.Message{
{Role: "system", Content: ""},
{Role: "user", Content: "Weather?"},
},
tools: weather,
want: "〈|EOS|〉<system>" +
"### Tools\n\n" +
"You may call functions to assist with the user query.\n" +
"All available function signatures are listed below:\n" +
"<available_tools>\n" + lagunaToolJSON + "\n</available_tools>" +
"</system>\n" +
"<user>Weather?</user>\n<assistant></think>",
},
{
name: "empty_first_system_thinking_enabled",
messages: []api.Message{
{Role: "system", Content: ""},
{Role: "user", Content: "Hi"},
},
think: think(true),
want: "〈|EOS|〉<system></system>\n<user>Hi</user>\n<assistant><think>",
},
{
name: "additional_system",
messages: []api.Message{
{Role: "system", Content: "Primary."},
{Role: "user", Content: "Hi"},
{Role: "system", Content: "Secondary."},
},
want: "〈|EOS|〉<system>Primary.</system>\n" +
"<user>Hi</user>\n" +
"<system>Secondary.</system>\n" +
"<assistant></think>",
},
{
name: "tools_in_header",
messages: []api.Message{
{Role: "system", Content: "Stay concise."},
{Role: "user", Content: "Weather?"},
},
tools: weather,
think: think(true),
want: "〈|EOS|〉<system>Stay concise.\n\n" +
"### Tools\n\n" +
"You may call functions to assist with the user query.\n" +
"All available function signatures are listed below:\n" +
"<available_tools>\n" + lagunaToolJSON + "\n</available_tools>" +
"</system>\n" +
"<user>Weather?</user>\n<assistant><think>",
},
{
name: "tools_default",
messages: []api.Message{{Role: "user", Content: "Weather?"}},
tools: weather,
want: "〈|EOS|〉<system>" + lagunaDefaultSystem + "\n\n" +
"### Tools\n\n" +
"You may call functions to assist with the user query.\n" +
"All available function signatures are listed below:\n" +
"<available_tools>\n" + lagunaToolJSON + "\n</available_tools>" +
"</system>\n" +
"<user>Weather?</user>\n<assistant></think>",
},
{
name: "multiple_tools_in_header",
messages: []api.Message{{Role: "user", Content: "Add then report weather"}},
tools: append(weather, lagunaMathTool()...),
want: "〈|EOS|〉<system>" + lagunaDefaultSystem + "\n\n" +
"### Tools\n\n" +
"You may call functions to assist with the user query.\n" +
"All available function signatures are listed below:\n" +
"<available_tools>\n" + lagunaToolJSON + "\n" + lagunaMathToolJSON + "\n</available_tools>" +
"</system>\n" +
"<user>Add then report weather</user>\n<assistant></think>",
},
{
name: "assistant_history",
messages: []api.Message{
{Role: "user", Content: "Add these."},
{
Role: "assistant",
Content: "\nCalling the tool.\n",
Thinking: "Need addition.",
ToolCalls: []api.ToolCall{{
Function: api.ToolCallFunction{
Name: "add",
Arguments: testArgsOrdered([]orderedArg{
{Key: "a", Value: 2},
{Key: "b", Value: 3},
}),
},
}},
},
{Role: "tool", Content: "5"},
{Role: "user", Content: "Thanks"},
},
think: think(true),
want: defaultHeader +
"<user>Add these.</user>\n" +
"<assistant>" +
"<think>Need addition.</think>" +
"\nCalling the tool.\n" +
"<tool_call>add" +
"<arg_key>a</arg_key><arg_value>2</arg_value>" +
"<arg_key>b</arg_key><arg_value>3</arg_value>" +
"</tool_call>" +
"</assistant>\n" +
"<tool_response>5</tool_response>\n" +
"<user>Thanks</user>\n<assistant><think>",
},
{
name: "assistant_reasoning_ignored_when_thinking_disabled",
messages: []api.Message{
{Role: "user", Content: "Explain"},
{Role: "assistant", Thinking: "Hidden plan.", Content: "Answer"},
{Role: "user", Content: "Next"},
},
want: defaultHeader +
"<user>Explain</user>\n" +
"<assistant></think>Answer</assistant>\n" +
"<user>Next</user>\n<assistant></think>",
},
{
name: "assistant_empty_reasoning_when_thinking_enabled",
messages: []api.Message{
{Role: "user", Content: "Explain"},
{Role: "assistant", Content: "Answer"},
{Role: "user", Content: "Next"},
},
think: think(true),
want: defaultHeader +
"<user>Explain</user>\n" +
"<assistant><think></think>Answer</assistant>\n" +
"<user>Next</user>\n<assistant><think>",
},
{
name: "assistant_preserves_content_whitespace",
messages: []api.Message{
{Role: "user", Content: "Explain"},
{Role: "assistant", Content: "\nAnswer\n"},
{Role: "user", Content: "Next"},
},
want: defaultHeader +
"<user>Explain</user>\n" +
"<assistant></think>\nAnswer\n</assistant>\n" +
"<user>Next</user>\n<assistant></think>",
},
{
name: "assistant_multiple_tool_calls_mixed_args",
messages: []api.Message{
{Role: "user", Content: "Do calls"},
{
Role: "assistant",
ToolCalls: []api.ToolCall{
{Function: api.ToolCallFunction{
Name: "echo",
Arguments: testArgsOrdered([]orderedArg{
{Key: "text", Value: "hello"},
{Key: "count", Value: 2},
}),
}},
{Function: api.ToolCallFunction{
Name: "configure",
Arguments: testArgsOrdered([]orderedArg{
{Key: "flag", Value: true},
{Key: "options", Value: map[string]any{"mode": "fast"}},
}),
}},
},
},
{Role: "user", Content: "Done?"},
},
want: defaultHeader +
"<user>Do calls</user>\n" +
"<assistant></think>" +
"<tool_call>echo" +
"<arg_key>text</arg_key><arg_value>hello</arg_value>" +
"<arg_key>count</arg_key><arg_value>2</arg_value>" +
"</tool_call>" +
"<tool_call>configure" +
"<arg_key>flag</arg_key><arg_value>true</arg_value>" +
"<arg_key>options</arg_key><arg_value>{\"mode\": \"fast\"}</arg_value>" +
"</tool_call>" +
"</assistant>\n" +
"<user>Done?</user>\n<assistant></think>",
},
{
name: "final_assistant_closes_then_generation_prompt",
messages: []api.Message{
{Role: "user", Content: "Complete this"},
{Role: "assistant", Content: "Partial"},
},
want: defaultHeader +
"<user>Complete this</user>\n" +
"<assistant></think>Partial</assistant>\n" +
"<assistant></think>",
},
}
renderer := &LagunaV8Renderer{}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got, err := renderer.Render(tt.messages, tt.tools, tt.think)
if err != nil {
t.Fatal(err)
}
if diff := cmp.Diff(tt.want, got); diff != "" {
t.Fatalf("renderer output mismatch vs template (-want +got):\n%s", diff)
}
if verifyJinja2 {
jinja := renderLagunaJinja2Template(t, lagunaV8Template, tt.messages, tt.tools, tt.think)
if diff := cmp.Diff(jinja, tt.want); diff != "" {
t.Fatalf("hardcoded expected mismatch vs Jinja2 template (-jinja +want):\n%s", diff)
}
if diff := cmp.Diff(jinja, got); diff != "" {
t.Fatalf("renderer output mismatch vs Jinja2 template (-jinja +got):\n%s", diff)
}
}
})
}
}
func TestLagunaRendererMatchesJinja2ExpandedParity(t *testing.T) {
if os.Getenv("VERIFY_JINJA2") == "" {
t.Skip("set VERIFY_JINJA2=1 to run expanded Jinja2 parity checks")
}
lagunaVerifyJinja2(t)
tests := []struct {
name string
messages []api.Message
tools []api.Tool
think *api.ThinkValue
}{
{
@@ -206,77 +646,198 @@ func TestLagunaRendererMatchesLocalJinjaControlFlow(t *testing.T) {
messages: []api.Message{{Role: "user", Content: "Answer directly."}},
think: &api.ThinkValue{Value: false},
},
{
name: "tools_and_assistant_history",
messages: []api.Message{
{Role: "system", Content: "Stay concise."},
{Role: "user", Content: "Weather?"},
{Role: "assistant", Content: "Calling.", Thinking: "Need weather.", ToolCalls: []api.ToolCall{{
Function: api.ToolCallFunction{
Name: "get_weather",
Arguments: testArgsOrdered([]orderedArg{{Key: "location", Value: "Paris"}}),
},
}}},
{Role: "tool", Content: "Sunny"},
{Role: "user", Content: "Thanks"},
},
tools: lagunaWeatherTool(),
think: &api.ThinkValue{Value: true},
},
}
renderer := &LagunaRenderer{}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got, err := renderer.Render(tt.messages, nil, tt.think)
if err != nil {
t.Fatal(err)
}
for _, modelDir := range []string{
"/Users/daniel/Models/poolside/laguna-xs-23-04-2026",
} {
want := renderLagunaChatTemplate(t, python, modelDir, tt.messages, tt.think)
if diff := cmp.Diff(want, got); diff != "" {
t.Fatalf("%s mismatch (-chat_template +renderer):\n%s", modelDir, diff)
}
variants := []struct {
name string
renderer Renderer
template string
}{
{name: "v2", renderer: &LagunaRenderer{}, template: lagunaV2Template},
{name: "v8", renderer: &LagunaV8Renderer{}, template: lagunaV8Template},
}
for _, variant := range variants {
t.Run(variant.name, func(t *testing.T) {
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got, err := variant.renderer.Render(tt.messages, tt.tools, tt.think)
if err != nil {
t.Fatal(err)
}
want := renderLagunaJinja2Template(t, variant.template, tt.messages, tt.tools, tt.think)
if diff := cmp.Diff(want, got); diff != "" {
t.Fatalf("renderer output mismatch vs Jinja2 template (-jinja +got):\n%s", diff)
}
})
}
})
}
}
func renderLagunaChatTemplate(t *testing.T, python, modelDir string, messages []api.Message, think *api.ThinkValue) string {
func lagunaVerifyJinja2(t *testing.T) bool {
t.Helper()
if os.Getenv("VERIFY_JINJA2") == "" {
return false
}
python := lagunaJinjaPython(t)
cmd := exec.Command(python, "-c", "from transformers.utils.chat_template_utils import _compile_jinja_template")
if out, err := cmd.CombinedOutput(); err != nil {
t.Fatalf("VERIFY_JINJA2=1 requires transformers chat template support in %s: %v\n%s", python, err, out)
}
return true
}
func lagunaJinjaPython(t *testing.T) string {
t.Helper()
python, err := exec.LookPath("python3")
if err != nil {
t.Fatal("VERIFY_JINJA2=1 requires python3 on PATH")
}
return python
}
func renderLagunaJinja2Template(t *testing.T, templateRelPath string, messages []api.Message, tools []api.Tool, think *api.ThinkValue) string {
t.Helper()
type templateMessage struct {
Role string `json:"role"`
Content string `json:"content"`
templatePath, err := filepath.Abs(templateRelPath)
if err != nil {
t.Fatalf("failed to get template path: %v", err)
}
templateMessages := make([]templateMessage, 0, len(messages))
type jinjaToolCall struct {
Function struct {
Name string `json:"name"`
Arguments json.RawMessage `json:"arguments"`
} `json:"function"`
}
type jinjaMessage struct {
Role string `json:"role"`
Content string `json:"content"`
Reasoning string `json:"reasoning,omitempty"`
ReasoningContent string `json:"reasoning_content,omitempty"`
ToolCalls []jinjaToolCall `json:"tool_calls,omitempty"`
}
jinjaMessages := make([]jinjaMessage, 0, len(messages))
for _, msg := range messages {
templateMessages = append(templateMessages, templateMessage{
Role: msg.Role,
Content: msg.Content,
})
jm := jinjaMessage{
Role: msg.Role,
Content: msg.Content,
Reasoning: msg.Thinking,
ReasoningContent: msg.Thinking,
}
for _, call := range msg.ToolCalls {
jc := jinjaToolCall{}
jc.Function.Name = call.Function.Name
raw, err := call.Function.Arguments.MarshalJSON()
if err != nil {
t.Fatalf("failed to marshal tool args: %v", err)
}
jc.Function.Arguments = json.RawMessage(raw)
jm.ToolCalls = append(jm.ToolCalls, jc)
}
jinjaMessages = append(jinjaMessages, jm)
}
messagesJSON, err := json.Marshal(templateMessages)
messagesJSON, err := json.Marshal(jinjaMessages)
if err != nil {
t.Fatalf("failed to marshal messages: %v", err)
}
enableThinking := "False"
if think != nil && think.Bool() {
enableThinking = "True"
toolsJSON := "None"
if len(tools) > 0 {
b, err := json.Marshal(tools)
if err != nil {
t.Fatalf("failed to marshal tools: %v", err)
}
toolsJSON = string(b)
}
enableThinking := "unset"
if think != nil {
if think.Bool() {
enableThinking = "true"
} else {
enableThinking = "false"
}
}
script := `
import json
import sys
from transformers import AutoTokenizer
from pathlib import Path
from transformers.utils.chat_template_utils import _compile_jinja_template
model_dir = sys.argv[1]
messages = json.loads(sys.argv[2])
enable_thinking = sys.argv[3] == "True"
tok = AutoTokenizer.from_pretrained(model_dir, trust_remote_code=True)
print(tok.apply_chat_template(
messages,
tokenize=False,
add_generation_prompt=True,
enable_thinking=enable_thinking,
), end="")
template_path, messages_json, tools_json, enable_thinking = sys.argv[1:5]
tmpl = _compile_jinja_template(Path(template_path).read_text())
kwargs = {
"messages": json.loads(messages_json),
"add_generation_prompt": True,
}
if tools_json != "None":
kwargs["tools"] = json.loads(tools_json)
if enable_thinking == "true":
kwargs["enable_thinking"] = True
elif enable_thinking == "false":
kwargs["enable_thinking"] = False
print(tmpl.render(**kwargs), end="")
`
cmd := exec.Command(python, "-c", script, modelDir, string(messagesJSON), enableThinking)
cmd := exec.Command(lagunaJinjaPython(t), "-c", script, templatePath, string(messagesJSON), toolsJSON, enableThinking)
var stdout, stderr strings.Builder
cmd.Stdout = &stdout
cmd.Stderr = &stderr
if err := cmd.Run(); err != nil {
t.Fatalf("chat_template render failed: %v\nstderr: %s", err, stderr.String())
t.Fatalf("python render failed: %v\nstderr: %s", err, stderr.String())
}
return stdout.String()
}
func TestLagunaTemplateFixturesMatchExpectedVersions(t *testing.T) {
v2, err := os.ReadFile(lagunaV2Template)
if err != nil {
t.Fatalf("failed to read %s: %v", lagunaV2Template, err)
}
v8, err := os.ReadFile(lagunaV8Template)
if err != nil {
t.Fatalf("failed to read %s: %v", lagunaV8Template, err)
}
if !strings.Contains(string(v2), "laguna_glm_thinking_v5/chat_template.jinja") {
t.Fatalf("%s does not look like the v2 Laguna template fixture", lagunaV2Template)
}
if !strings.Contains(string(v8), "laguna_glm_thinking_v8/chat_template.jinja") {
t.Fatalf("%s does not look like the v8 Laguna template fixture", lagunaV8Template)
}
if !strings.Contains(string(v2), "render_assistant_messages_raw") {
t.Fatalf("%s should retain the v2 raw assistant branch", lagunaV2Template)
}
if strings.Contains(string(v8), "render_assistant_messages_raw") {
t.Fatalf("%s unexpectedly contains the v2 raw assistant branch", lagunaV8Template)
}
if diff := cmp.Diff(string(v2), string(v8)); diff == "" {
t.Fatal("Laguna v2 and v8 template fixtures unexpectedly match")
}
}
func lagunaWeatherTool() []api.Tool {
return []api.Tool{{
Type: "function",
@@ -297,3 +858,33 @@ func lagunaWeatherTool() []api.Tool {
},
}}
}
func lagunaMathTool() []api.Tool {
return []api.Tool{{
Type: "function",
Function: api.ToolFunction{
Name: "add",
Description: "Add numbers",
Parameters: api.ToolFunctionParameters{
Type: "object",
Required: []string{"a", "b"},
Properties: testPropsOrdered([]orderedProp{
{
Key: "a",
Value: api.ToolProperty{
Type: api.PropertyType{"number"},
Description: "First number",
},
},
{
Key: "b",
Value: api.ToolProperty{
Type: api.PropertyType{"number"},
Description: "Second number",
},
},
}),
},
},
}}
}

View File

@@ -109,6 +109,8 @@ func rendererForName(name string) Renderer {
return &LFM2Renderer{IsThinking: true, useImgTags: RenderImgTags}
case "laguna":
return &LagunaRenderer{}
case "poolside-v1":
return &LagunaV8Renderer{}
case "cohere":
return &CohereRenderer{}
default:

View File

@@ -69,6 +69,7 @@ func TestLeadingBOSForRenderer(t *testing.T) {
{name: "lfm2", want: "<|startoftext|>"},
{name: "lfm2-thinking", want: "<|startoftext|>"},
{name: "laguna", want: "〈|EOS|〉"},
{name: "poolside-v1", want: "〈|EOS|〉"},
{name: "deepseek3.1", want: "<begin▁of▁sentence>"},
{name: "cogito", want: "<begin▁of▁sentence>"},
{name: "qwen3-coder", want: ""},

View File

@@ -0,0 +1,132 @@
{#- Iteration on laguna_glm_thinking_v5/chat_template.jinja -#}
{#- Adds a default system message (used when no system message is provided in `messages`). -#}
{{- "〈|EOS|〉" -}}
{%- set enable_thinking = enable_thinking | default(false) -%}
{%- set render_assistant_messages_raw = render_assistant_messages_raw | default(false) -%}
{%- set add_generation_prompt = add_generation_prompt | default(false) -%}
{#- ───── header (system message) ───── -#}
{%- set system_message = "You are a helpful, conversationally-fluent assistant made by Poolside. You are here to be helpful to users through natural language conversations." -%}
{%- if messages and messages[0].role == "system" -%}
{%- set system_message = messages[0].content -%}
{%- endif -%}
{%- if (system_message and system_message.strip()) or tools -%}
{{- "<system>\n" -}}
{%- if system_message and system_message.strip() -%}
{{- "\n" -}}
{{- system_message.rstrip() -}}
{%- endif -%}
{%- if tools -%}
{{- "\n\n### Tools\n\n" -}}
{%- set ns = namespace(tool_string="You may call functions to assist with the user query.\n"
~ "All available function signatures are listed below:\n"
~ "<available_tools>\n") -%}
{%- for tool in tools -%}
{%- set ns.tool_string = ns.tool_string ~ (tool | tojson) ~ "\n" -%}
{%- endfor -%}
{%- if enable_thinking -%}
{%- set tool_string = ns.tool_string + "</available_tools>\n\n" ~
"Wrap your thinking in '<think>', '</think>' tags, followed by a function call. For each function call, return an unescaped XML-like object with function name and arguments within '<tool_call>' and '</tool_call>' tags, like here:\n" ~
"<think> your thoughts here </think>\n" ~
"<tool_call>function-name\n<arg_key>argument-key</arg_key>\n<arg_value>value-of-argument-key</arg_value>\n" ~
"</tool_call>" -%}
{%- else -%}
{%- set tool_string = ns.tool_string + "</available_tools>\n\n" ~
"For each function call, return an unescaped XML-like object " ~
"with function name and arguments within '<tool_call>' and '</tool_call>' tags, like here:\n" ~
"<tool_call>function-name\n<arg_key>argument-key</arg_key>\n<arg_value>value-of-argument-key</arg_value>\n" ~
"</tool_call>" -%}
{%- endif -%}
{{- tool_string -}}
{%- endif -%}
{{- "\n</system>\n" -}}
{%- endif -%}
{#- ───── main loop ───── -#}
{%- for message in messages -%}
{%- set content = message.content if message.content is string else "" -%}
{%- if message.role == "user" -%}
{{- "<user>\n" + content + "\n</user>\n" -}}
{%- elif message.role == "assistant" -%}
{%- generation -%}
{{- "<assistant>\n" -}}
{%- if render_assistant_messages_raw -%}
{#- Raw mode: prepend the generation prompt token, then dump content verbatim. -#}
{#- The generation prompt is <think> when enable_thinking, </think> otherwise. -#}
{#- Only prepend if content doesn't already start with it. -#}
{%- if enable_thinking -%}
{%- if not content.startswith('<think>') -%}
{{- '<think>' -}}
{%- endif -%}
{%- else -%}
{%- if not content.startswith('</think>') -%}
{{- '</think>' -}}
{%- endif -%}
{%- endif -%}
{{- content -}}
{#- Append closing tag if content doesn't already end with it. -#}
{%- if not content.endswith('</assistant>\n') and not content.endswith('</assistant>') -%}
{{- '\n</assistant>' -}}
{%- endif -%}
{{- "\n" -}}
{%- else -%}
{#- Extract reasoning content from message.reasoning (vLLM field name) or message.reasoning_content, or from <think> tags -#}
{%- set reasoning_content = '' %}
{%- if message.reasoning is string %}
{%- set reasoning_content = message.reasoning %}
{%- elif message.reasoning_content is string %}
{%- set reasoning_content = message.reasoning_content %}
{%- endif %}
{#- Always strip <think> tags from content if present to avoid duplication -#}
{%- if '</think>' in content %}
{%- if not reasoning_content %}
{%- set reasoning_content = content.split('</think>')[0].rstrip('\n').split('<think>')[-1].lstrip('\n') %}
{%- endif %}
{%- set content = content.split('</think>')[-1].lstrip('\n') %}
{%- endif %}
{#- Display reasoning content for all messages -#}
{%- if reasoning_content -%}
{{- '<think>\n' + reasoning_content.strip() + '\n</think>\n' -}}
{%- else -%}
{{- '</think>\n' -}}
{%- endif -%}
{#- Display main content -#}
{%- if content.strip() -%}
{{- content.strip() ~ "\n" -}}
{%- endif -%}
{%- if message.tool_calls -%}
{%- for tool_call in message.tool_calls -%}
{%- set function_data = tool_call.function -%}
{{- '<tool_call>' + function_data.name }}
{% set _args = function_data.arguments %}
{%- for k, v in _args.items() -%}
{{- "<arg_key>" ~ k ~ "</arg_key>\n" -}}
{{- "<arg_value>"}}{{ v | tojson(ensure_ascii=False) if v is not string else v }}{{ "</arg_value>\n" -}}
{%- endfor -%}
{{- "</tool_call>\n" -}}
{%- endfor -%}
{%- endif -%}
{{- "</assistant>\n" -}}
{%- endif -%}
{%- endgeneration -%}
{%- elif message.role == "tool" -%}
{{- "<tool_response>\n" + content + "\n</tool_response>\n" -}}
{%- elif message.role == "system" and loop.index0 != 0 -%}
{#- Render additional system messages (skip the first one which is handled separately in the header) -#}
{{- "<system>\n" + content + "\n</system>\n" -}}
{%- endif -%}
{%- endfor -%}
{#- ───── generation prompt ───── -#}
{%- if add_generation_prompt -%}
{{- "<assistant>\n" -}}
{#- ───── Include reasoning mode directive ───── -#}
{%- if not enable_thinking %}
{{- '</think>' -}}
{%- else %}
{{- '<think>' -}}
{%- endif %}
{%- endif -%}

View File

@@ -0,0 +1,93 @@
{#- Iteration on laguna_glm_thinking_v8/chat_template.jinja -#}
{#- No formatting instructions -#}
{{- "〈|EOS|〉" -}}
{%- set enable_thinking = enable_thinking | default(false) -%}
{%- set add_generation_prompt = add_generation_prompt | default(false) -%}
{#- ───── header (system message) ───── -#}
{#- A caller-supplied system message with empty content opts out of the default below, producing no <system> block — used to train without a system message. -#}
{%- set system_message = "You are a helpful, conversationally-fluent assistant made by Poolside. You are here to be helpful to users through natural language conversations." -%}
{%- if messages and messages[0].role == "system" -%}
{%- set system_message = messages[0].content -%}
{%- set messages = messages[1:] -%}
{%- endif -%}
{%- set has_sys = system_message and system_message.strip() -%}
{%- if has_sys or tools or enable_thinking -%}
{{- "<system>" -}}
{%- if has_sys -%}
{{- system_message.rstrip() -}}
{%- if tools -%}{{- "\n\n" -}}{%- endif -%}
{%- endif -%}
{%- if tools -%}
{{- "### Tools\n\n" -}}
{{- "You may call functions to assist with the user query.\n" -}}
{{- "All available function signatures are listed below:\n" -}}
{{- "<available_tools>\n" -}}
{%- for tool in tools -%}
{{- (tool | tojson) ~ "\n" -}}
{%- endfor -%}
{{- "</available_tools>" -}}
{%- endif -%}
{{- "</system>\n" -}}
{%- endif -%}
{#- ───── main loop ───── -#}
{%- for message in messages -%}
{%- set content = message.content if message.content is string else "" -%}
{%- if message.role == "user" -%}
{{- "<user>" + content + "</user>\n" -}}
{%- elif message.role == "assistant" -%}
{%- generation -%}
{{- "<assistant>" -}}
{#- Extract reasoning content from message.reasoning (vLLM field name) or message.reasoning_content -#}
{%- set reasoning_content = '' -%}
{%- if message.reasoning is string -%}
{%- set reasoning_content = message.reasoning -%}
{%- elif message.reasoning_content is string -%}
{%- set reasoning_content = message.reasoning_content -%}
{%- endif -%}
{#- Display reasoning content for all messages if enable_thinking -#}
{%- if enable_thinking -%}
{{- '<think>' + reasoning_content + '</think>' -}}
{%- else -%}
{{- '</think>' -}}
{%- endif -%}
{#- Display main content (trailing newline only when no tool_calls follow) -#}
{%- if content -%}
{{- content -}}
{%- endif -%}
{%- if message.tool_calls -%}
{%- for tool_call in message.tool_calls -%}
{%- set function_data = tool_call.function -%}
{{- '<tool_call>' + function_data.name -}}
{%- set _args = function_data.arguments -%}
{%- for k, v in _args.items() -%}
{{- "<arg_key>" ~ k ~ "</arg_key>" -}}
{{- "<arg_value>" -}}{{- v | tojson(ensure_ascii=False) if v is not string else v -}}{{- "</arg_value>" -}}
{%- endfor -%}
{{- "</tool_call>" -}}
{%- endfor -%}
{%- endif -%}
{{- "</assistant>\n" -}}
{%- endgeneration -%}
{%- elif message.role == "tool" -%}
{{- "<tool_response>" + content + "</tool_response>\n" -}}
{%- elif message.role == "system" -%}
{#- Render additional system messages (the first one, if any, is handled separately in the header and was sliced off above) -#}
{{- "<system>" + content + "</system>\n" -}}
{%- endif -%}
{%- endfor -%}
{#- ───── generation prompt ───── -#}
{%- if add_generation_prompt -%}
{{- "<assistant>" -}}
{#- ───── Include reasoning mode directive ───── -#}
{%- if enable_thinking -%}
{{- '<think>' -}}
{%- else -%}
{{- '</think>' -}}
{%- endif -%}
{%- endif -%}

View File

@@ -380,6 +380,14 @@ func filesForModel(path string) ([]string, error) {
}
files = append(files, js...)
// Transformers stores a tokenizer's default template in this standalone
// file when it is not embedded in tokenizer_config.json.
chatTemplates, err := glob(filepath.Join(path, "chat_template.jinja"), "text/plain")
if err != nil {
return nil, err
}
files = append(files, chatTemplates...)
// add tokenizer.model if it exists (tokenizer.json is automatically picked up by the previous glob)
// tokenizer.model might be a unresolved git lfs reference; error if it is
if tks, _ := glob(filepath.Join(path, "tokenizer.model"), "application/octet-stream"); len(tks) > 0 {

View File

@@ -936,6 +936,7 @@ func TestFilesForModel(t *testing.T) {
"model-00002-of-00002.safetensors",
"config.json",
"tokenizer.json",
"chat_template.jinja",
}
for _, file := range files {
if err := os.WriteFile(filepath.Join(dir, file), []byte("test content"), 0o644); err != nil {
@@ -949,6 +950,7 @@ func TestFilesForModel(t *testing.T) {
"model-00002-of-00002.safetensors",
"config.json",
"tokenizer.json",
"chat_template.jinja",
},
},
{