diff --git a/core/services/failover/classify.go b/core/services/failover/classify.go new file mode 100644 index 000000000..58868c089 --- /dev/null +++ b/core/services/failover/classify.go @@ -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 +} diff --git a/core/services/failover/classify_test.go b/core/services/failover/classify_test.go new file mode 100644 index 000000000..455a9f993 --- /dev/null +++ b/core/services/failover/classify_test.go @@ -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), +) diff --git a/core/services/failover/failover_suite_test.go b/core/services/failover/failover_suite_test.go new file mode 100644 index 000000000..1a7bce9d5 --- /dev/null +++ b/core/services/failover/failover_suite_test.go @@ -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") +} diff --git a/core/services/failover/types.go b/core/services/failover/types.go new file mode 100644 index 000000000..dc4ac2002 --- /dev/null +++ b/core/services/failover/types.go @@ -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 +} diff --git a/core/services/failover/types_test.go b/core/services/failover/types_test.go new file mode 100644 index 000000000..abbe0bd09 --- /dev/null +++ b/core/services/failover/types_test.go @@ -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()) + }) +})