746 lines
21 KiB
Go
746 lines
21 KiB
Go
package ollama
|
|
|
|
import (
|
|
"bytes"
|
|
"context"
|
|
"encoding/json"
|
|
"errors"
|
|
"fmt"
|
|
"io"
|
|
"net/http"
|
|
"sort"
|
|
"strings"
|
|
"sync"
|
|
"time"
|
|
|
|
"github.com/local/glpi-neural-brain/internal/workqueue"
|
|
)
|
|
|
|
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
|
|
models map[string]string
|
|
|
|
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 requestPriorityKey struct{}
|
|
|
|
// WithLowPriority marks background work that may use Ollama only after all
|
|
// normal-priority waiters have had a chance to acquire a compatible node. An
|
|
// in-flight model call is never interrupted; prioritization applies at the next
|
|
// pool acquisition boundary.
|
|
func WithLowPriority(ctx context.Context) context.Context {
|
|
return context.WithValue(ctx, requestPriorityKey{}, true)
|
|
}
|
|
|
|
func isLowPriority(ctx context.Context) bool {
|
|
low, _ := ctx.Value(requestPriorityKey{}).(bool)
|
|
return low
|
|
}
|
|
|
|
type Client struct {
|
|
ChatModel, EmbeddingModel string
|
|
HTTP *http.Client
|
|
cfg PoolConfig
|
|
|
|
mu sync.Mutex
|
|
nodes []*nodeState
|
|
roundRobin uint64
|
|
healthReady bool
|
|
healthMu sync.Mutex
|
|
normalWaiters int
|
|
lowWaiters int
|
|
sharedLimiter *workqueue.Limiter
|
|
}
|
|
|
|
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) SetSharedLimiter(limiter *workqueue.Limiter) {
|
|
c.mu.Lock()
|
|
c.sharedLimiter = limiter
|
|
c.mu.Unlock()
|
|
}
|
|
|
|
// SetNodeMaxInflight changes the live per-node reservation ceiling. Existing
|
|
// requests are not cancelled when lowering the limit; new acquisitions wait
|
|
// until inflight falls below the new ceiling.
|
|
func (c *Client) SetNodeMaxInflight(limit int) {
|
|
if limit < 1 {
|
|
limit = 1
|
|
}
|
|
c.mu.Lock()
|
|
c.cfg.NodeMaxInflight = limit
|
|
c.mu.Unlock()
|
|
}
|
|
|
|
func (c *Client) NodeMaxInflight() int {
|
|
c.mu.Lock()
|
|
defer c.mu.Unlock()
|
|
return c.cfg.NodeMaxInflight
|
|
}
|
|
|
|
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 {
|
|
return c.ChatJSONModel(ctx, c.ChatModel, system, user, schema, target)
|
|
}
|
|
|
|
// ChatJSONModel executes a structured chat request with a specific model while
|
|
// keeping the same Ollama node pool, failover rules and shared work limiter.
|
|
// This allows article synthesis and article review to use different models
|
|
// without creating independent pools that could oversubscribe the same GPU.
|
|
func (c *Client) ChatJSONModel(ctx context.Context, model, system, user string, schema any, target any) error {
|
|
model = strings.TrimSpace(model)
|
|
if model == "" {
|
|
model = c.ChatModel
|
|
}
|
|
body := map[string]any{"model": model, "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.doJSONModel(ctx, "chat", model, "/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
|
|
}
|
|
|
|
// ModelStatus reports whether a dynamically selected chat model is present on
|
|
// healthy Ollama nodes. It is used by the analysis dashboard for dedicated
|
|
// author/reviewer models such as Gemma + Qwen.
|
|
func (c *Client) ModelStatus(model string) map[string]any {
|
|
model = strings.TrimSpace(model)
|
|
c.mu.Lock()
|
|
defer c.mu.Unlock()
|
|
now := time.Now()
|
|
healthyWithModel := 0
|
|
available := 0
|
|
nodes := make([]string, 0, len(c.nodes))
|
|
for _, n := range c.nodes {
|
|
if !n.healthy {
|
|
continue
|
|
}
|
|
if _, ok := findModelDigest(n.models, model); !ok {
|
|
continue
|
|
}
|
|
healthyWithModel++
|
|
nodes = append(nodes, n.cfg.Name)
|
|
if now.After(n.cooldownUntil) && n.inflight < c.cfg.NodeMaxInflight {
|
|
available++
|
|
}
|
|
}
|
|
sort.Strings(nodes)
|
|
return map[string]any{"model": model, "healthy_nodes_with_model": healthyWithModel, "available_nodes_with_model": available, "nodes": nodes}
|
|
}
|
|
|
|
func (c *Client) PoolStatus() map[string]any {
|
|
statuses := c.NodeStatuses()
|
|
c.mu.Lock()
|
|
normalWaiters, lowWaiters := c.normalWaiters, c.lowWaiters
|
|
nodeMaxInflight := c.cfg.NodeMaxInflight
|
|
c.mu.Unlock()
|
|
healthy, available := 0, 0
|
|
now := time.Now()
|
|
for _, s := range statuses {
|
|
if s.Healthy && s.Compatible {
|
|
healthy++
|
|
if now.After(s.CooldownUntil) && s.Inflight < nodeMaxInflight {
|
|
available++
|
|
}
|
|
}
|
|
}
|
|
status := map[string]any{"routing_mode": c.cfg.RoutingMode, "node_count": len(statuses), "healthy_nodes": healthy, "available_nodes": available, "node_max_inflight": nodeMaxInflight, "failover_enabled": c.cfg.FailoverEnabled, "normal_waiters": normalWaiters, "low_priority_waiters": lowWaiters, "nodes": statuses}
|
|
c.mu.Lock()
|
|
limiter := c.sharedLimiter
|
|
c.mu.Unlock()
|
|
if limiter != nil {
|
|
status["shared_queue"] = limiter.Status()
|
|
}
|
|
return status
|
|
}
|
|
|
|
func (c *Client) doJSON(ctx context.Context, capability, path string, in, out any) error {
|
|
return c.doJSONModel(ctx, capability, "", path, in, out)
|
|
}
|
|
|
|
func (c *Client) doJSONModel(ctx context.Context, capability, requiredModel, path string, in, out any) error {
|
|
if err := c.ensureHealth(ctx); err != nil {
|
|
return err
|
|
}
|
|
c.mu.Lock()
|
|
limiter := c.sharedLimiter
|
|
c.mu.Unlock()
|
|
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++ {
|
|
// Reserve actual Ollama capacity before entering the shared outbound-work
|
|
// limiter. Previously callers occupied a global queue slot while merely
|
|
// waiting for a busy Ollama node. With one GPU, the second model request
|
|
// could therefore consume the second shared slot and block unrelated
|
|
// SearXNG/fetch work (head-of-line blocking).
|
|
n, err := c.acquireNodeModel(ctx, capability, requiredModel, 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
|
|
|
|
queueKind := "ollama." + strings.ToLower(strings.TrimSpace(capability))
|
|
releaseQueue, queueErr := limiterAcquireKind(ctx, limiter, queueKind)
|
|
if queueErr != nil {
|
|
c.releaseNodeReservation(n)
|
|
return fmt.Errorf("ollama shared queue: %w", queueErr)
|
|
}
|
|
started := time.Now()
|
|
err = c.postNode(ctx, n, path, in, out)
|
|
releaseQueue()
|
|
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 limiterAcquireKind(ctx context.Context, limiter *workqueue.Limiter, kind string) (func(), error) {
|
|
if limiter == nil {
|
|
return func() {}, nil
|
|
}
|
|
return limiter.AcquireKind(ctx, kind)
|
|
}
|
|
|
|
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) {
|
|
return c.acquireNodeModel(ctx, capability, "", tried)
|
|
}
|
|
|
|
func (c *Client) acquireNodeModel(ctx context.Context, capability, requiredModel string, tried map[*nodeState]bool) (*nodeState, error) {
|
|
lowPriority := isLowPriority(ctx)
|
|
c.mu.Lock()
|
|
if lowPriority {
|
|
c.lowWaiters++
|
|
} else {
|
|
c.normalWaiters++
|
|
}
|
|
c.mu.Unlock()
|
|
defer func() {
|
|
c.mu.Lock()
|
|
if lowPriority {
|
|
c.lowWaiters--
|
|
} else {
|
|
c.normalWaiters--
|
|
}
|
|
c.mu.Unlock()
|
|
}()
|
|
for {
|
|
c.mu.Lock()
|
|
// Background jobs yield between every model call while an interactive or
|
|
// normal AI-THINK request is waiting for the pool.
|
|
if lowPriority && c.normalWaiters > 0 {
|
|
c.mu.Unlock()
|
|
select {
|
|
case <-ctx.Done():
|
|
return nil, ctx.Err()
|
|
case <-time.After(25 * time.Millisecond):
|
|
}
|
|
continue
|
|
}
|
|
candidates := make([]*nodeState, 0, len(c.nodes))
|
|
now := time.Now()
|
|
viable := 0
|
|
busy := false
|
|
cooling := false
|
|
var nextCooldown time.Time
|
|
for _, n := range c.nodes {
|
|
if tried[n] || !n.healthy {
|
|
continue
|
|
}
|
|
if strings.TrimSpace(requiredModel) == "" {
|
|
if !n.compatible {
|
|
continue
|
|
}
|
|
if capability == "embedding" && !n.embeddingModel {
|
|
continue
|
|
}
|
|
if capability == "chat" && !n.chatModel {
|
|
continue
|
|
}
|
|
} else {
|
|
if _, ok := findModelDigest(n.models, requiredModel); !ok {
|
|
continue
|
|
}
|
|
}
|
|
viable++
|
|
if now.Before(n.cooldownUntil) {
|
|
cooling = true
|
|
if nextCooldown.IsZero() || n.cooldownUntil.Before(nextCooldown) {
|
|
nextCooldown = 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 {
|
|
if strings.TrimSpace(requiredModel) != "" {
|
|
return nil, fmt.Errorf("no healthy Ollama node with model %q available", requiredModel)
|
|
}
|
|
return nil, errors.New("no additional compatible Ollama node available")
|
|
}
|
|
if !busy && !cooling {
|
|
if strings.TrimSpace(requiredModel) != "" {
|
|
return nil, fmt.Errorf("no healthy Ollama node with model %q available", requiredModel)
|
|
}
|
|
return nil, errors.New("no additional compatible Ollama node available")
|
|
}
|
|
|
|
// A healthy node in cooldown is temporarily unavailable, not unhealthy.
|
|
// Queue behind the reservation/cooldown instead of turning one timeout into
|
|
// a cascade of immediate "no healthy node" failures. The caller context is
|
|
// still the hard upper bound, so shutdowns and request deadlines remain
|
|
// responsive.
|
|
wait := 20 * time.Millisecond
|
|
if cooling && !nextCooldown.IsZero() {
|
|
until := time.Until(nextCooldown)
|
|
if until > 0 && until < 250*time.Millisecond {
|
|
wait = until
|
|
} else if until >= 250*time.Millisecond {
|
|
wait = 250 * time.Millisecond
|
|
}
|
|
}
|
|
select {
|
|
case <-ctx.Done():
|
|
return nil, ctx.Err()
|
|
case <-time.After(wait):
|
|
}
|
|
}
|
|
}
|
|
|
|
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) releaseNodeReservation(n *nodeState) {
|
|
if n == nil {
|
|
return
|
|
}
|
|
c.mu.Lock()
|
|
if n.inflight > 0 {
|
|
n.inflight--
|
|
}
|
|
c.mu.Unlock()
|
|
}
|
|
|
|
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.models = nil
|
|
n.healthy = false
|
|
n.chatModel = false
|
|
n.embeddingModel = false
|
|
n.lastError = r.err.Error()
|
|
continue
|
|
}
|
|
n.healthy = true
|
|
n.models = r.models
|
|
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
|
|
}
|