Files
glpi-neural-brain/internal/ollama/client.go
T
2026-08-04 05:27:55 +02:00

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
}