539 lines
15 KiB
Go
539 lines
15 KiB
Go
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
|
|
}
|