mirror of
https://github.com/ollama/ollama.git
synced 2026-07-23 17:21:27 -05:00
Compare commits
8 Commits
v0.32.2-rc
...
main
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
fce745fe5e | ||
|
|
1fd1ccf7ad | ||
|
|
efb7e3c55e | ||
|
|
b517b9bd01 | ||
|
|
479664e7aa | ||
|
|
a51df81573 | ||
|
|
a18c230189 | ||
|
|
e21d5327b0 |
1
.github/workflows/release.yaml
vendored
1
.github/workflows/release.yaml
vendored
@@ -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"'
|
||||
|
||||
1
.github/workflows/test-llamacpp-update.yaml
vendored
1
.github/workflows/test-llamacpp-update.yaml
vendored
@@ -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"'
|
||||
|
||||
@@ -1 +1 @@
|
||||
b10069
|
||||
b10091
|
||||
|
||||
@@ -1 +1 @@
|
||||
b7c3dd6d27f45b5365b08a840310187dc503f1db
|
||||
33c03c486c34a7dadab5339563612c9933c4a406
|
||||
350
agent/skills.go
350
agent/skills.go
@@ -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).
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
8
agent/testdata/import/release-notes/SKILL.md
vendored
Normal file
8
agent/testdata/import/release-notes/SKILL.md
vendored
Normal file
@@ -0,0 +1,8 @@
|
||||
---
|
||||
name: release-notes
|
||||
description: Draft concise release notes.
|
||||
---
|
||||
|
||||
# Release notes
|
||||
|
||||
Use short bullets.
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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"} {
|
||||
|
||||
@@ -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 |
|
||||
| --- | --- |
|
||||
|
||||
@@ -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/
|
||||
```
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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,
|
||||
|
||||
144
integration/chat_cases_test.go
Normal file
144
integration/chat_cases_test.go
Normal 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+":")
|
||||
}
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
|
||||
204
integration/embed_cases_test.go
Normal file
204
integration/embed_cases_test.go
Normal 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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -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{
|
||||
|
||||
@@ -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)
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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?",
|
||||
|
||||
@@ -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") != "" {
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
}
|
||||
@@ -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),
|
||||
)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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",
|
||||
}
|
||||
|
||||
48
integration/reg_fast_test.go
Normal file
48
integration/reg_fast_test.go
Normal 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))
|
||||
}
|
||||
143
integration/reg_groups_test.go
Normal file
143
integration/reg_groups_test.go
Normal 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)
|
||||
}
|
||||
231
integration/reg_library_test.go
Normal file
231
integration/reg_library_test.go
Normal 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))
|
||||
}
|
||||
132
integration/reg_release_test.go
Normal file
132
integration/reg_release_test.go
Normal 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))
|
||||
}
|
||||
9
integration/reg_required_test.go
Normal file
9
integration/reg_required_test.go
Normal 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
140
integration/reg_test.go
Normal 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)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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 |
|
||||
|
||||
38
llama/compat/llama-ollama-compat.cpp
vendored
38
llama/compat/llama-ollama-compat.cpp
vendored
@@ -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);
|
||||
|
||||
42
llama/compat/models/003-llama-cpp-laguna-metal.patch
Normal file
42
llama/compat/models/003-llama-cpp-laguna-metal.patch
Normal 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);
|
||||
@@ -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 {
|
||||
232
llama/compat/models/laguna.cpp
vendored
232
llama/compat/models/laguna.cpp
vendored
@@ -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);
|
||||
}
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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))
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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",
|
||||
},
|
||||
},
|
||||
}),
|
||||
},
|
||||
},
|
||||
}}
|
||||
}
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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: ""},
|
||||
|
||||
132
model/renderers/testdata/laguna_v2_chat_template.jinja2
vendored
Normal file
132
model/renderers/testdata/laguna_v2_chat_template.jinja2
vendored
Normal 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 -%}
|
||||
93
model/renderers/testdata/laguna_v8_chat_template.jinja2
vendored
Normal file
93
model/renderers/testdata/laguna_v8_chat_template.jinja2
vendored
Normal 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 -%}
|
||||
@@ -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 {
|
||||
|
||||
@@ -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",
|
||||
},
|
||||
},
|
||||
{
|
||||
|
||||
Reference in New Issue
Block a user