From 0761bd02c7774789e21176ab8c87b19eeea8edee Mon Sep 17 00:00:00 2001 From: localai-org-maint-bot Date: Tue, 18 Aug 2026 13:31:03 +0200 Subject: [PATCH] feat(chat): add end-to-end context compression (#11556) * feat(config): add context compression policy Define the opt-in model configuration contract before the chat middleware consumes it. Document each policy field so later request handling does not invent a second schema.\n\nRefs #9534\n\nAssisted-by: Codex:gpt-5 * fix(config): register compression fields The model editor metadata gate rejects new config fields without descriptions and suitable controls. Register the compression policy so operators can edit its six fields safely. Assisted-by: Codex:gpt-5 [monitoring-prs] * feat(chat): compress long contexts Long conversations currently fail once they reach the model context window. The opt-in policy now summarizes complete older turns before primary inference and preserves the newest tool chains. Both OpenAI and MCP chat routes share the same transformation. Usage metadata and metrics expose each compression event. Refs #9534 Assisted-by: Codex:gpt-5 --------- Co-authored-by: localai-org-maint-bot <306269227+localai-org-maint-bot@users.noreply.github.com> --- core/config/meta/registry.go | 58 +++++ core/config/model_config.go | 29 +++ core/config/model_config_test.go | 48 +++++ core/http/endpoints/localai/mcp.go | 4 +- core/http/endpoints/openai/chat.go | 38 +++- core/http/middleware/compression.go | 102 +++++++++ core/http/middleware/compression_test.go | 85 ++++++++ core/http/routes/localai.go | 11 +- core/http/routes/openai.go | 16 +- core/schema/openai.go | 14 +- .../compression/compression_suite_test.go | 13 ++ core/services/compression/inference.go | 62 ++++++ core/services/compression/metrics.go | 41 ++++ core/services/compression/service.go | 204 ++++++++++++++++++ core/services/compression/service_test.go | 174 +++++++++++++++ docs/content/features/context-compression.md | 58 +++++ pkg/tokens/count.go | 41 ++++ pkg/tokens/count_test.go | 26 +++ pkg/tokens/tokens_suite_test.go | 13 ++ 19 files changed, 1019 insertions(+), 18 deletions(-) create mode 100644 core/http/middleware/compression.go create mode 100644 core/http/middleware/compression_test.go create mode 100644 core/services/compression/compression_suite_test.go create mode 100644 core/services/compression/inference.go create mode 100644 core/services/compression/metrics.go create mode 100644 core/services/compression/service.go create mode 100644 core/services/compression/service_test.go create mode 100644 docs/content/features/context-compression.md create mode 100644 pkg/tokens/count.go create mode 100644 pkg/tokens/count_test.go create mode 100644 pkg/tokens/tokens_suite_test.go 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") +}