diff --git a/core/config/meta/registry.go b/core/config/meta/registry.go index 9e3711af2..fbf40fa8b 100644 --- a/core/config/meta/registry.go +++ b/core/config/meta/registry.go @@ -185,6 +185,64 @@ func DefaultRegistry() map[string]FieldMetaOverride { Advanced: true, Order: 22, }, + "compression.enabled": { + Section: "llm", + Label: "Context Compression", + Description: "Enable compression of chat history before it reaches the model context limit", + Component: "checkbox", + Advanced: true, + Order: 24, + }, + "compression.trigger_at_ratio": { + Section: "llm", + Label: "Compression Trigger Ratio", + Description: "Fraction of the model context window that starts compression", + Component: "slider", + Min: f64(0), + Max: f64(1), + Step: f64(0.05), + Advanced: true, + Order: 25, + }, + "compression.keep_tail_tokens": { + Section: "llm", + Label: "Compression Tail Tokens", + Description: "Number of newest conversation tokens to keep outside the summary", + Component: "number", + Min: f64(0), + Advanced: true, + Order: 26, + }, + "compression.max_summary_tokens": { + Section: "llm", + Label: "Maximum Summary Tokens", + Description: "Maximum number of tokens produced by context compression", + Component: "number", + Min: f64(0), + Advanced: true, + Order: 27, + }, + "compression.compressor_model": { + Section: "llm", + Label: "Compressor Model", + Description: "Chat model used to summarize context; empty uses this model", + Component: "model-select", + AutocompleteProvider: ProviderModelsChat, + Advanced: true, + Order: 28, + }, + "compression.on_post_compression_overflow": { + Section: "llm", + Label: "Post-compression Overflow", + Description: "Action to take when compressed context still exceeds the context limit", + Component: "select", + Options: []FieldOption{ + {Value: "drop_oldest_summary", Label: "Drop Oldest Summary"}, + {Value: "error", Label: "Return Error"}, + }, + Advanced: true, + Order: 29, + }, "cache_type_k": { Section: "llm", Label: "KV Cache Type (K)", diff --git a/core/config/model_config.go b/core/config/model_config.go index 76509598a..1cc7bc903 100644 --- a/core/config/model_config.go +++ b/core/config/model_config.go @@ -84,6 +84,7 @@ type ModelConfig struct { FunctionsConfig functions.FunctionsConfig `yaml:"function,omitempty" json:"function,omitempty"` ReasoningConfig reasoning.Config `yaml:"reasoning,omitempty" json:"reasoning,omitempty"` + Compression CompressionConfig `yaml:"compression,omitempty" json:"compression,omitempty"` // ReasoningEffort is the default reasoning effort (none|minimal|low|medium|high) // for this model. A per-request reasoning_effort overrides it. It is forwarded @@ -152,6 +153,18 @@ type ModelConfig struct { Limits LimitsConfig `yaml:"limits,omitempty" json:"limits,omitempty"` } +// CompressionConfig controls opt-in compression of chat history before inference. +// The request middleware consumes this configuration; keeping it on ModelConfig +// lets operators select a policy per context window and model workload. +type CompressionConfig struct { + Enabled bool `yaml:"enabled,omitempty" json:"enabled,omitempty"` + TriggerAtRatio float64 `yaml:"trigger_at_ratio,omitempty" json:"trigger_at_ratio,omitempty"` + KeepTailTokens int `yaml:"keep_tail_tokens,omitempty" json:"keep_tail_tokens,omitempty"` + MaxSummaryTokens int `yaml:"max_summary_tokens,omitempty" json:"max_summary_tokens,omitempty"` + CompressorModel string `yaml:"compressor_model,omitempty" json:"compressor_model,omitempty"` + OnPostCompressionOverflow string `yaml:"on_post_compression_overflow,omitempty" json:"on_post_compression_overflow,omitempty"` +} + // @Description Admission-control limits applied per request. The // admission middleware enforces these before invoking the handler; // requests that exceed a limit get 503 with a Retry-After hint so @@ -1538,6 +1551,22 @@ func (cfg *ModelConfig) SetDefaults(opts ...ConfigLoaderOption) { } func (c *ModelConfig) Validate() (bool, error) { + if c.Compression.Enabled { + if c.IsCloudProxyBackendPassthrough() { + return false, fmt.Errorf("compression: cloud-proxy passthrough is unsupported; configure proxy mode translate") + } + if c.Compression.TriggerAtRatio < 0 || c.Compression.TriggerAtRatio > 1 { + return false, fmt.Errorf("compression: trigger_at_ratio must be between 0 and 1") + } + if c.Compression.KeepTailTokens < 0 || c.Compression.MaxSummaryTokens < 0 { + return false, fmt.Errorf("compression: token limits cannot be negative") + } + switch c.Compression.OnPostCompressionOverflow { + case "", "error", "drop_oldest_summary": + default: + return false, fmt.Errorf("compression: unknown on_post_compression_overflow %q", c.Compression.OnPostCompressionOverflow) + } + } if c.IsAlias() && len(c.Artifacts) > 0 { return false, fmt.Errorf("alias model %q cannot declare artifacts", c.Name) } diff --git a/core/config/model_config_test.go b/core/config/model_config_test.go index 21741a061..0abc0c7b6 100644 --- a/core/config/model_config_test.go +++ b/core/config/model_config_test.go @@ -52,6 +52,54 @@ parameters: Expect(valid).To(BeTrue()) }) + It("round-trips context compression settings", func() { + raw := []byte(` +name: compressed-chat +backend: llama-cpp +parameters: + model: chat.gguf +compression: + enabled: true + trigger_at_ratio: 0.75 + keep_tail_tokens: 8000 + max_summary_tokens: 2048 + compressor_model: fast-summarizer + on_post_compression_overflow: error +`) + var cfg ModelConfig + Expect(yaml.Unmarshal(raw, &cfg)).To(Succeed()) + Expect(cfg.Compression.Enabled).To(BeTrue()) + Expect(cfg.Compression.TriggerAtRatio).To(Equal(0.75)) + Expect(cfg.Compression.KeepTailTokens).To(Equal(8000)) + Expect(cfg.Compression.MaxSummaryTokens).To(Equal(2048)) + Expect(cfg.Compression.CompressorModel).To(Equal("fast-summarizer")) + Expect(cfg.Compression.OnPostCompressionOverflow).To(Equal("error")) + }) + + It("rejects invalid context compression policies", func() { + cfg := ModelConfig{Compression: CompressionConfig{Enabled: true, TriggerAtRatio: 1.1}} + valid, err := cfg.Validate() + Expect(valid).To(BeFalse()) + Expect(err).To(MatchError(ContainSubstring("trigger_at_ratio"))) + + cfg.Compression.TriggerAtRatio = 0.75 + cfg.Compression.OnPostCompressionOverflow = "truncate_anything" + valid, err = cfg.Validate() + Expect(valid).To(BeFalse()) + Expect(err).To(MatchError(ContainSubstring("on_post_compression_overflow"))) + }) + + It("rejects context compression for cloud proxy passthrough", func() { + cfg := ModelConfig{ + Backend: "cloud-proxy", + Proxy: ProxyConfig{Mode: ProxyModePassthrough}, + Compression: CompressionConfig{Enabled: true}, + } + valid, err := cfg.Validate() + Expect(valid).To(BeFalse()) + Expect(err).To(MatchError(ContainSubstring("cloud-proxy passthrough"))) + }) + It("derives a managed snapshot filename without replacing the logical model", func() { const cacheKey = "0123456789abcdef0123456789abcdef0123456789abcdef0123456789abcdef" cfg := ModelConfig{ diff --git a/core/http/endpoints/localai/mcp.go b/core/http/endpoints/localai/mcp.go index 22b9ca183..f3905442d 100644 --- a/core/http/endpoints/localai/mcp.go +++ b/core/http/endpoints/localai/mcp.go @@ -57,7 +57,7 @@ type MCPErrorEvent struct { // @Param request body schema.OpenAIRequest true "query params" // @Success 200 {object} schema.OpenAIResponse "Response" // @Router /v1/mcp/chat/completions [post] -func MCPEndpoint(cl *config.ModelConfigLoader, ml *model.ModelLoader, evaluator *templates.Evaluator, appConfig *config.ApplicationConfig, natsClient mcpTools.MCPNATSClient) echo.HandlerFunc { +func MCPEndpoint(cl *config.ModelConfigLoader, ml *model.ModelLoader, evaluator *templates.Evaluator, appConfig *config.ApplicationConfig, natsClient mcpTools.MCPNATSClient, compressor middleware.ChatCompressor) echo.HandlerFunc { // The legacy /v1/mcp/chat/completions endpoint never opts into the // in-process LocalAI Assistant tool surface — pass nil holder so the // assistant branch in chat.go is unreachable from this code path. @@ -65,7 +65,7 @@ func MCPEndpoint(cl *config.ModelConfigLoader, ml *model.ModelLoader, evaluator // the per-model PII config and is kept for backward compatibility. // The request-side middleware on the main chat route handles // filtering for the standard /v1/chat/completions path. - chatHandler := openai.ChatEndpoint(cl, ml, evaluator, appConfig, natsClient, nil) + chatHandler := openai.ChatEndpoint(cl, ml, evaluator, appConfig, natsClient, nil, compressor) return func(c echo.Context) error { input, ok := c.Get(middleware.CONTEXT_LOCALS_KEY_LOCALAI_REQUEST).(*schema.OpenAIRequest) diff --git a/core/http/endpoints/openai/chat.go b/core/http/endpoints/openai/chat.go index f863631f6..fcec18932 100644 --- a/core/http/endpoints/openai/chat.go +++ b/core/http/endpoints/openai/chat.go @@ -129,7 +129,7 @@ func applyAutoparserOverride( // @Param request body schema.OpenAIRequest true "query params" // @Success 200 {object} schema.OpenAIResponse "Response" // @Router /v1/chat/completions [post] -func ChatEndpoint(cl *config.ModelConfigLoader, ml *model.ModelLoader, evaluator *templates.Evaluator, startupOptions *config.ApplicationConfig, natsClient mcpTools.MCPNATSClient, assistantHolder *mcpTools.LocalAIAssistantHolder) echo.HandlerFunc { +func ChatEndpoint(cl *config.ModelConfigLoader, ml *model.ModelLoader, evaluator *templates.Evaluator, startupOptions *config.ApplicationConfig, natsClient mcpTools.MCPNATSClient, assistantHolder *mcpTools.LocalAIAssistantHolder, compressor middleware.ChatCompressor) echo.HandlerFunc { return func(c echo.Context) error { var textContentToReturn string id := uuid.New().String() @@ -155,6 +155,9 @@ func ChatEndpoint(cl *config.ModelConfigLoader, ml *model.ModelLoader, evaluator // redaction already ran in the middleware; the response is // forwarded unmodified. if config.IsCloudProxyBackendPassthrough() { + if err := middleware.CompressChatRequest(c, compressor); err != nil { + return err + } return forwardCloudProxyOpenAIViaBackend(c, config, input, ml, startupOptions) } @@ -201,7 +204,8 @@ func ChatEndpoint(cl *config.ModelConfigLoader, ml *model.ModelLoader, evaluator // system message they're responsible for keeping the assistant // safe, so we leave it alone. if !hasSystemMessage(input.Messages) { - input.Messages = append([]schema.Message{{Role: "system", StringContent: assistantHolder.SystemPrompt()}}, input.Messages...) + prompt := assistantHolder.SystemPrompt() + input.Messages = append([]schema.Message{{Role: "system", Content: prompt, StringContent: prompt}}, input.Messages...) } xlog.Debug("LocalAI Assistant tools injected", "count", len(mcpFuncs)) @@ -244,6 +248,9 @@ func ChatEndpoint(cl *config.ModelConfigLoader, ml *model.ModelLoader, evaluator xlog.Error("Failed to parse MCP config", "error", mcpErr) } } + if err := middleware.CompressChatRequest(c, compressor); err != nil { + return err + } xlog.Debug("Tool call routing decision", "shouldUseFn", shouldUseFn, @@ -401,9 +408,17 @@ func ChatEndpoint(cl *config.ModelConfigLoader, ml *model.ModelLoader, evaluator for mcpStreamIter := 0; mcpStreamIter <= mcpStreamMaxIterations; mcpStreamIter++ { // Re-template on MCP iterations - if mcpStreamIter > 0 && !config.TemplateConfig.UseTokenizerTemplate { - predInput = evaluator.TemplateMessages(*input, input.Messages, config, funcs, shouldUseFn) - xlog.Debug("MCP stream re-templating", "iteration", mcpStreamIter) + if mcpStreamIter > 0 { + if err := middleware.CompressChatRequest(c, compressor); err != nil { + fmt.Fprintf(c.Response().Writer, "data: {\"error\":{\"message\":%q,\"type\":\"context_compression_error\"}}\n\n", err.Error()) + fmt.Fprintf(c.Response().Writer, "data: [DONE]\n\n") + c.Response().Flush() + return nil + } + if !config.TemplateConfig.UseTokenizerTemplate { + predInput = evaluator.TemplateMessages(*input, input.Messages, config, funcs, shouldUseFn) + xlog.Debug("MCP stream re-templating", "iteration", mcpStreamIter) + } } responses := make(chan schema.OpenAIResponse) @@ -658,6 +673,7 @@ func ChatEndpoint(cl *config.ModelConfigLoader, ml *model.ModelLoader, evaluator // tools-path worker not surfacing this value at all. if input.StreamOptions != nil && input.StreamOptions.IncludeUsage { trailerUsage := streamUsageFromTokenUsage(finalUsage, extraUsage) + trailerUsage.CompressionMeta = middleware.CompressionMetadata(c) trailer := streamUsageTrailerJSON(id, input.Model, created, trailerUsage) _, _ = fmt.Fprintf(c.Response().Writer, "data: %s\n\n", trailer) } @@ -683,9 +699,14 @@ func ChatEndpoint(cl *config.ModelConfigLoader, ml *model.ModelLoader, evaluator for mcpIteration := 0; mcpIteration <= mcpMaxIterations; mcpIteration++ { // Re-template on each MCP iteration since messages may have changed - if mcpIteration > 0 && !config.TemplateConfig.UseTokenizerTemplate { - predInput = evaluator.TemplateMessages(*input, input.Messages, config, funcs, shouldUseFn) - xlog.Debug("MCP re-templating", "iteration", mcpIteration, "prompt_len", len(predInput)) + if mcpIteration > 0 { + if err := middleware.CompressChatRequest(c, compressor); err != nil { + return err + } + if !config.TemplateConfig.UseTokenizerTemplate { + predInput = evaluator.TemplateMessages(*input, input.Messages, config, funcs, shouldUseFn) + xlog.Debug("MCP re-templating", "iteration", mcpIteration, "prompt_len", len(predInput)) + } } // Detect if thinking token is already in prompt or template @@ -1010,6 +1031,7 @@ func ChatEndpoint(cl *config.ModelConfigLoader, ml *model.ModelLoader, evaluator usage.TimingTokenGeneration = tokenUsage.TimingTokenGeneration usage.TimingPromptProcessing = tokenUsage.TimingPromptProcessing } + usage.CompressionMeta = middleware.CompressionMetadata(c) resp := &schema.OpenAIResponse{ ID: id, diff --git a/core/http/middleware/compression.go b/core/http/middleware/compression.go new file mode 100644 index 000000000..e4720cf33 --- /dev/null +++ b/core/http/middleware/compression.go @@ -0,0 +1,102 @@ +package middleware + +import ( + "context" + "net/http" + + "github.com/labstack/echo/v4" + "github.com/mudler/LocalAI/core/config" + "github.com/mudler/LocalAI/core/schema" + compressionservice "github.com/mudler/LocalAI/core/services/compression" + "github.com/mudler/LocalAI/pkg/tokens" +) + +const contextKeyCompressionMetadata = "COMPRESSION_METADATA" + +type ChatCompressor interface { + Transform(context.Context, config.CompressionConfig, int, string, []schema.Message) ([]schema.Message, *compressionservice.Metadata, error) +} + +func ContextCompression(compressor ChatCompressor) echo.MiddlewareFunc { + return func(next echo.HandlerFunc) echo.HandlerFunc { + return func(c echo.Context) error { + if err := CompressChatRequest(c, compressor); err != nil { + return err + } + return next(c) + } + } +} + +func CompressChatRequest(c echo.Context, compressor ChatCompressor) error { + cfg, ok := c.Get(CONTEXT_LOCALS_KEY_MODEL_CONFIG).(*config.ModelConfig) + if !ok || cfg == nil || !cfg.Compression.Enabled { + return nil + } + if cfg.IsCloudProxyBackendPassthrough() { + return echo.NewHTTPError(http.StatusBadRequest, "context compression is not supported by cloud-proxy passthrough models; configure translate mode or a local compressor model") + } + input, ok := c.Get(CONTEXT_LOCALS_KEY_LOCALAI_REQUEST).(*schema.OpenAIRequest) + if !ok { + return echo.NewHTTPError(http.StatusBadRequest, "context compression requires a chat request") + } + contextSize := config.DefaultContextSize + if cfg.ContextSize != nil && *cfg.ContextSize > 0 { + contextSize = *cfg.ContextSize + } + extraPayload := make(map[string]any) + if len(input.Functions) > 0 { + extraPayload["functions"] = input.Functions + } + if len(input.Tools) > 0 { + extraPayload["tools"] = input.Tools + } + if input.FunctionCall != nil { + extraPayload["function_call"] = input.FunctionCall + } + if input.ToolsChoice != nil { + extraPayload["tool_choice"] = input.ToolsChoice + } + if input.ResponseFormat != nil { + extraPayload["response_format"] = input.ResponseFormat + } + requestOverhead, err := tokens.CountPayload(extraPayload) + if err != nil { + return echo.NewHTTPError(http.StatusBadRequest, err.Error()) + } + if cfg.Maxtokens != nil && *cfg.Maxtokens > 0 { + requestOverhead += *cfg.Maxtokens + } + messages, meta, err := compressor.Transform(c.Request().Context(), cfg.Compression, contextSize-requestOverhead, cfg.ModelID(), input.Messages) + if err != nil { + if compressionservice.IsOverflow(err) { + return echo.NewHTTPError(http.StatusRequestEntityTooLarge, err.Error()) + } + return echo.NewHTTPError(http.StatusInternalServerError, err.Error()) + } + input.Messages = messages + if meta != nil { + meta.OriginalTokens += requestOverhead + meta.CompressedTokens += requestOverhead + if previous, ok := c.Get(contextKeyCompressionMetadata).(*compressionservice.Metadata); ok && previous != nil { + meta.OriginalTokens = previous.OriginalTokens + meta.DroppedTurns += previous.DroppedTurns + meta.SummaryTokens += previous.SummaryTokens + meta.OverflowRecoveries += previous.OverflowRecoveries + } + c.Set(contextKeyCompressionMetadata, meta) + } + return nil +} + +func CompressionMetadata(c echo.Context) *schema.CompressionMetadata { + meta, ok := c.Get(contextKeyCompressionMetadata).(*compressionservice.Metadata) + if !ok || meta == nil { + return nil + } + return &schema.CompressionMetadata{ + OriginalTokens: meta.OriginalTokens, CompressedTokens: meta.CompressedTokens, + DroppedTurns: meta.DroppedTurns, Compressor: meta.Compressor, + SummaryTokens: meta.SummaryTokens, OverflowRecoveries: meta.OverflowRecoveries, + } +} diff --git a/core/http/middleware/compression_test.go b/core/http/middleware/compression_test.go new file mode 100644 index 000000000..6b6de582d --- /dev/null +++ b/core/http/middleware/compression_test.go @@ -0,0 +1,85 @@ +package middleware_test + +import ( + "context" + "net/http" + "net/http/httptest" + + "github.com/labstack/echo/v4" + "github.com/mudler/LocalAI/core/config" + httpmiddleware "github.com/mudler/LocalAI/core/http/middleware" + "github.com/mudler/LocalAI/core/schema" + compressionservice "github.com/mudler/LocalAI/core/services/compression" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +type transformerFunc func(context.Context, config.CompressionConfig, int, string, []schema.Message) ([]schema.Message, *compressionservice.Metadata, error) + +func (f transformerFunc) Transform(ctx context.Context, policy config.CompressionConfig, contextSize int, model string, messages []schema.Message) ([]schema.Message, *compressionservice.Metadata, error) { + return f(ctx, policy, contextSize, model, messages) +} + +var _ = Describe("Context compression middleware", func() { + It("transforms requests and stamps metadata", func() { + e := echo.New() + req := httptest.NewRequest(http.MethodPost, "/v1/chat/completions", nil) + rec := httptest.NewRecorder() + c := e.NewContext(req, rec) + ctxSize := 32 + c.Set(httpmiddleware.CONTEXT_LOCALS_KEY_MODEL_CONFIG, &config.ModelConfig{Name: "primary", LLMConfig: config.LLMConfig{ContextSize: &ctxSize}, Compression: config.CompressionConfig{Enabled: true}}) + c.Set(httpmiddleware.CONTEXT_LOCALS_KEY_LOCALAI_REQUEST, &schema.OpenAIRequest{Messages: []schema.Message{{Role: "user", Content: "old"}}}) + wantMeta := &compressionservice.Metadata{OriginalTokens: 40, CompressedTokens: 10} + middleware := httpmiddleware.ContextCompression(transformerFunc(func(_ context.Context, _ config.CompressionConfig, gotSize int, model string, _ []schema.Message) ([]schema.Message, *compressionservice.Metadata, error) { + Expect(gotSize).To(Equal(32)) + Expect(model).To(Equal("primary")) + return []schema.Message{{Role: "system", Content: "summary"}}, wantMeta, nil + })) + handler := middleware(func(c echo.Context) error { + input := c.Get(httpmiddleware.CONTEXT_LOCALS_KEY_LOCALAI_REQUEST).(*schema.OpenAIRequest) + Expect(input.Messages).To(HaveLen(1)) + Expect(input.Messages[0].Role).To(Equal("system")) + Expect(httpmiddleware.CompressionMetadata(c)).NotTo(BeNil()) + Expect(httpmiddleware.CompressionMetadata(c).OriginalTokens).To(Equal(40)) + return c.NoContent(http.StatusNoContent) + }) + + Expect(handler(c)).To(Succeed()) + }) + + It("maps overflow to HTTP 413", func() { + e := echo.New() + c := e.NewContext(httptest.NewRequest(http.MethodPost, "/v1/chat/completions", nil), httptest.NewRecorder()) + ctxSize := 1 + c.Set(httpmiddleware.CONTEXT_LOCALS_KEY_MODEL_CONFIG, &config.ModelConfig{Name: "primary", LLMConfig: config.LLMConfig{ContextSize: &ctxSize}, Compression: config.CompressionConfig{Enabled: true}}) + c.Set(httpmiddleware.CONTEXT_LOCALS_KEY_LOCALAI_REQUEST, &schema.OpenAIRequest{Messages: []schema.Message{{Role: "user"}, {Role: "user"}}}) + realService := compressionservice.New(wordCountAdapter{}, compressionservice.SummarizerFunc(func(context.Context, string, []schema.Message, int) (string, int, error) { + return "still too large", 3, nil + })) + + err := httpmiddleware.ContextCompression(realService)(func(echo.Context) error { return nil })(c) + httpErr, ok := err.(*echo.HTTPError) + Expect(ok).To(BeTrue()) + Expect(httpErr.Code).To(Equal(http.StatusRequestEntityTooLarge)) + }) + + It("resolves the automatic context-size sentinel", func() { + e := echo.New() + c := e.NewContext(httptest.NewRequest(http.MethodPost, "/v1/chat/completions", nil), httptest.NewRecorder()) + autoContext := -1 + c.Set(httpmiddleware.CONTEXT_LOCALS_KEY_MODEL_CONFIG, &config.ModelConfig{Name: "primary", LLMConfig: config.LLMConfig{ContextSize: &autoContext}, Compression: config.CompressionConfig{Enabled: true}}) + c.Set(httpmiddleware.CONTEXT_LOCALS_KEY_LOCALAI_REQUEST, &schema.OpenAIRequest{Messages: []schema.Message{{Role: "user", Content: "hello"}}}) + + err := httpmiddleware.CompressChatRequest(c, transformerFunc(func(_ context.Context, _ config.CompressionConfig, gotSize int, _ string, messages []schema.Message) ([]schema.Message, *compressionservice.Metadata, error) { + Expect(gotSize).To(Equal(config.DefaultContextSize)) + return messages, nil, nil + })) + Expect(err).NotTo(HaveOccurred()) + }) +}) + +type wordCountAdapter struct{} + +func (wordCountAdapter) CountMessages(messages []schema.Message) (int, error) { + return len(messages) * 10, nil +} diff --git a/core/http/routes/localai.go b/core/http/routes/localai.go index 30a2803f7..1da4683db 100644 --- a/core/http/routes/localai.go +++ b/core/http/routes/localai.go @@ -9,12 +9,16 @@ import ( mcpTools "github.com/mudler/LocalAI/core/http/endpoints/mcp" "github.com/mudler/LocalAI/core/http/middleware" "github.com/mudler/LocalAI/core/schema" + compressionservice "github.com/mudler/LocalAI/core/services/compression" "github.com/mudler/LocalAI/core/services/galleryop" "github.com/mudler/LocalAI/core/services/monitoring" "github.com/mudler/LocalAI/core/services/nodes" + "github.com/mudler/LocalAI/core/services/routing/pii" + "github.com/mudler/LocalAI/core/services/routing/piiadapter" "github.com/mudler/LocalAI/core/templates" "github.com/mudler/LocalAI/internal" "github.com/mudler/LocalAI/pkg/model" + "github.com/mudler/LocalAI/pkg/tokens" echoswagger "github.com/swaggo/echo-swagger" ) @@ -447,11 +451,15 @@ func RegisterLocalAIRoutes(router *echo.Echo, // MCP endpoint - supports both streaming and non-streaming modes // Note: streaming mode is NOT compatible with the OpenAI apis. We have a set which streams more states. if evaluator != nil && !appConfig.DisableMCP { + chatCompressor := compressionservice.New( + compressionservice.CounterFunc(tokens.CountMessages), + compressionservice.NewInferenceSummarizer(cl, ml, appConfig), + ) var mcpNATS mcpTools.MCPNATSClient if d := app.Distributed(); d != nil { mcpNATS = d.Nats } - mcpStreamHandler := localai.MCPEndpoint(cl, ml, evaluator, appConfig, mcpNATS) + mcpStreamHandler := localai.MCPEndpoint(cl, ml, evaluator, appConfig, mcpNATS, chatCompressor) mcpStreamMiddleware := []echo.MiddlewareFunc{ requestExtractor.BuildFilteredFirstAvailableDefaultModel(config.BuildUsecaseFilterFn(config.FLAG_CHAT)), requestExtractor.SetModelAndConfig(func() schema.LocalAIRequest { return new(schema.OpenAIRequest) }), @@ -463,6 +471,7 @@ func RegisterLocalAIRoutes(router *echo.Echo, return next(c) } }, + pii.RequestMiddleware(app.PIIRedactor(), app.PIIEvents(), piiadapter.OpenAI(), app.FallbackUser(), pii.WithNERResolver(app.PIINERResolver()), pii.WithPolicyResolver(app.PIIPolicyResolver())), } router.POST("/v1/mcp/chat/completions", mcpStreamHandler, mcpStreamMiddleware...) router.POST("/mcp/v1/chat/completions", mcpStreamHandler, mcpStreamMiddleware...) diff --git a/core/http/routes/openai.go b/core/http/routes/openai.go index 24ee4e1e3..6a9012626 100644 --- a/core/http/routes/openai.go +++ b/core/http/routes/openai.go @@ -9,9 +9,11 @@ import ( "github.com/mudler/LocalAI/core/http/endpoints/openai" "github.com/mudler/LocalAI/core/http/middleware" "github.com/mudler/LocalAI/core/schema" + compressionservice "github.com/mudler/LocalAI/core/services/compression" "github.com/mudler/LocalAI/core/services/routing/pii" "github.com/mudler/LocalAI/core/services/routing/piiadapter" "github.com/mudler/LocalAI/core/services/routing/router" + "github.com/mudler/LocalAI/pkg/tokens" ) func RegisterOpenAIRoutes(app *echo.Echo, @@ -43,7 +45,11 @@ func RegisterOpenAIRoutes(app *echo.Echo, } // chat - chatHandler := openai.ChatEndpoint(application.ModelConfigLoader(), application.ModelLoader(), application.TemplatesEvaluator(), application.ApplicationConfig(), natsClient, application.LocalAIAssistant()) + chatCompressor := compressionservice.New( + compressionservice.CounterFunc(tokens.CountMessages), + compressionservice.NewInferenceSummarizer(application.ModelConfigLoader(), application.ModelLoader(), application.ApplicationConfig()), + ) + chatHandler := openai.ChatEndpoint(application.ModelConfigLoader(), application.ModelLoader(), application.TemplatesEvaluator(), application.ApplicationConfig(), natsClient, application.LocalAIAssistant(), chatCompressor) chatMiddleware := []echo.MiddlewareFunc{ nodeHeaderMiddleware, usageMiddleware, @@ -78,11 +84,11 @@ func RegisterOpenAIRoutes(app *echo.Echo, // saturated downstream gets rejected even when the requested // router-model has slack. middleware.AdmissionControl(application.AdmissionLimiter(), application.PIIEvents()), - // PII redaction runs INNERMOST, after RouteModel has resolved - // the actual served model. This is what makes per-model PII + // PII redaction runs after RouteModel has resolved the actual served + // model and before compression. This makes per-model PII // configs honour the routed target (e.g., a router fans out to - // claude-strict; that model's pii block applies, not the - // router model's). + // claude-strict; that model's pii block applies, not the router + // model's), and prevents the compressor from seeing unredacted input. pii.RequestMiddleware(application.PIIRedactor(), application.PIIEvents(), piiadapter.OpenAI(), application.FallbackUser(), pii.WithNERResolver(application.PIINERResolver()), pii.WithPolicyResolver(application.PIIPolicyResolver())), } app.POST("/v1/chat/completions", chatHandler, chatMiddleware...) diff --git a/core/schema/openai.go b/core/schema/openai.go index 752a27853..04c90548c 100644 --- a/core/schema/openai.go +++ b/core/schema/openai.go @@ -33,8 +33,18 @@ type OpenAIUsage struct { OutputTokens int `json:"output_tokens,omitempty"` InputTokensDetails *InputTokensDetails `json:"input_tokens_details,omitempty"` // Extra timing data, disabled by default as is't not a part of OpenAI specification - TimingPromptProcessing float64 `json:"timing_prompt_processing,omitempty"` - TimingTokenGeneration float64 `json:"timing_token_generation,omitempty"` + TimingPromptProcessing float64 `json:"timing_prompt_processing,omitempty"` + TimingTokenGeneration float64 `json:"timing_token_generation,omitempty"` + CompressionMeta *CompressionMetadata `json:"compression_meta,omitempty"` +} + +type CompressionMetadata struct { + OriginalTokens int `json:"original_tokens"` + CompressedTokens int `json:"compressed_tokens"` + DroppedTurns int `json:"dropped_turns"` + Compressor string `json:"compressor"` + SummaryTokens int `json:"summary_tokens"` + OverflowRecoveries int `json:"overflow_recoveries"` } type Item struct { diff --git a/core/services/compression/compression_suite_test.go b/core/services/compression/compression_suite_test.go new file mode 100644 index 000000000..cd0aa55c8 --- /dev/null +++ b/core/services/compression/compression_suite_test.go @@ -0,0 +1,13 @@ +package compression + +import ( + "testing" + + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +func TestCompression(t *testing.T) { + RegisterFailHandler(Fail) + RunSpecs(t, "Context compression test suite") +} diff --git a/core/services/compression/inference.go b/core/services/compression/inference.go new file mode 100644 index 000000000..af6a587fe --- /dev/null +++ b/core/services/compression/inference.go @@ -0,0 +1,62 @@ +package compression + +import ( + "context" + "encoding/json" + "fmt" + "strings" + + "github.com/mudler/LocalAI/core/backend" + "github.com/mudler/LocalAI/core/config" + "github.com/mudler/LocalAI/core/schema" + "github.com/mudler/LocalAI/pkg/functions" + "github.com/mudler/LocalAI/pkg/model" +) + +const summarizerInstruction = `Summarize the following conversation for an AI agent to continue coherently. Preserve names, numbers, decisions, URLs, error messages, tool names, and tool results. Drop pleasantries and repetition.` + +type InferenceSummarizer struct { + configs *config.ModelConfigLoader + models *model.ModelLoader + app *config.ApplicationConfig +} + +func NewInferenceSummarizer(configs *config.ModelConfigLoader, models *model.ModelLoader, app *config.ApplicationConfig) *InferenceSummarizer { + return &InferenceSummarizer{configs: configs, models: models, app: app} +} + +func (s *InferenceSummarizer) Summarize(ctx context.Context, modelName string, messages []schema.Message, maxTokens int) (string, int, error) { + cfg, err := s.configs.LoadModelConfigFileByNameDefaultOptions(modelName, s.app) + if err != nil { + return "", 0, fmt.Errorf("load compressor model %q: %w", modelName, err) + } + if cfg == nil { + return "", 0, fmt.Errorf("load compressor model %q: configuration not found", modelName) + } + runtimeCfg := *cfg + cfg = &runtimeCfg + cfg.Maxtokens = &maxTokens + temperature := 0.0 + cfg.Temperature = &temperature + payload, err := json.Marshal(messages) + if err != nil { + return "", 0, fmt.Errorf("encode conversation: %w", err) + } + prompt := fmt.Sprintf("%s\n\nMaximum summary length: %d tokens.\n\nConversation:\n%s", summarizerInstruction, maxTokens, payload) + fn, err := backend.ModelInference(ctx, prompt, nil, nil, nil, nil, s.models, cfg, s.configs, s.app, nil, "", "", nil, nil, nil, nil) + if err != nil { + return "", 0, err + } + response, err := fn() + if err != nil { + return "", 0, err + } + summary := strings.TrimSpace(functions.ContentFromChatDeltas(response.ChatDeltas)) + if summary == "" { + summary = strings.TrimSpace(response.Response) + } + if summary == "" { + return "", 0, fmt.Errorf("compressor model %q returned an empty summary", modelName) + } + return summary, response.Usage.Completion, nil +} diff --git a/core/services/compression/metrics.go b/core/services/compression/metrics.go new file mode 100644 index 000000000..2d07b0bf6 --- /dev/null +++ b/core/services/compression/metrics.go @@ -0,0 +1,41 @@ +package compression + +import ( + "context" + "sync" + "time" + + "go.opentelemetry.io/otel" + "go.opentelemetry.io/otel/attribute" + "go.opentelemetry.io/otel/metric" +) + +var ( + metricsOnce sync.Once + events metric.Int64Counter + ratio metric.Float64Histogram + duration metric.Float64Histogram +) + +func record(model, result string, started time.Time, original, compressed int) { + metricsOnce.Do(func() { + meter := otel.Meter("github.com/mudler/LocalAI") + events, _ = meter.Int64Counter("localai_compression_events_total", metric.WithDescription("Chat context compression attempts by model and result")) + ratio, _ = meter.Float64Histogram("localai_compression_ratio", metric.WithDescription("Original-to-compressed chat token ratio by model")) + duration, _ = meter.Float64Histogram("localai_compression_duration_seconds", metric.WithDescription("Chat context compression duration by model")) + }) + attrs := metric.WithAttributes(attribute.String("model", model), attribute.String("result", result)) + if events != nil { + events.Add(context.Background(), 1, attrs) + } + if started.IsZero() { + return + } + modelAttr := metric.WithAttributes(attribute.String("model", model)) + if duration != nil { + duration.Record(context.Background(), time.Since(started).Seconds(), modelAttr) + } + if ratio != nil && compressed > 0 { + ratio.Record(context.Background(), float64(original)/float64(compressed), modelAttr) + } +} diff --git a/core/services/compression/service.go b/core/services/compression/service.go new file mode 100644 index 000000000..c5776902d --- /dev/null +++ b/core/services/compression/service.go @@ -0,0 +1,204 @@ +package compression + +import ( + "context" + "errors" + "fmt" + "strings" + "time" + + "github.com/mudler/LocalAI/core/config" + "github.com/mudler/LocalAI/core/schema" +) + +const summaryPrefix = "[COMPRESSED: " + +type Counter interface { + CountMessages([]schema.Message) (int, error) +} + +type CounterFunc func([]schema.Message) (int, error) + +func (f CounterFunc) CountMessages(messages []schema.Message) (int, error) { return f(messages) } + +type Summarizer interface { + Summarize(context.Context, string, []schema.Message, int) (string, int, error) +} + +type SummarizerFunc func(context.Context, string, []schema.Message, int) (string, int, error) + +func (f SummarizerFunc) Summarize(ctx context.Context, model string, messages []schema.Message, maxTokens int) (string, int, error) { + return f(ctx, model, messages, maxTokens) +} + +type Metadata struct { + OriginalTokens int `json:"original_tokens"` + CompressedTokens int `json:"compressed_tokens"` + DroppedTurns int `json:"dropped_turns"` + Compressor string `json:"compressor"` + SummaryTokens int `json:"summary_tokens"` + OverflowRecoveries int `json:"overflow_recoveries"` +} + +type overflowError struct{ tokens, limit int } + +func (e *overflowError) Error() string { + return fmt.Sprintf("compressed request (%d tokens) exceeds context size (%d)", e.tokens, e.limit) +} + +func IsOverflow(err error) bool { + var target *overflowError + return errors.As(err, &target) +} + +type Service struct { + counter Counter + summarizer Summarizer +} + +func New(counter Counter, summarizer Summarizer) *Service { + return &Service{counter: counter, summarizer: summarizer} +} + +func (s *Service) Transform(ctx context.Context, policy config.CompressionConfig, contextSize int, primaryModel string, messages []schema.Message) ([]schema.Message, *Metadata, error) { + if !policy.Enabled || len(messages) == 0 { + record(primaryModel, "skipped", time.Time{}, 0, 0) + return messages, nil, nil + } + originalTokens, err := s.counter.CountMessages(messages) + if err != nil { + record(primaryModel, "error", time.Now(), 0, 0) + return nil, nil, err + } + if contextSize <= 0 { + record(primaryModel, "error", time.Now(), originalTokens, originalTokens) + return nil, nil, &overflowError{tokens: originalTokens, limit: contextSize} + } + ratio := policy.TriggerAtRatio + if ratio == 0 { + ratio = .75 + } + if originalTokens < int(float64(contextSize)*ratio) { + record(primaryModel, "skipped", time.Time{}, originalTokens, originalTokens) + return messages, nil, nil + } + started := time.Now() + if len(messages) < 2 { + if originalTokens <= contextSize { + record(primaryModel, "skipped", started, originalTokens, originalTokens) + return messages, nil, nil + } + record(primaryModel, "error", started, originalTokens, originalTokens) + return nil, nil, &overflowError{tokens: originalTokens, limit: contextSize} + } + + prefixEnd := 0 + for prefixEnd < len(messages) && isPreservedPrefix(messages[prefixEnd]) { + prefixEnd++ + } + prefix := messages[:prefixEnd] + head, tail := partition(messages[prefixEnd:], policy.KeepTailTokens, s.counter) + if len(head) == 0 { + if originalTokens <= contextSize { + record(primaryModel, "skipped", started, originalTokens, originalTokens) + return messages, nil, nil + } + record(primaryModel, "error", started, originalTokens, originalTokens) + return nil, nil, &overflowError{tokens: originalTokens, limit: contextSize} + } + compressor := policy.CompressorModel + if compressor == "" { + compressor = primaryModel + } + maxSummaryTokens := policy.MaxSummaryTokens + if maxSummaryTokens == 0 { + maxSummaryTokens = 512 + } + summary, summaryTokens, err := s.summarizer.Summarize(ctx, compressor, head, maxSummaryTokens) + if err != nil { + record(primaryModel, "error", started, originalTokens, 0) + return nil, nil, fmt.Errorf("compress chat history: %w", err) + } + content := summaryPrefix + strings.TrimSpace(summary) + "]" + result := append([]schema.Message(nil), prefix...) + result = append(result, schema.Message{Role: "system", Content: content, StringContent: content}) + result = append(result, tail...) + compressedTokens, err := s.counter.CountMessages(result) + if err != nil { + record(primaryModel, "error", started, originalTokens, 0) + return nil, nil, err + } + meta := &Metadata{ + OriginalTokens: originalTokens, CompressedTokens: compressedTokens, + DroppedTurns: len(head), Compressor: compressor, SummaryTokens: summaryTokens, + } + if compressedTokens <= contextSize { + record(primaryModel, "success", started, originalTokens, compressedTokens) + return result, meta, nil + } + if policy.OnPostCompressionOverflow != "drop_oldest_summary" { + record(primaryModel, "error", started, originalTokens, compressedTokens) + return nil, nil, &overflowError{tokens: compressedTokens, limit: contextSize} + } + + for recoveries := 0; recoveries < 2; recoveries++ { + idx := oldestSummary(result) + if idx < 0 { + break + } + result = append(result[:idx], result[idx+1:]...) + meta.OverflowRecoveries++ + compressedTokens, err = s.counter.CountMessages(result) + if err != nil { + return nil, nil, err + } + meta.CompressedTokens = compressedTokens + if compressedTokens <= contextSize { + record(primaryModel, "success", started, originalTokens, compressedTokens) + return result, meta, nil + } + } + record(primaryModel, "error", started, originalTokens, compressedTokens) + return nil, nil, &overflowError{tokens: compressedTokens, limit: contextSize} +} + +func isPreservedPrefix(message schema.Message) bool { + return message.Role == "system" || message.Role == "developer" +} + +func partition(messages []schema.Message, keepTokens int, counter Counter) ([]schema.Message, []schema.Message) { + if keepTokens <= 0 { + keepTokens = 2048 + } + cut := len(messages) + for cut > 1 { + start := cut - 1 + if messages[start].Role == "tool" { + for start > 0 && messages[start-1].Role == "tool" { + start-- + } + if start > 0 && len(messages[start-1].ToolCalls) > 0 { + start-- + } + } + candidate := messages[start:] + tokens, err := counter.CountMessages(candidate) + // The newest complete unit is mandatory even when it exceeds the + // preferred tail budget. Summarizing the active user request changes + // its meaning; the post-compression fit check returns 413 instead. + if cut != len(messages) && (err != nil || tokens > keepTokens) { + break + } + cut = start + } + return messages[:cut], messages[cut:] +} + +func oldestSummary(messages []schema.Message) int { + for i, message := range messages { + if message.Role == "system" && strings.HasPrefix(message.StringContent, summaryPrefix) { + return i + } + } + return -1 +} diff --git a/core/services/compression/service_test.go b/core/services/compression/service_test.go new file mode 100644 index 000000000..0ed9c5f2d --- /dev/null +++ b/core/services/compression/service_test.go @@ -0,0 +1,174 @@ +package compression + +import ( + "context" + "strings" + + "github.com/mudler/LocalAI/core/config" + "github.com/mudler/LocalAI/core/schema" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +type wordCounter struct{} + +func (wordCounter) CountMessages(messages []schema.Message) (int, error) { + total := 0 + for _, message := range messages { + total += len(strings.Fields(message.StringContent)) + } + return total, nil +} + +var _ = Describe("Context compression", func() { + It("skips when disabled or below the threshold", func() { + called := false + service := New(wordCounter{}, SummarizerFunc(func(context.Context, string, []schema.Message, int) (string, int, error) { + called = true + return "summary", 1, nil + })) + messages := []schema.Message{{Role: "user", StringContent: "one two", Content: "one two"}} + + got, meta, err := service.Transform(context.Background(), config.CompressionConfig{}, 10, "primary", messages) + Expect(err).NotTo(HaveOccurred()) + Expect(meta).To(BeNil()) + Expect(called).To(BeFalse()) + Expect(got).To(HaveLen(1)) + got, meta, err = service.Transform(context.Background(), config.CompressionConfig{Enabled: true, TriggerAtRatio: .75}, 10, "primary", messages) + Expect(err).NotTo(HaveOccurred()) + Expect(meta).To(BeNil()) + Expect(called).To(BeFalse()) + Expect(got).To(HaveLen(1)) + }) + + It("compresses the head and preserves a tool chain in the tail", func() { + var compressed []schema.Message + service := New(wordCounter{}, SummarizerFunc(func(_ context.Context, model string, messages []schema.Message, maxTokens int) (string, int, error) { + Expect(model).To(Equal("fast")) + Expect(maxTokens).To(Equal(4)) + compressed = append([]schema.Message(nil), messages...) + return "facts retained", 2, nil + })) + messages := []schema.Message{ + {Role: "system", Content: "system rules", StringContent: "system rules"}, + {Role: "user", Content: "old question with many details", StringContent: "old question with many details"}, + {Role: "assistant", Content: "old answer with many details", StringContent: "old answer with many details"}, + {Role: "assistant", Content: "", StringContent: "call", ToolCalls: []schema.ToolCall{{ID: "c1", Type: "function", FunctionCall: schema.FunctionCall{Name: "lookup", Arguments: `{}`}}}}, + {Role: "tool", ToolCallID: "c1", Content: "tool result", StringContent: "tool result"}, + {Role: "user", Content: "latest question", StringContent: "latest question"}, + } + policy := config.CompressionConfig{Enabled: true, TriggerAtRatio: .4, KeepTailTokens: 4, MaxSummaryTokens: 4, CompressorModel: "fast"} + + got, meta, err := service.Transform(context.Background(), policy, 20, "primary", messages) + Expect(err).NotTo(HaveOccurred()) + Expect(compressed).To(HaveLen(4)) + Expect(got).To(HaveLen(3)) + Expect(got[0].Role).To(Equal("system")) + Expect(got[0].StringContent).To(Equal("system rules")) + Expect(got[1].StringContent).To(Equal("[COMPRESSED: facts retained]")) + Expect(compressed[2].ToolCalls[0].ID).To(Equal("c1")) + Expect(compressed[3].ToolCallID).To(Equal("c1")) + Expect(got[2].StringContent).To(Equal("latest question")) + Expect(meta).NotTo(BeNil()) + Expect(*meta).To(Equal(Metadata{OriginalTokens: 17, CompressedTokens: 7, DroppedTurns: 4, Compressor: "fast", SummaryTokens: 2})) + }) + + It("returns overflow when configured to error", func() { + service := New(wordCounter{}, SummarizerFunc(func(context.Context, string, []schema.Message, int) (string, int, error) { + return "summary remains much too large", 5, nil + })) + messages := []schema.Message{ + {Role: "user", Content: "one two three four five", StringContent: "one two three four five"}, + {Role: "user", Content: "six seven eight nine ten", StringContent: "six seven eight nine ten"}, + } + policy := config.CompressionConfig{Enabled: true, TriggerAtRatio: .5, KeepTailTokens: 5, MaxSummaryTokens: 5, OnPostCompressionOverflow: "error"} + + _, _, err := service.Transform(context.Background(), policy, 8, "primary", messages) + Expect(err).To(HaveOccurred()) + Expect(IsOverflow(err)).To(BeTrue()) + }) + + It("drops the oldest summary for overflow recovery", func() { + service := New(wordCounter{}, SummarizerFunc(func(context.Context, string, []schema.Message, int) (string, int, error) { + return "short summary", 2, nil + })) + messages := []schema.Message{ + {Role: "system", Content: "[COMPRESSED: stale summary facts]", StringContent: "[COMPRESSED: stale summary facts]"}, + {Role: "user", Content: "old one two three", StringContent: "old one two three"}, + {Role: "user", Content: "new four five six", StringContent: "new four five six"}, + } + policy := config.CompressionConfig{Enabled: true, TriggerAtRatio: .5, KeepTailTokens: 4, MaxSummaryTokens: 2, OnPostCompressionOverflow: "drop_oldest_summary"} + + got, meta, err := service.Transform(context.Background(), policy, 6, "primary", messages) + Expect(err).NotTo(HaveOccurred()) + Expect(meta).NotTo(BeNil()) + Expect(meta.OverflowRecoveries).To(Equal(2)) + for _, message := range got { + Expect(message.StringContent).NotTo(Equal("[COMPRESSED: stale summary facts]")) + } + }) + + It("never summarizes the only message", func() { + called := false + service := New(wordCounter{}, SummarizerFunc(func(context.Context, string, []schema.Message, int) (string, int, error) { + called = true + return "", 0, nil + })) + messages := []schema.Message{{Role: "user", Content: "one two three four five", StringContent: "one two three four five"}} + + _, _, err := service.Transform(context.Background(), config.CompressionConfig{Enabled: true, TriggerAtRatio: .5}, 4, "primary", messages) + Expect(err).To(HaveOccurred()) + Expect(IsOverflow(err)).To(BeTrue()) + Expect(called).To(BeFalse()) + }) + + It("keeps a single message that fits the hard limit", func() { + called := false + service := New(wordCounter{}, SummarizerFunc(func(context.Context, string, []schema.Message, int) (string, int, error) { + called = true + return "", 0, nil + })) + messages := []schema.Message{{Role: "user", Content: "one two three four five", StringContent: "one two three four five"}} + + got, meta, err := service.Transform(context.Background(), config.CompressionConfig{Enabled: true, TriggerAtRatio: .5}, 6, "primary", messages) + Expect(err).NotTo(HaveOccurred()) + Expect(got).To(Equal(messages)) + Expect(meta).To(BeNil()) + Expect(called).To(BeFalse()) + }) + + It("retains the newest message beyond the tail budget", func() { + service := New(wordCounter{}, SummarizerFunc(func(context.Context, string, []schema.Message, int) (string, int, error) { + return "short", 1, nil + })) + messages := []schema.Message{ + {Role: "user", Content: "old one two three", StringContent: "old one two three"}, + {Role: "user", Content: "active four five six seven", StringContent: "active four five six seven"}, + } + + got, _, err := service.Transform(context.Background(), config.CompressionConfig{Enabled: true, TriggerAtRatio: .5, KeepTailTokens: 1}, 10, "primary", messages) + Expect(err).NotTo(HaveOccurred()) + Expect(got).To(HaveLen(2)) + Expect(got[1].StringContent).To(Equal("active four five six seven")) + }) + + It("does not recursively summarize an existing summary", func() { + service := New(wordCounter{}, SummarizerFunc(func(_ context.Context, _ string, messages []schema.Message, _ int) (string, int, error) { + for _, message := range messages { + Expect(message.StringContent).NotTo(HavePrefix(summaryPrefix)) + } + return "new summary", 2, nil + })) + messages := []schema.Message{ + {Role: "system", Content: "safety rules", StringContent: "safety rules"}, + {Role: "system", Content: "[COMPRESSED: old facts]", StringContent: "[COMPRESSED: old facts]"}, + {Role: "user", Content: "old question details", StringContent: "old question details"}, + {Role: "user", Content: "active question", StringContent: "active question"}, + } + + got, _, err := service.Transform(context.Background(), config.CompressionConfig{Enabled: true, TriggerAtRatio: .5, KeepTailTokens: 2}, 12, "primary", messages) + Expect(err).NotTo(HaveOccurred()) + Expect(got[0].StringContent).To(Equal("safety rules")) + Expect(got[1].StringContent).To(Equal("[COMPRESSED: old facts]")) + }) +}) diff --git a/docs/content/features/context-compression.md b/docs/content/features/context-compression.md new file mode 100644 index 000000000..acc411579 --- /dev/null +++ b/docs/content/features/context-compression.md @@ -0,0 +1,58 @@ +--- +title: "Context compression" +description: "Configure automatic compression for long chat histories" +--- + +Context compression is an opt-in, per-model policy for chat requests that approach +the model context limit. The configuration is disabled by default and does not +change existing requests unless `enabled` is true. + +```yaml +name: long-context-chat +backend: llama-cpp +parameters: + model: chat-model.gguf +compression: + enabled: true + trigger_at_ratio: 0.75 + keep_tail_tokens: 8000 + max_summary_tokens: 2048 + compressor_model: fast-summarizer + on_post_compression_overflow: drop_oldest_summary +``` + +The chat middleware counts the request before inference. Requests below the configured +ratio pass through unchanged. Requests above it replace the oldest complete turns with +a system summary while retaining the newest messages and keeping assistant tool calls +with their tool results. + +Token counts use a conservative byte-level upper-bound estimate so compression never downloads a +tokenizer vocabulary in the request path. Tool schemas and the configured maximum +completion length are included in the context budget. + +- `trigger_at_ratio` selects the fraction of `context_size` that starts compression. +- `keep_tail_tokens` protects the newest part of the conversation from compression. +- `max_summary_tokens` limits the generated summary. +- `compressor_model` selects a secondary model. An empty value selects the primary model. +- `on_post_compression_overflow` selects `drop_oldest_summary` or `error` when the compressed request still exceeds the context limit. + +When omitted, `trigger_at_ratio` defaults to `0.75`, `keep_tail_tokens` to `2048`, +`max_summary_tokens` to `512`, and `on_post_compression_overflow` to `error`. + +Compression applies to `/v1/chat/completions`, `/chat/completions`, and the LocalAI +MCP chat-completion routes. Non-streaming responses include `usage.compression_meta`. +Streaming responses include the same metadata in the trailing usage chunk when the +request sets `stream_options.include_usage`. + +Compression is not supported with `cloud-proxy` passthrough mode because LocalAI +cannot safely rewrite an opaque provider payload. Configure cloud proxy translation +mode to use context compression. + +The compressor model must be installed and configured. If `compressor_model` is empty, +LocalAI uses the primary model. A compressor failure returns an error instead of sending +an over-limit request to the primary model. The `drop_oldest_summary` overflow policy +removes up to two existing summary messages; if the request still does not fit, LocalAI +returns HTTP 413. + +The `/metrics` endpoint exports `localai_compression_events_total`, +`localai_compression_ratio`, and `localai_compression_duration_seconds`. diff --git a/pkg/tokens/count.go b/pkg/tokens/count.go new file mode 100644 index 000000000..3cf0d7be1 --- /dev/null +++ b/pkg/tokens/count.go @@ -0,0 +1,41 @@ +package tokens + +import ( + "encoding/json" + "fmt" + + "github.com/mudler/LocalAI/core/schema" +) + +// CountMessages returns a stable OpenAI-compatible estimate that includes +// roles, multimodal text, tool calls, and tool results. Exact backend counts +// are model-specific; the safety margin comes from the configurable trigger. +func CountMessages(messages []schema.Message) (int, error) { + total := 0 + for _, message := range messages { + payload, err := json.Marshal(message) + if err != nil { + return 0, fmt.Errorf("encode %s message: %w", message.Role, err) + } + // Use a conservative offline estimate. A vocabulary download in the + // request path can hang firewalled installations, while LocalAI must + // decide whether to compress before any model is loaded. One token per + // JSON byte is a safe upper-bound estimate for byte-fallback tokenizers. + // It intentionally triggers compression early instead of risking a late + // backend context rejection for code, identifiers, or multilingual text. + total += 4 + len(payload) + } + return total + 2, nil +} + +// CountPayload estimates non-message request data such as tool schemas. +func CountPayload(value any) (int, error) { + payload, err := json.Marshal(value) + if err != nil { + return 0, fmt.Errorf("encode token payload: %w", err) + } + if string(payload) == "{}" || string(payload) == "null" { + return 0, nil + } + return len(payload), nil +} diff --git a/pkg/tokens/count_test.go b/pkg/tokens/count_test.go new file mode 100644 index 000000000..a2061f58b --- /dev/null +++ b/pkg/tokens/count_test.go @@ -0,0 +1,26 @@ +package tokens + +import ( + "github.com/mudler/LocalAI/core/schema" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +var _ = Describe("Token counting", func() { + It("includes roles, content, and tool payloads", func() { + messages := []schema.Message{ + {Role: "system", Content: "keep this"}, + {Role: "assistant", Content: "", ToolCalls: []schema.ToolCall{{ID: "call-1", Type: "function", FunctionCall: schema.FunctionCall{Name: "lookup", Arguments: `{"id":42}`}}}}, + {Role: "tool", ToolCallID: "call-1", Content: "https://example.com/result"}, + } + + got, err := CountMessages(messages) + Expect(err).NotTo(HaveOccurred()) + Expect(got).To(BeNumerically(">=", 15)) + }) + + It("rejects unsupported content", func() { + _, err := CountMessages([]schema.Message{{Role: "user", Content: make(chan int)}}) + Expect(err).To(HaveOccurred()) + }) +}) diff --git a/pkg/tokens/tokens_suite_test.go b/pkg/tokens/tokens_suite_test.go new file mode 100644 index 000000000..6305df7d4 --- /dev/null +++ b/pkg/tokens/tokens_suite_test.go @@ -0,0 +1,13 @@ +package tokens + +import ( + "testing" + + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +func TestTokens(t *testing.T) { + RegisterFailHandler(Fail) + RunSpecs(t, "Token counting test suite") +}