diff --git a/.gitignore b/.gitignore index 15b3016..fd8cfff 100644 --- a/.gitignore +++ b/.gitignore @@ -19,3 +19,6 @@ build-cmake/ build-kf5/ build-kf6/ CMakeFiles/ + +# Large ML model files (downloaded during install) +data/models/ diff --git a/Makefile b/Makefile index d17e416..5674ae5 100644 --- a/Makefile +++ b/Makefile @@ -8,7 +8,7 @@ DATA_DIR ?= $(HOME)/.local/share/ollie BUILD_DIR := build JOBS ?= $(shell nproc 2>/dev/null || echo 2) -.PHONY: all build go tools kde install install-data test clean uninstall help +.PHONY: all build go tools kde install install-data install-models test clean uninstall help # Build everything, run tests, then install all: build test install @@ -55,7 +55,7 @@ test: go test ./... # Install everything from build dir -install: install-data +install: build install-data mkdir -p $(BINDIR) $(LIBDIR) install -m755 $(BUILD_DIR)/bin/olliesrv $(BINDIR)/olliesrv install -m755 $(BUILD_DIR)/bin/ollie-9p $(BINDIR)/ollie-9p @@ -72,20 +72,42 @@ install: install-data install -m755 kde/lib9p/libollie9p.so $(LIBDIR)/libollie9p.so; \ fi -# Install data files (agents, prompts, skills, workflows, scripts) -install-data: - mkdir -p $(CONFIG_DIR)/agents $(CONFIG_DIR)/prompts $(CONFIG_DIR)/skills $(CONFIG_DIR)/workflows $(CONFIG_DIR)/optmem +# Install data files (agents, prompts, skills, workflows, scripts, tools, models) +install-data: install-models + mkdir -p $(CONFIG_DIR)/agents $(CONFIG_DIR)/prompts $(CONFIG_DIR)/skills $(CONFIG_DIR)/workflows $(CONFIG_DIR)/optmem $(CONFIG_DIR)/tools @test -f $(CONFIG_DIR)/backends.conf || install -Dm600 data/backends.conf $(CONFIG_DIR)/backends.conf install -Dm755 third_party/optmem/memo $(CONFIG_DIR)/optmem/memo cp -a data/agents/. $(CONFIG_DIR)/agents/ cp -a data/prompts/. $(CONFIG_DIR)/prompts/ cp -a data/skills/. $(CONFIG_DIR)/skills/ + cp data/tools/*.meta $(CONFIG_DIR)/tools/ + for f in data/tools/*; do [ -f "$$f" ] && [ -x "$$f" ] && cp "$$f" $(CONFIG_DIR)/tools/; done; true + @test -d data/tools/_lib && cp -a data/tools/_lib $(CONFIG_DIR)/tools/ || true install -Dm755 data/workflows/* $(CONFIG_DIR)/workflows/ install -Dm644 cmd/toolsrv/internal/sandbox/sandbox.yaml $(CONFIG_DIR)/sandbox.yaml install -Dm755 data/scripts/ollie-remount $(BINDIR)/ollie-remount install -Dm755 data/scripts/logseq-cli $(BINDIR)/logseq-cli install -Dm755 data/scripts/o $(BINDIR)/o +# Install embedding model for skill matching +install-models: + mkdir -p $(DATA_DIR)/models + @if [ ! -f $(DATA_DIR)/models/model.onnx ]; then \ + echo "Downloading embedding model..."; \ + curl -fsSL -o $(DATA_DIR)/models/model.onnx \ + "https://huggingface.co/sentence-transformers/all-MiniLM-L6-v2/resolve/main/onnx/model.onnx"; \ + curl -fsSL -o $(DATA_DIR)/models/tokenizer.json \ + "https://huggingface.co/sentence-transformers/all-MiniLM-L6-v2/resolve/main/tokenizer.json"; \ + fi + @if [ ! -f $(DATA_DIR)/models/libonnxruntime.so ]; then \ + echo "Downloading ONNX runtime..."; \ + curl -fsSL -o /tmp/onnxruntime.tgz \ + "https://github.com/microsoft/onnxruntime/releases/download/v1.29.0/onnxruntime-linux-x64-1.29.0.tgz"; \ + tar -xzf /tmp/onnxruntime.tgz -C /tmp; \ + cp /tmp/onnxruntime-linux-x64-1.29.0/lib/libonnxruntime.so.1.29.0 $(DATA_DIR)/models/libonnxruntime.so; \ + rm -rf /tmp/onnxruntime.tgz /tmp/onnxruntime-linux-x64-1.29.0; \ + fi + # Remove installed files uninstall: rm -f $(BINDIR)/olliesrv $(BINDIR)/ollie-9p diff --git a/cmd/olliesrv/internal/agent/skill_match.go b/cmd/olliesrv/internal/agent/skill_match.go new file mode 100644 index 0000000..bc53fe9 --- /dev/null +++ b/cmd/olliesrv/internal/agent/skill_match.go @@ -0,0 +1,168 @@ +package agent + +import ( + "encoding/json" + "os" + "path/filepath" + "strings" + "sync" + + "ollie/embedding" + "ollie/skills" + "ollie/util" +) + +// Skill matching configuration. +const ( + // skillMatchThreshold is the minimum cosine similarity to include a skill. + skillMatchThreshold = 0.35 + // skillMatchLimit is the maximum number of skills to inject per turn. + skillMatchLimit = 3 + + // toolMatchThreshold is the minimum cosine similarity to include a tool hint. + toolMatchThreshold = 0.35 + // toolMatchLimit is the maximum number of tool hints to inject per turn. + toolMatchLimit = 5 +) + +var ( + skillIndex *skills.Index + skillInitErr error + skillOnce sync.Once + + toolIndex *skills.Index + toolInitErr error + toolOnce sync.Once +) + +// matchSkills finds skills relevant to the user's input and returns +// their content formatted for injection. Initializes the skill index +// on first call. +func matchSkills(input string) string { + skillOnce.Do(func() { + idx, err := skills.NewIndex(skills.DefaultModelDir(), skills.DefaultSkillDirs()) + if err != nil { + skillInitErr = err + return + } + skillIndex = idx + }) + + if skillIndex == nil { + return "" + } + + results, err := skillIndex.Match(input, skillMatchThreshold, skillMatchLimit) + if err != nil { + return "" + } + if len(results) == 0 { + return "" + } + + var sb strings.Builder + sb.WriteString("\n") + for _, r := range results { + sb.WriteString("\n") + // Include full skill content (already includes frontmatter) + sb.WriteString(r.Skill.Content) + if !strings.HasSuffix(r.Skill.Content, "\n") { + sb.WriteString("\n") + } + sb.WriteString("\n") + } + sb.WriteString("\n") + return sb.String() +} + +// matchTools finds tools relevant to the user's input and returns +// a hint block telling the model to call them. +func matchTools(input string) string { + toolOnce.Do(func() { + idx, err := newToolIndex() + if err != nil { + toolInitErr = err + return + } + toolIndex = idx + }) + + if toolIndex == nil { + return "" + } + + results, err := toolIndex.Match(input, toolMatchThreshold, toolMatchLimit) + if err != nil { + return "" + } + if len(results) == 0 { + return "" + } + + var sb strings.Builder + sb.WriteString("\n") + sb.WriteString("Relevant tools for this request (auto-load on first call):\n\n") + for _, r := range results { + sb.WriteString("→ `") + sb.WriteString(r.Skill.Name) + sb.WriteString("({...})` — ") + sb.WriteString(r.Skill.Description) + sb.WriteString("\n") + } + sb.WriteString("\nDo NOT say \"I don't have this tool\" — just CALL IT.\n") + sb.WriteString("\n") + return sb.String() +} + +func formatScore(s float32) string { + // Format as percentage + pct := int(s * 100) + if pct > 99 { + pct = 99 + } + return string([]byte{'0' + byte(pct/10), '0' + byte(pct%10), '%'}) +} + +// newToolIndex creates a skill-compatible index from tool .meta files. +func newToolIndex() (*skills.Index, error) { + model, err := embedding.LoadModel(skills.DefaultModelDir()) + if err != nil { + return nil, err + } + + toolsDir := filepath.Join(util.CfgDir(), "tools") + entries, err := os.ReadDir(toolsDir) + if err != nil { + model.Close() + return nil, err + } + + var toolSkills []skills.Skill + for _, entry := range entries { + if entry.IsDir() || !strings.HasSuffix(entry.Name(), ".meta") { + continue + } + metaPath := filepath.Join(toolsDir, entry.Name()) + data, err := os.ReadFile(metaPath) + if err != nil { + continue + } + var meta struct { + Description string `json:"description"` + } + if json.Unmarshal(data, &meta) != nil || meta.Description == "" { + continue + } + toolName := strings.TrimSuffix(entry.Name(), ".meta") + toolSkills = append(toolSkills, skills.Skill{ + Name: toolName, + Description: meta.Description, + }) + } + + return skills.NewIndexFromSkills(model, toolSkills) +} diff --git a/cmd/olliesrv/internal/agent/turn.go b/cmd/olliesrv/internal/agent/turn.go index a8c2a79..f163be7 100644 --- a/cmd/olliesrv/internal/agent/turn.go +++ b/cmd/olliesrv/internal/agent/turn.go @@ -95,9 +95,19 @@ func (ag *Agent) executeTurn(ctx context.Context, input string) string { input = "Run the memory_wake tool now. Follow its output completely before addressing my request.\n\n" + input } - // Prepend resolved user prompts (active global rules) to the user input. + // Build context block: user prompts + matched skills + tool hints + var contextParts []string if ag.runtime.UserPrompt != "" { - input = "\n" + ag.runtime.UserPrompt + "\n\n\n" + input + contextParts = append(contextParts, ag.runtime.UserPrompt) + } + if toolHints := matchTools(input); toolHints != "" { + contextParts = append(contextParts, toolHints) + } + if skillContent := matchSkills(input); skillContent != "" { + contextParts = append(contextParts, skillContent) + } + if len(contextParts) > 0 { + input = "\n" + strings.Join(contextParts, "\n") + "\n\n\n" + input } ag.emit(Event{Role: "user", Content: input}) diff --git a/data/tools/memory_nap.meta b/data/tools/memory_nap.meta index 30decba..a5e2a3a 100644 --- a/data/tools/memory_nap.meta +++ b/data/tools/memory_nap.meta @@ -1,13 +1,22 @@ { - "description": "Submit a memory compression (nap) to OptMem.", + "description": "Compress memories when prompted. Submit the compression text.", "prompt": "## memory_nap\n\nSubmit a compression for a memory block range. Called when memory_wake or memory_remember prints a compression prompt.\n\n**Calling convention:**\n```\nmemory_nap(range=\"0-1\", text=\"compressed one-line summary\")\n```\n- `range`: block range as printed by the compression prompt (e.g. \"0-1\")\n- `text`: the compressed one-line summary (max 280 bytes)\n\nIf more compressions remain, the tool prints the next one.", "cmd": "input=$(cat); range=$(printf '%s' \"$input\" | jq -er '.range'); text=$(printf '%s' \"$input\" | jq -er '.text'); MEMO_TOOLS=1 MEMORY_DIR=\"${XDG_DATA_HOME:-$HOME/.local/share}/ollie/optmem\" exec \"${XDG_CONFIG_HOME:-$HOME/.config}/ollie/optmem/memo\" nap \"$range\" \"$text\"", "args": { "type": "object", - "required": ["range", "text"], + "required": [ + "range", + "text" + ], "properties": { - "range": {"type": "string", "description": "Block range (e.g. \"0-1\")"}, - "text": {"type": "string", "description": "Compressed one-line summary (max 280 bytes)"} + "range": { + "type": "string", + "description": "Block range (e.g. \"0-1\")" + }, + "text": { + "type": "string", + "description": "Compressed one-line summary (max 280 bytes)" + } } }, "scope": "global" diff --git a/data/tools/memory_recall.meta b/data/tools/memory_recall.meta index 657d8f5..0aa529d 100644 --- a/data/tools/memory_recall.meta +++ b/data/tools/memory_recall.meta @@ -1,12 +1,17 @@ { - "description": "Search stored memories for relevant context.", + "description": "Search, recall, or retrieve saved memories. Find what you stored previously.", "prompt": "## memory_recall\n\nSearch stored memories for relevant context using OptMem's bounded memory index.\n\n**Calling convention:**\n```\nmemory_recall(query=\"keyword\")\n```\n- `query`: regular expression or short search term.\n\n**Returns:** matching memory records from the persistent OptMem store.", "cmd": "input=$(cat); query=$(printf '%s' \"$input\" | jq -er '.query'); MEMO_TOOLS=1 MEMORY_DIR=\"${XDG_DATA_HOME:-$HOME/.local/share}/ollie/optmem\" exec \"${XDG_CONFIG_HOME:-$HOME/.config}/ollie/optmem/memo\" recall \"$query\"", "args": { "type": "object", - "required": ["query"], + "required": [ + "query" + ], "properties": { - "query": {"type": "string", "description": "Search keyword or regular expression"} + "query": { + "type": "string", + "description": "Search keyword or regular expression" + } } }, "tier": "cold", diff --git a/data/tools/memory_remember.meta b/data/tools/memory_remember.meta index ce5fa8a..a1338a8 100644 --- a/data/tools/memory_remember.meta +++ b/data/tools/memory_remember.meta @@ -1,14 +1,27 @@ { - "description": "Persist a fact that would otherwise be lost when the session ends.", + "description": "Save, store, or remember a fact for later. Persists across sessions.", "prompt": "## memory_remember\n\nPersist a durable fact in OptMem's append-only memory store.\n\n**Calling convention:**\n```\nmemory_remember(title=\"...\", tags=\"...\", body=\"...\")\n```\n- `title`: short noun phrase\n- `tags`: comma-separated tags\n- `body`: one standalone fact, no more than 280 bytes\n\nThe title, tags, and body are stored as one searchable memory record.", "cmd": "input=$(cat); title=$(printf '%s' \"$input\" | jq -er '.title'); tags=$(printf '%s' \"$input\" | jq -er '.tags'); body=$(printf '%s' \"$input\" | jq -er '.body'); record=\"[$tags] $title: $body\"; bytes=$(printf '%s' \"$record\" | wc -c); [ \"$bytes\" -le 280 ] || { printf 'memory exceeds OptMem limit: %s bytes (maximum 280)\\n' \"$bytes\" >&2; exit 1; }; MEMO_TOOLS=1 MEMORY_DIR=\"${XDG_DATA_HOME:-$HOME/.local/share}/ollie/optmem\" exec \"${XDG_CONFIG_HOME:-$HOME/.config}/ollie/optmem/memo\" note \"$record\"", "args": { "type": "object", - "required": ["title", "tags", "body"], + "required": [ + "title", + "tags", + "body" + ], "properties": { - "title": {"type": "string", "description": "Short noun phrase"}, - "tags": {"type": "string", "description": "Comma-separated tags"}, - "body": {"type": "string", "description": "Standalone fact, maximum 280 bytes after formatting"} + "title": { + "type": "string", + "description": "Short noun phrase" + }, + "tags": { + "type": "string", + "description": "Comma-separated tags" + }, + "body": { + "type": "string", + "description": "Standalone fact, maximum 280 bytes after formatting" + } } }, "scope": "global" diff --git a/data/tools/memory_wake.meta b/data/tools/memory_wake.meta index 418b6db..9eed65a 100644 --- a/data/tools/memory_wake.meta +++ b/data/tools/memory_wake.meta @@ -1,12 +1,18 @@ { - "description": "Load the bounded OptMem context at session startup.", + "description": "Wake up and load your persistent memory at session start. Call this first in every session.", "prompt": "## memory_wake\n\nLoad the bounded persistent memory context. Run this once before other tools at the start of every top-level session. Follow any printed OptMem compression instruction before continuing. Do not run from a sub-agent.", "cmd": "input=$(cat); part=$(printf '%s' \"$input\" | jq -r '.part // empty'); T=$(printf '%s' \"$input\" | jq -r '.T // empty'); args=''; [ -n \"$part\" ] && args=\"$part\"; [ -n \"$T\" ] && args=\"$args $T\"; MEMO_TOOLS=1 MEMORY_DIR=\"${XDG_DATA_HOME:-$HOME/.local/share}/ollie/optmem\" exec \"${XDG_CONFIG_HOME:-$HOME/.config}/ollie/optmem/memo\" wake $args", "args": { "type": "object", "properties": { - "part": {"type": "integer", "description": "Page number for paginated memory (default: 1)"}, - "T": {"type": "integer", "description": "Memory snapshot count (used with part for pagination)"} + "part": { + "type": "integer", + "description": "Page number for paginated memory (default: 1)" + }, + "T": { + "type": "integer", + "description": "Memory snapshot count (used with part for pagination)" + } } }, "tier": "cold", diff --git a/data/tools/memory_zoom.meta b/data/tools/memory_zoom.meta index 8e2fb99..e4d2a0f 100644 --- a/data/tools/memory_zoom.meta +++ b/data/tools/memory_zoom.meta @@ -1,12 +1,17 @@ { - "description": "Expand a memory tree node into its two halves.", + "description": "Expand a memory summary to see its details or raw memories.", "prompt": "## memory_zoom\n\nOpen a memory tree node into its two halves, down to raw memories.\n\n**Calling convention:**\n```\nmemory_zoom(range=\"0-3\")\n```\n- `range`: block range as printed by memory_wake (e.g. \"0-3\")\n\nReturns the two child nodes (summaries or raw memories).", "cmd": "input=$(cat); range=$(printf '%s' \"$input\" | jq -er '.range'); MEMO_TOOLS=1 MEMORY_DIR=\"${XDG_DATA_HOME:-$HOME/.local/share}/ollie/optmem\" exec \"${XDG_CONFIG_HOME:-$HOME/.config}/ollie/optmem/memo\" zoom \"$range\"", "args": { "type": "object", - "required": ["range"], + "required": [ + "range" + ], "properties": { - "range": {"type": "string", "description": "Block range to expand (e.g. \"0-3\")"} + "range": { + "type": "string", + "description": "Block range to expand (e.g. \"0-3\")" + } } }, "tier": "cold", diff --git a/doc/embedding.conf.sample b/doc/embedding.conf.sample new file mode 100644 index 0000000..ab1f2c2 --- /dev/null +++ b/doc/embedding.conf.sample @@ -0,0 +1,7 @@ +# Skill directories for embedding-based matching (one per line). +# Paths are searched in order; first match wins for duplicate skill names. +# Supports ~ and $VAR expansion. Lines starting with # are comments. +# +# Default (if this file is absent): ~/.config/ollie/skills + +~/.config/ollie/skills diff --git a/embedding/embedding.go b/embedding/embedding.go new file mode 100644 index 0000000..2d3e463 --- /dev/null +++ b/embedding/embedding.go @@ -0,0 +1,358 @@ +// Package embedding provides text embedding using ONNX-based sentence transformers. +// It loads an all-MiniLM-L6-v2 model and tokenizer, and provides functions to +// embed text and compute similarity scores. +package embedding + +import ( + "encoding/json" + "fmt" + "math" + "os" + "path/filepath" + "strings" + "sync" + "unicode" + + ort "github.com/yalue/onnxruntime_go" +) + +// Model holds the ONNX model and tokenizer state. +type Model struct { + session *ort.DynamicAdvancedSession + tokenizer *Tokenizer + mu sync.Mutex +} + +// Vector is a 384-dimensional embedding vector (MiniLM output size). +type Vector []float32 + +var initOnce sync.Once +var initErr error + +// LoadModel loads the ONNX model and tokenizer from the given directory. +// The directory must contain model.onnx, tokenizer.json, and libonnxruntime.so. +func LoadModel(modelDir string) (*Model, error) { + modelPath := filepath.Join(modelDir, "model.onnx") + tokenizerPath := filepath.Join(modelDir, "tokenizer.json") + libPath := filepath.Join(modelDir, "libonnxruntime.so") + + // Initialize ONNX runtime once + initOnce.Do(func() { + ort.SetSharedLibraryPath(libPath) + initErr = ort.InitializeEnvironment() + }) + if initErr != nil { + return nil, fmt.Errorf("init onnx environment: %w", initErr) + } + + // Load tokenizer + tok, err := loadTokenizer(tokenizerPath) + if err != nil { + return nil, fmt.Errorf("load tokenizer: %w", err) + } + + // Create ONNX session + inputs := []string{"input_ids", "attention_mask", "token_type_ids"} + outputs := []string{"last_hidden_state"} + + session, err := ort.NewDynamicAdvancedSession(modelPath, inputs, outputs, nil) + if err != nil { + return nil, fmt.Errorf("create onnx session: %w", err) + } + + return &Model{ + session: session, + tokenizer: tok, + }, nil +} + +// Close releases model resources. +func (m *Model) Close() error { + m.mu.Lock() + defer m.mu.Unlock() + if m.session != nil { + return m.session.Destroy() + } + return nil +} + +// Embed returns the embedding vector for the given text. +func (m *Model) Embed(text string) (Vector, error) { + m.mu.Lock() + defer m.mu.Unlock() + + // Tokenize + inputIDs, attentionMask, tokenTypeIDs := m.tokenizer.Encode(text) + seqLen := int64(len(inputIDs)) + + // Create input tensors with shape [1, seqLen] + shape := ort.Shape{1, seqLen} + + inputIDsTensor, err := ort.NewTensor(shape, toInt64(inputIDs)) + if err != nil { + return nil, fmt.Errorf("create input_ids tensor: %w", err) + } + defer inputIDsTensor.Destroy() + + attentionTensor, err := ort.NewTensor(shape, toInt64(attentionMask)) + if err != nil { + return nil, fmt.Errorf("create attention_mask tensor: %w", err) + } + defer attentionTensor.Destroy() + + tokenTypeTensor, err := ort.NewTensor(shape, toInt64(tokenTypeIDs)) + if err != nil { + return nil, fmt.Errorf("create token_type_ids tensor: %w", err) + } + defer tokenTypeTensor.Destroy() + + // Create output tensor with shape [1, seqLen, 384] + outputShape := ort.Shape{1, seqLen, 384} + outputTensor, err := ort.NewEmptyTensor[float32](outputShape) + if err != nil { + return nil, fmt.Errorf("create output tensor: %w", err) + } + defer outputTensor.Destroy() + + // Run inference + err = m.session.Run( + []ort.ArbitraryTensor{inputIDsTensor, attentionTensor, tokenTypeTensor}, + []ort.ArbitraryTensor{outputTensor}, + ) + if err != nil { + return nil, fmt.Errorf("run inference: %w", err) + } + + // Mean pooling: average over sequence length, taking attention mask into account + output := outputTensor.GetData() + return meanPool(output, attentionMask, int(seqLen), 384), nil +} + +// EmbedBatch embeds multiple texts efficiently. +func (m *Model) EmbedBatch(texts []string) ([]Vector, error) { + results := make([]Vector, len(texts)) + for i, text := range texts { + vec, err := m.Embed(text) + if err != nil { + return nil, fmt.Errorf("embed text %d: %w", i, err) + } + results[i] = vec + } + return results, nil +} + +// CosineSimilarity computes the cosine similarity between two vectors. +func CosineSimilarity(a, b Vector) float32 { + if len(a) != len(b) { + return 0 + } + var dot, normA, normB float64 + for i := range a { + dot += float64(a[i]) * float64(b[i]) + normA += float64(a[i]) * float64(a[i]) + normB += float64(b[i]) * float64(b[i]) + } + if normA == 0 || normB == 0 { + return 0 + } + return float32(dot / (math.Sqrt(normA) * math.Sqrt(normB))) +} + +// meanPool performs mean pooling over the sequence dimension with attention masking. +func meanPool(output []float32, mask []int, seqLen, hiddenSize int) Vector { + result := make(Vector, hiddenSize) + var count float32 + for i := 0; i < seqLen; i++ { + if mask[i] == 0 { + continue + } + count++ + for j := 0; j < hiddenSize; j++ { + result[j] += output[i*hiddenSize+j] + } + } + if count > 0 { + for j := range result { + result[j] /= count + } + } + // L2 normalize + var norm float64 + for _, v := range result { + norm += float64(v) * float64(v) + } + norm = math.Sqrt(norm) + if norm > 0 { + for j := range result { + result[j] = float32(float64(result[j]) / norm) + } + } + return result +} + +func toInt64(ints []int) []int64 { + out := make([]int64, len(ints)) + for i, v := range ints { + out[i] = int64(v) + } + return out +} + +// --- Tokenizer --- + +// Tokenizer handles WordPiece tokenization for BERT-style models. +type Tokenizer struct { + vocab map[string]int + maxLen int + clsID int + sepID int + padID int + unkID int + lowercase bool +} + +// tokenizerJSON is the HuggingFace tokenizers JSON format. +type tokenizerJSON struct { + Truncation *struct { + MaxLength int `json:"max_length"` + } `json:"truncation"` + Normalizer *struct { + Lowercase bool `json:"lowercase"` + } `json:"normalizer"` + Model struct { + Vocab map[string]int `json:"vocab"` + } `json:"model"` + AddedTokens []struct { + ID int `json:"id"` + Content string `json:"content"` + } `json:"added_tokens"` +} + +func loadTokenizer(path string) (*Tokenizer, error) { + data, err := os.ReadFile(path) + if err != nil { + return nil, err + } + var tj tokenizerJSON + if err := json.Unmarshal(data, &tj); err != nil { + return nil, err + } + + tok := &Tokenizer{ + vocab: tj.Model.Vocab, + maxLen: 128, // default + unkID: 100, // [UNK] + clsID: 101, // [CLS] + sepID: 102, // [SEP] + padID: 0, // [PAD] + } + if tj.Truncation != nil { + tok.maxLen = tj.Truncation.MaxLength + } + if tj.Normalizer != nil { + tok.lowercase = tj.Normalizer.Lowercase + } + + // Override IDs from added_tokens if present + for _, at := range tj.AddedTokens { + switch at.Content { + case "[PAD]": + tok.padID = at.ID + case "[UNK]": + tok.unkID = at.ID + case "[CLS]": + tok.clsID = at.ID + case "[SEP]": + tok.sepID = at.ID + } + } + + return tok, nil +} + +// Encode tokenizes text and returns input_ids, attention_mask, and token_type_ids. +func (t *Tokenizer) Encode(text string) (inputIDs, attentionMask, tokenTypeIDs []int) { + if t.lowercase { + text = strings.ToLower(text) + } + + // Basic whitespace + punctuation tokenization + words := tokenizeBasic(text) + + // WordPiece tokenization + var tokens []int + tokens = append(tokens, t.clsID) + for _, word := range words { + wordTokens := t.tokenizeWord(word) + tokens = append(tokens, wordTokens...) + } + tokens = append(tokens, t.sepID) + + // Truncate if needed (keep [CLS] and [SEP]) + if len(tokens) > t.maxLen { + tokens = append(tokens[:t.maxLen-1], t.sepID) + } + + // Build masks + seqLen := len(tokens) + inputIDs = tokens + attentionMask = make([]int, seqLen) + tokenTypeIDs = make([]int, seqLen) + for i := 0; i < seqLen; i++ { + attentionMask[i] = 1 + tokenTypeIDs[i] = 0 + } + + return inputIDs, attentionMask, tokenTypeIDs +} + +func (t *Tokenizer) tokenizeWord(word string) []int { + var tokens []int + remaining := word + for len(remaining) > 0 { + found := false + for end := len(remaining); end > 0; end-- { + subword := remaining[:end] + if len(tokens) > 0 { + subword = "##" + subword + } + if id, ok := t.vocab[subword]; ok { + tokens = append(tokens, id) + remaining = remaining[end:] + found = true + break + } + } + if !found { + tokens = append(tokens, t.unkID) + break + } + } + return tokens +} + +// tokenizeBasic splits on whitespace and punctuation. +func tokenizeBasic(text string) []string { + var words []string + var current strings.Builder + for _, r := range text { + if unicode.IsSpace(r) { + if current.Len() > 0 { + words = append(words, current.String()) + current.Reset() + } + } else if unicode.IsPunct(r) { + if current.Len() > 0 { + words = append(words, current.String()) + current.Reset() + } + words = append(words, string(r)) + } else { + current.WriteRune(r) + } + } + if current.Len() > 0 { + words = append(words, current.String()) + } + return words +} diff --git a/embedding/embedding_test.go b/embedding/embedding_test.go new file mode 100644 index 0000000..a0fcfd4 --- /dev/null +++ b/embedding/embedding_test.go @@ -0,0 +1,129 @@ +package embedding + +import ( + "os" + "path/filepath" + "testing" +) + +func TestEmbedAndSimilarity(t *testing.T) { + // Find model directory + modelDir := os.Getenv("OLLIE_MODEL_DIR") + if modelDir == "" { + // Try common locations + home, _ := os.UserHomeDir() + candidates := []string{ + "../data/models", + filepath.Join(home, ".local/share/ollie/models"), + } + for _, c := range candidates { + if _, err := os.Stat(filepath.Join(c, "model.onnx")); err == nil { + modelDir = c + break + } + } + } + if modelDir == "" { + t.Skip("model not found, set OLLIE_MODEL_DIR") + } + + model, err := LoadModel(modelDir) + if err != nil { + t.Fatalf("LoadModel: %v", err) + } + defer model.Close() + + // Test embedding + vec, err := model.Embed("hello world") + if err != nil { + t.Fatalf("Embed: %v", err) + } + if len(vec) != 384 { + t.Fatalf("expected 384 dimensions, got %d", len(vec)) + } + + // Test similarity - similar sentences should have high similarity + v1, _ := model.Embed("I love programming in Go") + v2, _ := model.Embed("Go is my favorite programming language") + v3, _ := model.Embed("The weather is nice today") + + sim12 := CosineSimilarity(v1, v2) + sim13 := CosineSimilarity(v1, v3) + + t.Logf("Similar sentences: %.4f", sim12) + t.Logf("Dissimilar sentences: %.4f", sim13) + + if sim12 < 0.5 { + t.Errorf("expected similar sentences to have similarity > 0.5, got %.4f", sim12) + } + if sim13 > sim12 { + t.Errorf("expected similar sentences to have higher similarity than dissimilar") + } +} + +func TestSkillMatching(t *testing.T) { + modelDir := os.Getenv("OLLIE_MODEL_DIR") + if modelDir == "" { + home, _ := os.UserHomeDir() + candidates := []string{ + "../data/models", + filepath.Join(home, ".local/share/ollie/models"), + } + for _, c := range candidates { + if _, err := os.Stat(filepath.Join(c, "model.onnx")); err == nil { + modelDir = c + break + } + } + } + if modelDir == "" { + t.Skip("model not found") + } + + model, err := LoadModel(modelDir) + if err != nil { + t.Fatalf("LoadModel: %v", err) + } + defer model.Close() + + // Simulate skill descriptions + skills := []string{ + "Interact with pascom Atlassian (Jira/Confluence). Use for issues, pages, and project management.", + "Interact with GitHub repositories, issues, PRs, workflows using the gh CLI.", + "Execute bash commands on remote dev server with synced mobydick workspace.", + "Write ast-grep rules for AST-based structural code search and analysis.", + } + + skillVecs, err := model.EmbedBatch(skills) + if err != nil { + t.Fatalf("EmbedBatch: %v", err) + } + + // Test queries + queries := []struct { + query string + expected int // index of expected best match + }{ + {"read jira ticket PR-12345", 0}, // should match Atlassian + {"create a github issue", 1}, // should match GitHub + {"run make on the remote server", 2}, // should match remote-bash + {"find all function calls in Go", 3}, // should match ast-grep + } + + for _, tc := range queries { + qvec, _ := model.Embed(tc.query) + best := -1 + bestSim := float32(-1) + for i, sv := range skillVecs { + sim := CosineSimilarity(qvec, sv) + t.Logf("%q vs skill[%d]: %.4f", tc.query, i, sim) + if sim > bestSim { + bestSim = sim + best = i + } + } + if best != tc.expected { + t.Errorf("%q: expected skill %d, got %d (sim=%.4f)", tc.query, tc.expected, best, bestSim) + } + } +} diff --git a/go.mod b/go.mod index 828888d..fb14a42 100644 --- a/go.mod +++ b/go.mod @@ -24,6 +24,8 @@ require ( ollie/virtfs v0.0.0 ) +require github.com/yalue/onnxruntime_go v1.35.0 // indirect + require ( github.com/PuerkitoBio/goquery v1.9.2 // indirect github.com/andybalholm/cascadia v1.3.2 // indirect diff --git a/go.sum b/go.sum index 6de9381..6311f6e 100644 --- a/go.sum +++ b/go.sum @@ -69,6 +69,8 @@ github.com/tree-sitter/tree-sitter-rust v0.24.2 h1:NL4nF67ib21RMzzfvkmXlVwe45vvh github.com/tree-sitter/tree-sitter-rust v0.24.2/go.mod h1:hfeGWic9BAfgTrc7Xf6FaOAguCFJRo3RBbs7QJ6D7MI= github.com/tree-sitter/tree-sitter-typescript v0.23.2 h1:/Odvphn18PniVixb9e97X0DbNVsU6Qocv9mfkyzdXwU= github.com/tree-sitter/tree-sitter-typescript v0.23.2/go.mod h1:zjzMXT/Ulffel2xfOcAkQQkiAkmgnbtPGlFQw/5X4xA= +github.com/yalue/onnxruntime_go v1.35.0 h1:IEIqLmh1r2LfN4U4hksRPh0711t3d4a5FQi95TzRQ4I= +github.com/yalue/onnxruntime_go v1.35.0/go.mod h1:b4X26A8pekNb1ACJ58wAXgNKeUCGEAQ9dmACut9Sm/4= github.com/yuin/goldmark v1.4.13/go.mod h1:6yULJ656Px+3vBD8DxQVa3kxgyrAnzto9xy5taEt/CY= github.com/yuin/goldmark v1.7.1 h1:3bajkSilaCbjdKVsKdZjZCLBNPL9pYzrCakKaf4U49U= github.com/yuin/goldmark v1.7.1/go.mod h1:uzxRWxtg69N339t3louHJ7+O03ezfj6PlliRlaOzY1E= diff --git a/skills/skills.go b/skills/skills.go new file mode 100644 index 0000000..333280e --- /dev/null +++ b/skills/skills.go @@ -0,0 +1,265 @@ +// Package skills provides skill discovery, embedding, and matching. +package skills + +import ( + "bufio" + "bytes" + "fmt" + "os" + "path/filepath" + "sort" + "strings" + "sync" + + "ollie/embedding" + "ollie/util" +) + +// Skill represents a discovered skill with its metadata and content. +type Skill struct { + Name string + Description string + Path string // path to SKILL.md + Content string // full content (including frontmatter) +} + +// Index holds precomputed skill embeddings for fast matching. +type Index struct { + model *embedding.Model + skills []Skill + vecs []embedding.Vector + mu sync.RWMutex +} + +// NewIndex creates a new skill index with precomputed embeddings. +// modelDir is the path to the directory containing the ONNX model. +// skillDirs is a list of directories to scan for skills. +func NewIndex(modelDir string, skillDirs []string) (*Index, error) { + model, err := embedding.LoadModel(modelDir) + if err != nil { + return nil, fmt.Errorf("load embedding model: %w", err) + } + + idx := &Index{model: model} + if err := idx.loadSkills(skillDirs); err != nil { + model.Close() + return nil, err + } + return idx, nil +} + +// NewIndexFromSkills creates an index from pre-loaded skills. +// Takes ownership of the model. +func NewIndexFromSkills(model *embedding.Model, skills []Skill) (*Index, error) { + idx := &Index{model: model, skills: skills} + + // Compute embeddings + vecs := make([]embedding.Vector, len(skills)) + for i, skill := range skills { + vec, err := model.Embed(skill.Description) + if err != nil { + return nil, fmt.Errorf("embed skill %q: %w", skill.Name, err) + } + vecs[i] = vec + } + idx.vecs = vecs + + return idx, nil +} + +// Close releases resources. +func (idx *Index) Close() error { + if idx.model != nil { + return idx.model.Close() + } + return nil +} + +// Match returns skills matching the query, sorted by relevance. +// threshold is the minimum cosine similarity (0-1) to include. +// limit is the maximum number of skills to return (0 for no limit). +func (idx *Index) Match(query string, threshold float32, limit int) ([]MatchResult, error) { + idx.mu.RLock() + defer idx.mu.RUnlock() + + if len(idx.skills) == 0 { + return nil, nil + } + + qvec, err := idx.model.Embed(query) + if err != nil { + return nil, fmt.Errorf("embed query: %w", err) + } + + var results []MatchResult + for i, skill := range idx.skills { + sim := embedding.CosineSimilarity(qvec, idx.vecs[i]) + if sim >= threshold { + results = append(results, MatchResult{ + Skill: skill, + Score: sim, + }) + } + } + + // Sort by score descending + sort.Slice(results, func(i, j int) bool { + return results[i].Score > results[j].Score + }) + + if limit > 0 && len(results) > limit { + results = results[:limit] + } + return results, nil +} + +// MatchResult holds a matched skill and its similarity score. +type MatchResult struct { + Skill Skill + Score float32 +} + +// All returns all indexed skills. +func (idx *Index) All() []Skill { + idx.mu.RLock() + defer idx.mu.RUnlock() + return append([]Skill(nil), idx.skills...) +} + +// Reload rescans skill directories and updates embeddings. +func (idx *Index) Reload(skillDirs []string) error { + idx.mu.Lock() + defer idx.mu.Unlock() + return idx.loadSkillsLocked(skillDirs) +} + +func (idx *Index) loadSkills(dirs []string) error { + idx.mu.Lock() + defer idx.mu.Unlock() + return idx.loadSkillsLocked(dirs) +} + +func (idx *Index) loadSkillsLocked(dirs []string) error { + var skills []Skill + seen := make(map[string]bool) + + for _, dir := range dirs { + dir = util.ExpandHome(dir) + entries, err := os.ReadDir(dir) + if err != nil { + continue // skip missing directories + } + for _, entry := range entries { + if !entry.IsDir() { + continue + } + name := entry.Name() + if seen[name] { + continue // first dir wins + } + skillPath := filepath.Join(dir, name, "SKILL.md") + skill, err := loadSkill(skillPath) + if err != nil { + continue // skip invalid skills + } + seen[name] = true + skills = append(skills, skill) + } + } + + // Compute embeddings + vecs := make([]embedding.Vector, len(skills)) + for i, skill := range skills { + vec, err := idx.model.Embed(skill.Description) + if err != nil { + return fmt.Errorf("embed skill %q: %w", skill.Name, err) + } + vecs[i] = vec + } + + idx.skills = skills + idx.vecs = vecs + return nil +} + +func loadSkill(path string) (Skill, error) { + data, err := os.ReadFile(path) + if err != nil { + return Skill{}, err + } + + name, desc, err := parseFrontmatter(data) + if err != nil { + return Skill{}, err + } + + return Skill{ + Name: name, + Description: desc, + Path: path, + Content: string(data), + }, nil +} + +// parseFrontmatter extracts name and description from YAML frontmatter. +func parseFrontmatter(data []byte) (name, description string, err error) { + scanner := bufio.NewScanner(bytes.NewReader(data)) + + // First line must be --- + if !scanner.Scan() || strings.TrimSpace(scanner.Text()) != "---" { + return "", "", fmt.Errorf("missing frontmatter") + } + + // Read until closing --- + for scanner.Scan() { + line := scanner.Text() + if strings.TrimSpace(line) == "---" { + break + } + if key, val, ok := strings.Cut(line, ":"); ok { + key = strings.TrimSpace(key) + val = strings.TrimSpace(val) + switch key { + case "name": + name = val + case "description": + description = val + } + } + } + + if name == "" { + return "", "", fmt.Errorf("missing name in frontmatter") + } + if description == "" { + return "", "", fmt.Errorf("missing description in frontmatter") + } + return name, description, nil +} + +// DefaultSkillDirs returns the skill directories to search. +// Reads from ~/.config/ollie/embedding.conf if it exists (one path per line). +// Supports ~ and $VAR expansion. Otherwise, defaults to ~/.config/ollie/skills. +func DefaultSkillDirs() []string { + confPath := filepath.Join(util.CfgDir(), "embedding.conf") + data, err := os.ReadFile(confPath) + if err == nil { + var dirs []string + for _, line := range strings.Split(string(data), "\n") { + line = strings.TrimSpace(line) + if line != "" && !strings.HasPrefix(line, "#") { + line = os.ExpandEnv(line) + dirs = append(dirs, util.ExpandHome(line)) + } + } + if len(dirs) > 0 { + return dirs + } + } + return []string{filepath.Join(util.CfgDir(), "skills")} +} + +// DefaultModelDir returns the default model directory. +func DefaultModelDir() string { + return filepath.Join(util.DataDir(), "models") +} diff --git a/skills/skills_test.go b/skills/skills_test.go new file mode 100644 index 0000000..cc9a2dc --- /dev/null +++ b/skills/skills_test.go @@ -0,0 +1,160 @@ +package skills + +import ( + "os" + "path/filepath" + "testing" +) + +func TestLoadSkill(t *testing.T) { + // Create a temp skill + dir := t.TempDir() + skillDir := filepath.Join(dir, "test-skill") + os.MkdirAll(skillDir, 0755) + content := `--- +name: test-skill +description: A test skill for unit testing. Use when testing skill loading. +--- + +# Test Skill + +This is test content. +` + os.WriteFile(filepath.Join(skillDir, "SKILL.md"), []byte(content), 0644) + + skill, err := loadSkill(filepath.Join(skillDir, "SKILL.md")) + if err != nil { + t.Fatalf("loadSkill: %v", err) + } + if skill.Name != "test-skill" { + t.Errorf("name = %q, want %q", skill.Name, "test-skill") + } + if skill.Description != "A test skill for unit testing. Use when testing skill loading." { + t.Errorf("description = %q", skill.Description) + } +} + +func TestParseFrontmatter(t *testing.T) { + tests := []struct { + name string + input string + wantN string + wantD string + wantErr bool + }{ + { + name: "valid", + input: `--- +name: my-skill +description: My description +--- +content`, + wantN: "my-skill", + wantD: "My description", + }, + { + name: "no frontmatter", + input: "# Just content", + wantErr: true, + }, + { + name: "missing name", + input: `--- +description: test +---`, + wantErr: true, + }, + { + name: "missing description", + input: `--- +name: test +---`, + wantErr: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + n, d, err := parseFrontmatter([]byte(tt.input)) + if tt.wantErr { + if err == nil { + t.Error("expected error") + } + return + } + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if n != tt.wantN { + t.Errorf("name = %q, want %q", n, tt.wantN) + } + if d != tt.wantD { + t.Errorf("description = %q, want %q", d, tt.wantD) + } + }) + } +} + +func TestIndexMatch(t *testing.T) { + // Find model directory + home, _ := os.UserHomeDir() + modelDir := "" + candidates := []string{ + "../data/models", + filepath.Join(home, ".local/share/ollie/models"), + } + for _, c := range candidates { + if _, err := os.Stat(filepath.Join(c, "model.onnx")); err == nil { + modelDir = c + break + } + } + if modelDir == "" { + t.Skip("model not found") + } + + // Create temp skill directory + dir := t.TempDir() + + skills := []struct { + name string + desc string + }{ + {"jira-cli", "Interact with Jira for issue tracking. Use for tickets, sprints, and project management."}, + {"github-cli", "Interact with GitHub repositories, issues, and pull requests."}, + {"bash-exec", "Execute bash commands locally or remotely."}, + } + + for _, s := range skills { + skillDir := filepath.Join(dir, s.name) + os.MkdirAll(skillDir, 0755) + content := "---\nname: " + s.name + "\ndescription: " + s.desc + "\n---\n# " + s.name + os.WriteFile(filepath.Join(skillDir, "SKILL.md"), []byte(content), 0644) + } + + idx, err := NewIndex(modelDir, []string{dir}) + if err != nil { + t.Fatalf("NewIndex: %v", err) + } + defer idx.Close() + + // Test matching + results, err := idx.Match("create a jira ticket for the bug", 0.1, 3) + if err != nil { + t.Fatalf("Match: %v", err) + } + + if len(results) == 0 { + t.Fatal("expected results") + } + + t.Logf("Query: 'create a jira ticket for the bug'") + for _, r := range results { + t.Logf(" %s: %.4f", r.Skill.Name, r.Score) + } + + // Jira should be top result + if results[0].Skill.Name != "jira-cli" { + t.Errorf("expected jira-cli as top result, got %s", results[0].Skill.Name) + } +}