mirror of
https://github.com/mudler/LocalAI.git
synced 2026-09-29 17:44:30 -04:00
feat(failover): add types and retryable error classification
Assisted-by: Claude:claude-opus-5-5
This commit is contained in:
1 parent
a6bbc9e01e
commit
e13e7ea8fa
5 files changed
+256
No files matched your search
@@ -0,0 +1,75 @@
|
||||
package failover
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"net/http"
|
||||
"regexp"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"github.com/labstack/echo/v4"
|
||||
"google.golang.org/grpc/codes"
|
||||
grpcstatus "google.golang.org/grpc/status"
|
||||
)
|
||||
|
||||
// cloud-proxy translate mode reports upstream failures as plain text, so the
|
||||
// status is only visible in the message.
|
||||
var upstreamStatusRe = regexp.MustCompile(`upstream (\d{3})`)
|
||||
|
||||
// Errors that the next target would reject in the same way.
|
||||
var requestErrorMarkers = []string{
|
||||
"exceeds the available context size",
|
||||
"is larger than the max context size",
|
||||
"maximum context length",
|
||||
}
|
||||
|
||||
// IsRetryable reports whether a failed attempt should move to the next
|
||||
// target. status is the HTTP status a handler wrote, or 0 when it returned err
|
||||
// without writing.
|
||||
func IsRetryable(err error, status int) bool {
|
||||
if errors.Is(err, context.Canceled) {
|
||||
return false
|
||||
}
|
||||
if status != 0 {
|
||||
return retryableStatus(status)
|
||||
}
|
||||
if err == nil {
|
||||
return false
|
||||
}
|
||||
var he *echo.HTTPError
|
||||
if errors.As(err, &he) {
|
||||
return retryableStatus(he.Code)
|
||||
}
|
||||
if errors.Is(err, context.DeadlineExceeded) {
|
||||
return true
|
||||
}
|
||||
if st, ok := grpcstatus.FromError(err); ok {
|
||||
switch st.Code() {
|
||||
case codes.Unavailable, codes.Internal, codes.DeadlineExceeded, codes.Unknown:
|
||||
return !isRequestError(st.Message())
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
msg := err.Error()
|
||||
if m := upstreamStatusRe.FindStringSubmatch(msg); m != nil {
|
||||
code, _ := strconv.Atoi(m[1])
|
||||
return retryableStatus(code)
|
||||
}
|
||||
// Anything else is usually a dial or load failure of this target.
|
||||
return !isRequestError(msg)
|
||||
}
|
||||
|
||||
func retryableStatus(code int) bool {
|
||||
return code >= 500 && code != http.StatusNotImplemented
|
||||
}
|
||||
|
||||
func isRequestError(msg string) bool {
|
||||
for _, m := range requestErrorMarkers {
|
||||
if strings.Contains(msg, m) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
@@ -0,0 +1,39 @@
|
||||
package failover
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/http"
|
||||
|
||||
"github.com/labstack/echo/v4"
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
"google.golang.org/grpc/codes"
|
||||
grpcstatus "google.golang.org/grpc/status"
|
||||
)
|
||||
|
||||
var _ = DescribeTable("IsRetryable",
|
||||
func(err error, status int, want bool) {
|
||||
Expect(IsRetryable(err, status)).To(Equal(want))
|
||||
},
|
||||
Entry("nil error, no status", nil, 0, false),
|
||||
Entry("held 503", nil, http.StatusServiceUnavailable, true),
|
||||
Entry("held 500", nil, http.StatusInternalServerError, true),
|
||||
Entry("held 501", nil, http.StatusNotImplemented, false),
|
||||
Entry("client cancel", context.Canceled, 0, false),
|
||||
Entry("wrapped client cancel", fmt.Errorf("predict: %w", context.Canceled), 0, false),
|
||||
Entry("deadline", context.DeadlineExceeded, 0, true),
|
||||
Entry("echo 502", echo.NewHTTPError(http.StatusBadGateway, "x"), 0, true),
|
||||
Entry("echo 400", echo.NewHTTPError(http.StatusBadRequest, "x"), 0, false),
|
||||
Entry("echo 404", echo.NewHTTPError(http.StatusNotFound, "x"), 0, false),
|
||||
Entry("grpc unavailable", grpcstatus.Error(codes.Unavailable, "x"), 0, true),
|
||||
Entry("grpc internal", grpcstatus.Error(codes.Internal, "x"), 0, true),
|
||||
Entry("grpc deadline", grpcstatus.Error(codes.DeadlineExceeded, "x"), 0, true),
|
||||
Entry("grpc unknown", grpcstatus.Error(codes.Unknown, "x"), 0, true),
|
||||
Entry("grpc invalid argument", grpcstatus.Error(codes.InvalidArgument, "x"), 0, false),
|
||||
Entry("cloud-proxy upstream 503", errors.New("cloud-proxy: upstream 503: no healthy nodes"), 0, true),
|
||||
Entry("cloud-proxy upstream 429 stays 4xx", errors.New("cloud-proxy: upstream 429: slow down"), 0, false),
|
||||
Entry("context overflow", errors.New("the request exceeds the available context size"), 0, false),
|
||||
Entry("dial error", errors.New("dial tcp 10.0.0.1:8080: connect: connection refused"), 0, true),
|
||||
)
|
||||
@@ -0,0 +1,13 @@
|
||||
package failover
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
)
|
||||
|
||||
func TestFailover(t *testing.T) {
|
||||
RegisterFailHandler(Fail)
|
||||
RunSpecs(t, "Failover test suite")
|
||||
}
|
||||
@@ -0,0 +1,107 @@
|
||||
// Package failover serves a model name from an ordered chain of target
|
||||
// models, moving to the next target when one fails and back when it
|
||||
// recovers.
|
||||
package failover
|
||||
|
||||
import (
|
||||
"slices"
|
||||
"time"
|
||||
|
||||
"github.com/mudler/LocalAI/core/config"
|
||||
)
|
||||
|
||||
type TargetState string
|
||||
|
||||
const (
|
||||
StateHealthy TargetState = "healthy"
|
||||
StateDown TargetState = "down"
|
||||
StateRecovering TargetState = "recovering"
|
||||
StateMissing TargetState = "missing"
|
||||
)
|
||||
|
||||
type ChainState string
|
||||
|
||||
const (
|
||||
ChainPrimary ChainState = "primary"
|
||||
ChainFallback ChainState = "fallback"
|
||||
ChainDegraded ChainState = "degraded"
|
||||
)
|
||||
|
||||
type Kind string
|
||||
|
||||
const (
|
||||
KindLocal Kind = "local"
|
||||
KindRemote Kind = "remote"
|
||||
)
|
||||
|
||||
type Reason string
|
||||
|
||||
const (
|
||||
ReasonTrip Reason = "trip"
|
||||
ReasonRecovery Reason = "recovery"
|
||||
ReasonManual Reason = "manual"
|
||||
ReasonDegraded Reason = "degraded"
|
||||
ReasonMissing Reason = "missing"
|
||||
ReasonInitial Reason = "initial"
|
||||
)
|
||||
|
||||
type EventType string
|
||||
|
||||
const (
|
||||
EventChainSwitched EventType = "chain.switched"
|
||||
EventTargetState EventType = "target.state"
|
||||
)
|
||||
|
||||
// Event is one change of a target state or of a chain's active target.
|
||||
type Event struct {
|
||||
Type EventType `json:"type"`
|
||||
Chain string `json:"chain,omitempty"`
|
||||
Target string `json:"target,omitempty"`
|
||||
From string `json:"from"`
|
||||
To string `json:"to"`
|
||||
State string `json:"state,omitempty"`
|
||||
Reason Reason `json:"reason"`
|
||||
Error string `json:"error,omitempty"`
|
||||
At time.Time `json:"at"`
|
||||
}
|
||||
|
||||
type TargetStatus struct {
|
||||
Model string `json:"model"`
|
||||
Kind Kind `json:"kind"`
|
||||
Warm bool `json:"warm"`
|
||||
State TargetState `json:"state"`
|
||||
ConsecutiveOK int `json:"consecutive_ok"`
|
||||
LastProbe *time.Time `json:"last_probe,omitempty"`
|
||||
LastError string `json:"last_error,omitempty"`
|
||||
}
|
||||
|
||||
type ChainStatus struct {
|
||||
Name string `json:"name"`
|
||||
State ChainState `json:"state"`
|
||||
Active string `json:"active"`
|
||||
ActiveSince time.Time `json:"active_since"`
|
||||
Pinned *string `json:"pinned"`
|
||||
Targets []TargetStatus `json:"targets"`
|
||||
}
|
||||
|
||||
// KindOf decides how a target is probed: proxy backends forward to another
|
||||
// server and are checked over HTTP, everything else runs in this instance.
|
||||
func KindOf(cfg config.ModelConfig) Kind {
|
||||
switch cfg.Backend {
|
||||
case "cloud-proxy", "localai-proxy":
|
||||
return KindRemote
|
||||
}
|
||||
return KindLocal
|
||||
}
|
||||
|
||||
// MergePinned adds warm failover targets to the config-pinned model list, so
|
||||
// the watchdog never evicts them.
|
||||
func MergePinned(pinned, warm []string) []string {
|
||||
out := slices.Clone(pinned)
|
||||
for _, w := range warm {
|
||||
if !slices.Contains(out, w) {
|
||||
out = append(out, w)
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
@@ -0,0 +1,22 @@
|
||||
package failover
|
||||
|
||||
import (
|
||||
"github.com/mudler/LocalAI/core/config"
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
)
|
||||
|
||||
var _ = Describe("KindOf", func() {
|
||||
It("treats proxy backends as remote", func() {
|
||||
Expect(KindOf(config.ModelConfig{Backend: "cloud-proxy"})).To(Equal(KindRemote))
|
||||
Expect(KindOf(config.ModelConfig{Backend: "localai-proxy"})).To(Equal(KindRemote))
|
||||
Expect(KindOf(config.ModelConfig{Backend: "llama-cpp"})).To(Equal(KindLocal))
|
||||
})
|
||||
})
|
||||
|
||||
var _ = Describe("MergePinned", func() {
|
||||
It("adds warm targets without duplicates", func() {
|
||||
Expect(MergePinned([]string{"a", "b"}, []string{"b", "c"})).To(Equal([]string{"a", "b", "c"}))
|
||||
Expect(MergePinned(nil, nil)).To(BeEmpty())
|
||||
})
|
||||
})
|
||||
Reference in new issue
Block a user