package ollama import ( "bytes" "context" "encoding/json" "errors" "fmt" "io" "net/http" "sort" "strings" "sync" "time" ) type NodeConfig struct { Name string URL string Weight int } type PoolConfig struct { Nodes []NodeConfig RoutingMode string NodeMaxInflight int HealthInterval time.Duration FailureCooldown time.Duration RequestTimeout time.Duration FailoverEnabled bool FailoverAttempts int RequireSameModelDigest bool RequireEmbeddingModel bool } type NodeStatus struct { Name string `json:"name"` URL string `json:"url"` Weight int `json:"weight"` Healthy bool `json:"healthy"` Compatible bool `json:"compatible"` ChatModel bool `json:"chat_model"` EmbeddingModel bool `json:"embedding_model"` ChatDigest string `json:"chat_digest,omitempty"` EmbeddingDigest string `json:"embedding_digest,omitempty"` Inflight int `json:"inflight"` Requests uint64 `json:"requests"` Failures uint64 `json:"failures"` AverageDurationMS float64 `json:"average_duration_ms"` CooldownUntil time.Time `json:"cooldown_until,omitempty"` LastCheck time.Time `json:"last_check,omitempty"` LastError string `json:"last_error,omitempty"` } type nodeState struct { cfg NodeConfig healthy bool compatible bool chatModel bool embeddingModel bool chatDigest string embeddingDigest string inflight int requests uint64 failures uint64 totalDuration time.Duration averageDuration time.Duration cooldownUntil time.Time lastCheck time.Time lastError string } type Client struct { ChatModel, EmbeddingModel string HTTP *http.Client cfg PoolConfig mu sync.Mutex nodes []*nodeState roundRobin uint64 healthReady bool healthMu sync.Mutex } type requestError struct { err error retryable bool statusCode int } func (e *requestError) Error() string { return e.err.Error() } func (e *requestError) Unwrap() error { return e.err } func New(base, chat, embed string) *Client { return NewPool(PoolConfig{Nodes: []NodeConfig{{Name: "ollama-1", URL: base, Weight: 1}}, RoutingMode: "least_inflight", NodeMaxInflight: 1, HealthInterval: 15 * time.Second, FailureCooldown: 30 * time.Second, RequestTimeout: 8 * time.Minute, FailoverEnabled: true, FailoverAttempts: 1, RequireSameModelDigest: true, RequireEmbeddingModel: true}, chat, embed) } func NewPool(cfg PoolConfig, chat, embed string) *Client { if cfg.NodeMaxInflight < 1 { cfg.NodeMaxInflight = 1 } if cfg.HealthInterval <= 0 { cfg.HealthInterval = 15 * time.Second } if cfg.FailureCooldown <= 0 { cfg.FailureCooldown = 30 * time.Second } if cfg.RequestTimeout <= 0 { cfg.RequestTimeout = 8 * time.Minute } if cfg.RoutingMode == "" { cfg.RoutingMode = "least_inflight" } c := &Client{ChatModel: chat, EmbeddingModel: embed, cfg: cfg, HTTP: &http.Client{Timeout: cfg.RequestTimeout}} for i, raw := range cfg.Nodes { raw.URL = strings.TrimRight(strings.TrimSpace(raw.URL), "/") if raw.Name == "" { raw.Name = fmt.Sprintf("ollama-%d", i+1) } if raw.Weight < 1 { raw.Weight = 1 } c.nodes = append(c.nodes, &nodeState{cfg: raw}) } return c } func (c *Client) Start(ctx context.Context) { go func() { _ = c.refreshHealth(ctx) ticker := time.NewTicker(c.cfg.HealthInterval) defer ticker.Stop() for { select { case <-ctx.Done(): return case <-ticker.C: _ = c.refreshHealth(ctx) } } }() } func (c *Client) Embed(ctx context.Context, texts []string) ([][]float64, error) { body := map[string]any{"model": c.EmbeddingModel, "input": texts, "truncate": true} var out struct { Embeddings [][]float64 `json:"embeddings"` } if err := c.doJSON(ctx, "embedding", "/api/embed", body, &out); err != nil { return nil, err } if len(out.Embeddings) != len(texts) { return nil, fmt.Errorf("ollama returned %d embeddings for %d inputs", len(out.Embeddings), len(texts)) } return out.Embeddings, nil } func (c *Client) ChatJSON(ctx context.Context, system, user string, schema any, target any) error { body := map[string]any{"model": c.ChatModel, "messages": []map[string]string{{"role": "system", "content": system}, {"role": "user", "content": user}}, "stream": false, "think": false, "format": schema, "options": map[string]any{"temperature": 0}} var env struct { Message struct { Content string `json:"content"` } `json:"message"` } if err := c.doJSON(ctx, "chat", "/api/chat", body, &env); err != nil { return err } if strings.TrimSpace(env.Message.Content) == "" { return fmt.Errorf("empty Ollama response") } if err := json.Unmarshal([]byte(env.Message.Content), target); err != nil { return fmt.Errorf("decode structured response: %w", err) } return nil } func (c *Client) Ping(ctx context.Context) error { if err := c.ensureHealth(ctx); err != nil { return err } c.mu.Lock() defer c.mu.Unlock() for _, n := range c.nodes { if n.healthy && n.compatible && n.chatModel && (!c.cfg.RequireEmbeddingModel || n.embeddingModel) { return nil } } return errors.New("no healthy compatible Ollama node") } func (c *Client) RoutingMode() string { return c.cfg.RoutingMode } func (c *Client) NodeStatuses() []NodeStatus { c.mu.Lock() defer c.mu.Unlock() out := make([]NodeStatus, 0, len(c.nodes)) for _, n := range c.nodes { avg := float64(n.averageDuration.Milliseconds()) out = append(out, NodeStatus{Name: n.cfg.Name, URL: n.cfg.URL, Weight: n.cfg.Weight, Healthy: n.healthy, Compatible: n.compatible, ChatModel: n.chatModel, EmbeddingModel: n.embeddingModel, ChatDigest: n.chatDigest, EmbeddingDigest: n.embeddingDigest, Inflight: n.inflight, Requests: n.requests, Failures: n.failures, AverageDurationMS: avg, CooldownUntil: n.cooldownUntil, LastCheck: n.lastCheck, LastError: n.lastError}) } sort.Slice(out, func(i, j int) bool { return out[i].Name < out[j].Name }) return out } func (c *Client) PoolStatus() map[string]any { statuses := c.NodeStatuses() healthy, available := 0, 0 now := time.Now() for _, s := range statuses { if s.Healthy && s.Compatible { healthy++ if now.After(s.CooldownUntil) && s.Inflight < c.cfg.NodeMaxInflight { available++ } } } return map[string]any{"routing_mode": c.cfg.RoutingMode, "node_count": len(statuses), "healthy_nodes": healthy, "available_nodes": available, "node_max_inflight": c.cfg.NodeMaxInflight, "failover_enabled": c.cfg.FailoverEnabled, "nodes": statuses} } func (c *Client) doJSON(ctx context.Context, capability, path string, in, out any) error { if err := c.ensureHealth(ctx); err != nil { return err } attemptLimit := c.cfg.FailoverAttempts if attemptLimit <= 0 || attemptLimit > len(c.nodes) { attemptLimit = len(c.nodes) } if !c.cfg.FailoverEnabled && attemptLimit > 1 { attemptLimit = 1 } tried := make(map[*nodeState]bool, attemptLimit) var errs []string for attempt := 0; attempt < attemptLimit; attempt++ { n, err := c.acquireNode(ctx, capability, tried) if err != nil { if len(errs) > 0 { return fmt.Errorf("ollama pool failed: %s; %w", strings.Join(errs, "; "), err) } return err } tried[n] = true started := time.Now() err = c.postNode(ctx, n, path, in, out) c.releaseNode(n, time.Since(started), err) if err == nil { return nil } errs = append(errs, n.cfg.Name+": "+err.Error()) var re *requestError if !errors.As(err, &re) || !re.retryable || !c.cfg.FailoverEnabled { break } } return fmt.Errorf("ollama pool request failed: %s", strings.Join(errs, "; ")) } func (c *Client) ensureHealth(ctx context.Context) error { c.mu.Lock() ready := c.healthReady c.mu.Unlock() if ready { return nil } return c.refreshHealth(ctx) } func (c *Client) acquireNode(ctx context.Context, capability string, tried map[*nodeState]bool) (*nodeState, error) { for { c.mu.Lock() candidates := make([]*nodeState, 0, len(c.nodes)) now := time.Now() viable := 0 busy := false for _, n := range c.nodes { if tried[n] || !n.healthy || !n.compatible { continue } if capability == "embedding" && !n.embeddingModel { continue } if capability == "chat" && !n.chatModel { continue } viable++ if now.Before(n.cooldownUntil) { continue } if n.inflight >= c.cfg.NodeMaxInflight { busy = true continue } candidates = append(candidates, n) } if len(candidates) > 0 { n := c.chooseNode(candidates) n.inflight++ c.mu.Unlock() return n, nil } c.mu.Unlock() if viable == 0 || !busy { return nil, errors.New("no additional compatible Ollama node available") } select { case <-ctx.Done(): return nil, ctx.Err() case <-time.After(20 * time.Millisecond): } } } func (c *Client) chooseNode(nodes []*nodeState) *nodeState { switch c.cfg.RoutingMode { case "round_robin": sort.Slice(nodes, func(i, j int) bool { return nodes[i].cfg.Name < nodes[j].cfg.Name }) n := nodes[c.roundRobin%uint64(len(nodes))] c.roundRobin++ return n case "weighted": sort.Slice(nodes, func(i, j int) bool { a := float64(nodes[i].requests+uint64(nodes[i].inflight)) / float64(nodes[i].cfg.Weight) b := float64(nodes[j].requests+uint64(nodes[j].inflight)) / float64(nodes[j].cfg.Weight) if a == b { return nodes[i].cfg.Name < nodes[j].cfg.Name } return a < b }) return nodes[0] case "fastest_recent": sort.Slice(nodes, func(i, j int) bool { if nodes[i].requests == 0 || nodes[j].requests == 0 { if nodes[i].requests == nodes[j].requests { return nodes[i].cfg.Name < nodes[j].cfg.Name } return nodes[i].requests == 0 } a := nodes[i].averageDuration b := nodes[j].averageDuration if a == b { return nodes[i].cfg.Name < nodes[j].cfg.Name } return a < b }) return nodes[0] default: // least_inflight sort.Slice(nodes, func(i, j int) bool { if nodes[i].inflight != nodes[j].inflight { return nodes[i].inflight < nodes[j].inflight } if nodes[i].requests != nodes[j].requests { return nodes[i].requests < nodes[j].requests } return nodes[i].cfg.Name < nodes[j].cfg.Name }) return nodes[0] } } func (c *Client) releaseNode(n *nodeState, duration time.Duration, err error) { c.mu.Lock() defer c.mu.Unlock() if n.inflight > 0 { n.inflight-- } n.requests++ n.totalDuration += duration if n.averageDuration == 0 { n.averageDuration = duration } else { n.averageDuration = time.Duration(float64(n.averageDuration)*0.8 + float64(duration)*0.2) } if err == nil { n.lastError = "" return } n.failures++ n.lastError = err.Error() var re *requestError if errors.As(err, &re) && re.retryable { n.cooldownUntil = time.Now().Add(c.cfg.FailureCooldown) } } func (c *Client) postNode(ctx context.Context, n *nodeState, path string, in, out any) error { b, err := json.Marshal(in) if err != nil { return err } req, err := http.NewRequestWithContext(ctx, http.MethodPost, n.cfg.URL+path, bytes.NewReader(b)) if err != nil { return err } req.Header.Set("Content-Type", "application/json") resp, err := c.HTTP.Do(req) if err != nil { return &requestError{err: err, retryable: true} } defer resp.Body.Close() data, err := io.ReadAll(io.LimitReader(resp.Body, 16<<20)) if err != nil { return &requestError{err: err, retryable: true} } if resp.StatusCode/100 != 2 { retryable := resp.StatusCode == http.StatusRequestTimeout || resp.StatusCode == http.StatusTooManyRequests || resp.StatusCode >= 500 return &requestError{err: fmt.Errorf("ollama %s HTTP %d: %s", path, resp.StatusCode, strings.TrimSpace(string(data))), retryable: retryable, statusCode: resp.StatusCode} } if err := json.Unmarshal(data, out); err != nil { return &requestError{err: fmt.Errorf("decode Ollama response: %w", err), retryable: true} } return nil } func (c *Client) refreshHealth(ctx context.Context) error { c.healthMu.Lock() defer c.healthMu.Unlock() type result struct { n *nodeState models map[string]string err error } results := make(chan result, len(c.nodes)) for _, n := range c.nodes { go func(n *nodeState) { models, err := c.fetchTags(ctx, n.cfg.URL) results <- result{n: n, models: models, err: err} }(n) } collected := make([]result, 0, len(c.nodes)) for range c.nodes { collected = append(collected, <-results) } c.mu.Lock() defer c.mu.Unlock() chatDigests := map[string]struct{}{} embedDigests := map[string]struct{}{} now := time.Now().UTC() for _, r := range collected { n := r.n n.lastCheck = now n.compatible = false if r.err != nil { n.healthy = false n.chatModel = false n.embeddingModel = false n.lastError = r.err.Error() continue } n.healthy = true n.lastError = "" n.chatDigest, n.chatModel = findModelDigest(r.models, c.ChatModel) n.embeddingDigest, n.embeddingModel = findModelDigest(r.models, c.EmbeddingModel) if n.chatModel && n.chatDigest != "" { chatDigests[n.chatDigest] = struct{}{} } if n.embeddingModel && n.embeddingDigest != "" { embedDigests[n.embeddingDigest] = struct{}{} } } divergentChat := c.cfg.RequireSameModelDigest && len(chatDigests) > 1 divergentEmbed := c.cfg.RequireSameModelDigest && len(embedDigests) > 1 for _, n := range c.nodes { if !n.healthy || !n.chatModel || (c.cfg.RequireEmbeddingModel && !n.embeddingModel) { continue } if c.cfg.RequireSameModelDigest && (n.chatDigest == "" || (c.cfg.RequireEmbeddingModel && n.embeddingDigest == "")) { n.lastError = "model digest missing while strict digest validation is enabled" continue } if divergentChat || divergentEmbed { n.lastError = "model digest mismatch inside Ollama pool" continue } n.compatible = true } c.healthReady = true for _, n := range c.nodes { if n.healthy && n.compatible { return nil } } return errors.New("no healthy compatible Ollama node") } func (c *Client) fetchTags(ctx context.Context, base string) (map[string]string, error) { reqCtx, cancel := context.WithTimeout(ctx, minDuration(c.cfg.RequestTimeout, 20*time.Second)) defer cancel() req, err := http.NewRequestWithContext(reqCtx, http.MethodGet, base+"/api/tags", nil) if err != nil { return nil, err } resp, err := c.HTTP.Do(req) if err != nil { return nil, err } defer resp.Body.Close() if resp.StatusCode/100 != 2 { return nil, fmt.Errorf("ollama tags HTTP %d", resp.StatusCode) } var env struct { Models []struct { Name string `json:"name"` Model string `json:"model"` Digest string `json:"digest"` } `json:"models"` } if err := json.NewDecoder(io.LimitReader(resp.Body, 4<<20)).Decode(&env); err != nil { return nil, err } out := make(map[string]string) for _, m := range env.Models { name := m.Name if name == "" { name = m.Model } out[name] = m.Digest } return out, nil } func findModelDigest(models map[string]string, wanted string) (string, bool) { wanted = strings.TrimSpace(wanted) for name, digest := range models { if name == wanted || strings.TrimSuffix(name, ":latest") == strings.TrimSuffix(wanted, ":latest") { return digest, true } } return "", false } func minDuration(a, b time.Duration) time.Duration { if a < b { return a } return b }