diff --git a/config/config.go b/config/config.go index 92daf99f..00dcdfce 100644 --- a/config/config.go +++ b/config/config.go @@ -47,6 +47,12 @@ type Config struct { // VirtualModels declares redirects, load balancers, and access policies as // infrastructure-as-code. They override admin-store rows of the same source. VirtualModels []VirtualModelConfig `yaml:"virtual_models"` + + // ModelNormalizer declares canonical model aliases rewritten by the gateway + // before provider dispatch. Rules map an alias to a provider/model target + // and optionally pin a thinking policy. The MODEL_NORMALIZER env var + // (JSON array) merges over this list and wins per alias. + ModelNormalizer []ModelNormalizerRule `yaml:"model_normalizer"` } // LoadResult is returned by Load and bundles the application config with the raw @@ -231,6 +237,12 @@ func Load() (*LoadResult, error) { if err := applyVirtualModelsEnv(cfg, strict); err != nil { return nil, err } + if err := applyModelNormalizerEnv(cfg, strict); err != nil { + return nil, err + } + if err := validateModelNormalizerRules(cfg.ModelNormalizer); err != nil { + return nil, err + } if err := applyTaggingEnv(cfg); err != nil { return nil, err } diff --git a/config/modelnormalizer.go b/config/modelnormalizer.go new file mode 100644 index 00000000..e44f2d8c --- /dev/null +++ b/config/modelnormalizer.go @@ -0,0 +1,72 @@ +package config + +import ( + "fmt" + "os" + "strings" +) + +// ModelNormalizerRule declares one model normalization rule in config.yaml or +// the MODEL_NORMALIZER env var. It maps a canonical client-facing alias to a +// provider/model target and optionally pins a thinking policy, mirroring the +// behavior of the internal/modelnormalizer package. Rules are read-only at +// runtime: they live alongside virtual_models as infrastructure-as-code. +type ModelNormalizerRule struct { + // Alias is the client-facing model ID that triggers this rule. + Alias string `yaml:"alias" json:"alias"` + + // Target is the provider/model selector the request is rewritten to, e.g. + // "kimicode/kimi-for-coding". + Target string `yaml:"target" json:"target"` + + // Thinking pins the thinking policy: "enabled", "disabled", or "passthrough". + // Empty is treated as passthrough (no field is injected). + Thinking string `yaml:"thinking,omitempty" json:"thinking,omitempty"` + + // ContextWindow is the advertised context window for /v1/models metadata. + ContextWindow *int `yaml:"context_window,omitempty" json:"context_window,omitempty"` + + // Modes declares the model's kinds for registry metadata, e.g. ["chat"]. + Modes []string `yaml:"modes,omitempty" json:"modes,omitempty"` +} + +const envModelNormalizer = "MODEL_NORMALIZER" + +// validateModelNormalizerRules checks that every rule has a non-empty alias +// and target and a known thinking policy. A rule with an invalid policy would +// silently rewrite requests to an undefined state, so fail fast at load time. +func validateModelNormalizerRules(rules []ModelNormalizerRule) error { + for i, r := range rules { + if strings.TrimSpace(r.Alias) == "" { + return fmt.Errorf("model_normalizer[%d]: alias is required", i) + } + if strings.TrimSpace(r.Target) == "" { + return fmt.Errorf("model_normalizer[%d] (%s): target is required", i, r.Alias) + } + switch strings.TrimSpace(r.Thinking) { + case "", "passthrough", "enabled", "disabled": + default: + return fmt.Errorf("model_normalizer[%d] (%s): thinking must be one of: enabled, disabled, passthrough; got %q", i, r.Alias, r.Thinking) + } + } + return nil +} + +// applyModelNormalizerEnv parses the MODEL_NORMALIZER env var — a JSON array +// of model normalization rules — and merges it over the YAML-declared list. +// Env entries override YAML entries with the same alias (case-insensitive), +// consistent with the rest of the config pipeline where env always wins. +func applyModelNormalizerEnv(cfg *Config, strict bool) error { + raw := strings.TrimSpace(os.Getenv(envModelNormalizer)) + if raw == "" { + return nil + } + var fromEnv []ModelNormalizerRule + if err := decodeIaCJSON(envModelNormalizer, raw, &fromEnv, strict); err != nil { + return fmt.Errorf("invalid %s: %w", envModelNormalizer, err) + } + cfg.ModelNormalizer = mergeByKey(cfg.ModelNormalizer, fromEnv, func(rule ModelNormalizerRule) string { + return canonicalTextKey(rule.Alias) + }) + return nil +} \ No newline at end of file diff --git a/config/modelnormalizer_test.go b/config/modelnormalizer_test.go new file mode 100644 index 00000000..6ef0f286 --- /dev/null +++ b/config/modelnormalizer_test.go @@ -0,0 +1,156 @@ +package config + +import ( + "fmt" + "strings" + "testing" +) + +func TestApplyModelNormalizerEnv_ParsesAndMerges(t *testing.T) { + cfg := &Config{ModelNormalizer: []ModelNormalizerRule{ + {Alias: "kimi-k2.6", Target: "kimicode/kimi-for-coding", Thinking: "disabled"}, + {Alias: "kimi-k2.7-code", Target: "kimicode/kimi-for-coding"}, + }} + t.Setenv(envModelNormalizer, `[ + {"alias":"kimi-k2.6","target":"kimicode/kimi-for-coding","thinking":"passthrough"}, + {"alias":"kimi-k3","target":"kimicode/k3","context_window":1048576,"modes":["chat"]} + ]`) + + if err := applyModelNormalizerEnv(cfg, true); err != nil { + t.Fatalf("applyModelNormalizerEnv() error = %v", err) + } + if len(cfg.ModelNormalizer) != 3 { + t.Fatalf("merged len = %d, want 3", len(cfg.ModelNormalizer)) + } + // "kimi-k2.6" is overridden in place (env wins) and keeps its position. + k26 := cfg.ModelNormalizer[0] + if k26.Alias != "kimi-k2.6" || k26.Thinking != "passthrough" { + t.Fatalf("env did not override kimi-k2.6: %#v", k26) + } + // "kimi-k2.7-code" is untouched; "kimi-k3" is appended. + if cfg.ModelNormalizer[1].Alias != "kimi-k2.7-code" || cfg.ModelNormalizer[2].Alias != "kimi-k3" { + t.Fatalf("merge order wrong: %#v", cfg.ModelNormalizer) + } + // Verify the new entry carries metadata. + k3 := cfg.ModelNormalizer[2] + if k3.ContextWindow == nil || *k3.ContextWindow != 1048576 { + t.Fatalf("kimi-k3 context_window = %v, want 1048576", k3.ContextWindow) + } + if len(k3.Modes) != 1 || k3.Modes[0] != "chat" { + t.Fatalf("kimi-k3 modes = %v, want [chat]", k3.Modes) + } +} + +func TestApplyModelNormalizerEnv_Invalid(t *testing.T) { + cfg := &Config{} + t.Setenv(envModelNormalizer, `{not valid json`) + if err := applyModelNormalizerEnv(cfg, true); err == nil { + t.Fatalf("applyModelNormalizerEnv() error = nil, want parse error") + } +} + +// The env layer overrides YAML entry by entry, so a typo must fail loudly rather +// than let a malformed env entry silently win over a correct YAML one. +func TestApplyModelNormalizerEnv_RejectsUnknownField(t *testing.T) { + cfg := &Config{} + t.Setenv(envModelNormalizer, `[{"alias":"kimi-k2.6","targets":"kimicode/kimi-for-coding"}]`) + + err := applyModelNormalizerEnv(cfg, true) + if err == nil { + t.Fatal("applyModelNormalizerEnv() error = nil, want unknown-field error") + } + if !strings.Contains(err.Error(), "targets") { + t.Fatalf("applyModelNormalizerEnv() error = %q, want it to name the unknown field", err) + } +} + +// json.Decoder stops after the first value and leaves the rest unread, so trailing +// data must be rejected explicitly — silently applying half an env var is the failure +// this path exists to prevent. Structural, therefore fatal in both modes. +func TestApplyModelNormalizerEnv_RejectsTrailingData(t *testing.T) { + trailing := map[string]string{ + "garbage suffix": `[{"alias":"a","target":"b"}] and then some junk`, + "second JSON value": `[{"alias":"a","target":"b"}] {"alias":"c","target":"d"}`, + "second JSON on a line": "[{\"alias\":\"a\",\"target\":\"b\"}]\n{\"alias\":\"c\",\"target\":\"d\"}", + } + for name, raw := range trailing { + for _, strict := range []bool{true, false} { + t.Run(fmt.Sprintf("%s/strict=%v", name, strict), func(t *testing.T) { + cfg := &Config{} + t.Setenv(envModelNormalizer, raw) + + err := applyModelNormalizerEnv(cfg, strict) + if err == nil { + t.Fatal("applyModelNormalizerEnv() error = nil, want trailing-data error") + } + if !strings.Contains(err.Error(), "unexpected data after the JSON value") { + t.Fatalf("applyModelNormalizerEnv() error = %q, want a trailing-data error", err) + } + }) + } + } +} + +func TestValidateModelNormalizerRules(t *testing.T) { + tests := []struct { + name string + rules []ModelNormalizerRule + wantErr string + }{ + {name: "nil is valid"}, + {name: "empty slice is valid", rules: []ModelNormalizerRule{}}, + {name: "valid rule", rules: []ModelNormalizerRule{{Alias: "a", Target: "b"}}}, + {name: "valid with thinking", rules: []ModelNormalizerRule{{Alias: "a", Target: "b", Thinking: "disabled"}}}, + {name: "valid with passthrough", rules: []ModelNormalizerRule{{Alias: "a", Target: "b", Thinking: "passthrough"}}}, + {name: "missing alias", rules: []ModelNormalizerRule{{Target: "b"}}, wantErr: "alias is required"}, + {name: "missing target", rules: []ModelNormalizerRule{{Alias: "a"}}, wantErr: "target is required"}, + {name: "invalid thinking", rules: []ModelNormalizerRule{{Alias: "a", Target: "b", Thinking: "bogus"}}, wantErr: "thinking must be one of"}, + {name: "whitespace alias", rules: []ModelNormalizerRule{{Alias: " ", Target: "b"}}, wantErr: "alias is required"}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + err := validateModelNormalizerRules(tt.rules) + if tt.wantErr == "" { + if err != nil { + t.Fatalf("validateModelNormalizerRules() error = %v, want nil", err) + } + return + } + if err == nil { + t.Fatalf("validateModelNormalizerRules() error = nil, want %q", tt.wantErr) + } + if !strings.Contains(err.Error(), tt.wantErr) { + t.Fatalf("validateModelNormalizerRules() error = %q, want to contain %q", err, tt.wantErr) + } + }) + } +} + +func TestApplyModelNormalizerEnv_NoOpWhenUnset(t *testing.T) { + cfg := &Config{ModelNormalizer: []ModelNormalizerRule{ + {Alias: "kimi-k2.6", Target: "kimicode/kimi-for-coding"}, + }} + if err := applyModelNormalizerEnv(cfg, true); err != nil { + t.Fatalf("applyModelNormalizerEnv() error = %v", err) + } + if len(cfg.ModelNormalizer) != 1 { + t.Fatalf("len = %d, want 1 (unchanged)", len(cfg.ModelNormalizer)) + } +} + +func TestApplyModelNormalizerEnv_CaseInsensitiveAliasKey(t *testing.T) { + cfg := &Config{ModelNormalizer: []ModelNormalizerRule{ + {Alias: "Kimi-K2.6", Target: "kimicode/kimi-for-coding", Thinking: "disabled"}, + }} + t.Setenv(envModelNormalizer, `[{"alias":"kimi-k2.6","target":"kimicode/kimi-for-coding","thinking":"passthrough"}]`) + + if err := applyModelNormalizerEnv(cfg, true); err != nil { + t.Fatalf("applyModelNormalizerEnv() error = %v", err) + } + if len(cfg.ModelNormalizer) != 1 { + t.Fatalf("len = %d, want 1 (case-insensitive alias key)", len(cfg.ModelNormalizer)) + } + if cfg.ModelNormalizer[0].Thinking != "passthrough" { + t.Fatalf("thinking = %q, want passthrough", cfg.ModelNormalizer[0].Thinking) + } +} diff --git a/internal/app/app.go b/internal/app/app.go index 07bcb70a..d18d7bad 100644 --- a/internal/app/app.go +++ b/internal/app/app.go @@ -34,6 +34,7 @@ import ( "github.com/enterpilot/gomodel/internal/live" "github.com/enterpilot/gomodel/internal/llmclient" "github.com/enterpilot/gomodel/internal/mcpgateway" + "github.com/enterpilot/gomodel/internal/modelnormalizer" "github.com/enterpilot/gomodel/internal/pricingoverrides" "github.com/enterpilot/gomodel/internal/providers" "github.com/enterpilot/gomodel/internal/providers/health" @@ -746,6 +747,7 @@ func New(ctx context.Context, cfg Config) (*App, error) { TranslatedRequestPatcher: translatedRequestPatcher, BatchRequestPreparer: batchRequestPreparer, ExposedModelLister: vm, + ModelNormalizer: modelnormalizer.BuildFromConfig(appCfg.ModelNormalizer), KeepOnlyAliasesAtModelsEndpoint: appCfg.Models.KeepOnlyAliasesAtModelsEndpoint, PassthroughSemanticEnrichers: cfg.Factory.PassthroughSemanticEnrichers(), BatchStore: batchResult.Store, diff --git a/internal/modelnormalizer/config.go b/internal/modelnormalizer/config.go new file mode 100644 index 00000000..958d834a --- /dev/null +++ b/internal/modelnormalizer/config.go @@ -0,0 +1,53 @@ +package modelnormalizer + +import ( + "github.com/enterpilot/gomodel/config" + + "github.com/enterpilot/gomodel/internal/core" +) + +// BuildFromConfig constructs a Normalizer from a config.Config slice. Empty or +// blank rules are silently dropped. The result is nil when no valid rules are +// declared, so callers can pass it directly to server/gateway hooks without +// nil-checking downstream. +func BuildFromConfig(rules []config.ModelNormalizerRule) *Normalizer { + if len(rules) == 0 { + return nil + } + converted := make([]Rule, 0, len(rules)) + for _, r := range rules { + converted = append(converted, Rule{ + Alias: r.Alias, + Target: r.Target, + Thinking: ThinkingPolicy(r.Thinking), + ContextWindow: r.ContextWindow, + Modes: r.Modes, + }) + } + return New(converted) +} + +// ChainedExposedModelLister wraps two ExposedModels functions so the +// secondary's output is appended to the primary. The server layer adapts this +// onto its ExposedModelLister interface. Used when a normalizer is configured +// alongside another lister (e.g. virtual models) so canonical aliases still +// appear in /v1/models. +type ChainedExposedModelLister struct { + Primary func() []core.Model + Secondary func() []core.Model +} + +// ExposedModels concatenates both sources. A nil primary collapses to the +// secondary; a nil secondary collapses to the primary; both nil returns nil. +func (c ChainedExposedModelLister) ExposedModels() []core.Model { + if c.Primary == nil { + if c.Secondary == nil { + return nil + } + return c.Secondary() + } + if c.Secondary == nil { + return c.Primary() + } + return append(c.Primary(), c.Secondary()...) +} \ No newline at end of file diff --git a/internal/modelnormalizer/config_test.go b/internal/modelnormalizer/config_test.go new file mode 100644 index 00000000..4bf59093 --- /dev/null +++ b/internal/modelnormalizer/config_test.go @@ -0,0 +1,125 @@ +package modelnormalizer + +import ( + "testing" + + "github.com/enterpilot/gomodel/config" + + "github.com/enterpilot/gomodel/internal/core" +) + +func TestBuildFromConfig(t *testing.T) { + cw := 262144 + cfgRules := []config.ModelNormalizerRule{ + {Alias: "kimi-k2.6", Target: "kimicode/kimi-for-coding", Thinking: "disabled", ContextWindow: &cw, Modes: []string{"chat"}}, + {Alias: "kimi-k3", Target: "kimicode/k3"}, + } + + n := BuildFromConfig(cfgRules) + if n == nil { + t.Fatal("BuildFromConfig() returned nil for valid rules") + } + + r, ok := n.Lookup("kimi-k2.6") + if !ok { + t.Fatal("Lookup kimi-k2.6 not found") + } + if r.Target != "kimicode/kimi-for-coding" { + t.Fatalf("Target = %q, want %q", r.Target, "kimicode/kimi-for-coding") + } + if r.Thinking != ThinkingDisabled { + t.Fatalf("Thinking = %q, want %q", r.Thinking, ThinkingDisabled) + } + if r.ContextWindow == nil || *r.ContextWindow != 262144 { + t.Fatalf("ContextWindow = %v, want 262144", r.ContextWindow) + } + if len(r.Modes) != 1 || r.Modes[0] != "chat" { + t.Fatalf("Modes = %v, want [chat]", r.Modes) + } +} + +func TestBuildFromConfig_NilForEmpty(t *testing.T) { + if n := BuildFromConfig(nil); n != nil { + t.Fatalf("BuildFromConfig(nil) = %v, want nil", n) + } + if n := BuildFromConfig([]config.ModelNormalizerRule{}); n != nil { + t.Fatalf("BuildFromConfig(empty) = %v, want nil", n) + } +} + +func TestBuildFromConfig_SkipsBlankRules(t *testing.T) { + cfgRules := []config.ModelNormalizerRule{ + {Alias: "", Target: "x"}, + {Alias: "a", Target: ""}, + {Alias: "valid", Target: "b"}, + } + n := BuildFromConfig(cfgRules) + if n == nil { + t.Fatal("BuildFromConfig() returned nil, want non-nil with one valid rule") + } + aliases := n.Aliases() + if len(aliases) != 1 || aliases[0] != "valid" { + t.Fatalf("Aliases() = %v, want [valid]", aliases) + } +} + +func TestExposedModels_SatisfiesListModelsHook(t *testing.T) { + cw := 262144 + n := New([]Rule{ + {Alias: "kimi-k2.6", Target: "kimicode/kimi-for-coding", Modes: []string{"chat"}, ContextWindow: &cw}, + {Alias: "bge_m3_embed", Target: "kimicode/bge_m3_embed", Modes: []string{"embedding"}}, + }) + + models := n.ExposedModels() + if len(models) != 2 { + t.Fatalf("ExposedModels() len = %d, want 2", len(models)) + } + + byID := make(map[string]core.Model, len(models)) + for _, m := range models { + byID[m.ID] = m + } + + embed := byID["bge_m3_embed"] + if embed.Metadata == nil { + t.Fatal("bge_m3_embed metadata is nil") + } + if len(embed.Metadata.Categories) != 1 || embed.Metadata.Categories[0] != core.CategoryEmbedding { + t.Fatalf("bge_m3_embed categories = %v, want [embedding]", embed.Metadata.Categories) + } +} + +func TestExposedModels_NilNormalizer(t *testing.T) { + var n *Normalizer + if models := n.ExposedModels(); models != nil { + t.Fatalf("nil Normalizer ExposedModels() = %v, want nil", models) + } +} + +func TestChainedExposedModelLister_Concatenates(t *testing.T) { + primary := func() []core.Model { return []core.Model{{ID: "a"}, {ID: "b"}} } + secondary := func() []core.Model { return []core.Model{{ID: "c"}} } + + chain := ChainedExposedModelLister{Primary: primary, Secondary: secondary} + got := chain.ExposedModels() + if len(got) != 3 { + t.Fatalf("len = %d, want 3", len(got)) + } + if got[0].ID != "a" || got[1].ID != "b" || got[2].ID != "c" { + t.Fatalf("order = [%s, %s, %s], want [a, b, c]", got[0].ID, got[1].ID, got[2].ID) + } +} + +func TestChainedExposedModelLister_NilPrimaries(t *testing.T) { + secondary := func() []core.Model { return []core.Model{{ID: "c"}} } + + if got := (ChainedExposedModelLister{Secondary: secondary}).ExposedModels(); len(got) != 1 { + t.Fatalf("nil primary: len = %d, want 1", len(got)) + } + if got := (ChainedExposedModelLister{Primary: secondary}).ExposedModels(); len(got) != 1 { + t.Fatalf("nil secondary: len = %d, want 1", len(got)) + } + if got := (ChainedExposedModelLister{}).ExposedModels(); got != nil { + t.Fatalf("both nil: got %v, want nil", got) + } +} \ No newline at end of file diff --git a/internal/modelnormalizer/normalizer.go b/internal/modelnormalizer/normalizer.go new file mode 100644 index 00000000..ae19eba2 --- /dev/null +++ b/internal/modelnormalizer/normalizer.go @@ -0,0 +1,263 @@ +// Package modelnormalizer provides a data-driven, gateway-edge model +// normalizer that rewrites chat requests before provider dispatch. It resolves +// canonical client-facing model aliases (e.g. "kimi-k2.6") to concrete +// provider/model targets and injects per-alias thinking policies, without +// touching any provider package. +// +// The normalizer runs as the first step of the chat prepare path (before +// model resolution), so the resolver and every downstream provider only see +// the already-rewritten request. Unknown model IDs pass through unchanged +// (Postel's law). +package modelnormalizer + +import ( + "fmt" + "sort" + "strings" + + "github.com/goccy/go-json" + + "github.com/enterpilot/gomodel/internal/core" +) + +// ThinkingPolicy controls how the normalizer sets the `thinking` extension on +// the rewritten request. +type ThinkingPolicy string + +const ( + // ThinkingPassthrough leaves any existing thinking field untouched and + // injects nothing. The provider sees whatever the client sent (or nothing). + ThinkingPassthrough ThinkingPolicy = "passthrough" + // ThinkingEnabled injects {"thinking":{"type":"enabled"}} so the upstream + // activates extended reasoning regardless of the client payload. + ThinkingEnabled ThinkingPolicy = "enabled" + // ThinkingDisabled injects {"thinking":{"type":"disabled"}} so the upstream + // suppresses extended reasoning regardless of the client payload. + ThinkingDisabled ThinkingPolicy = "disabled" +) + +// Valid reports whether the policy is a known value. Empty string is treated +// as passthrough (the zero value). +func (p ThinkingPolicy) Valid() bool { + switch p { + case "", ThinkingPassthrough, ThinkingEnabled, ThinkingDisabled: + return true + default: + return false + } +} + +// Rule declares one model normalization rule. It maps an alias to a +// provider-qualified target and optionally pins a thinking policy and/or +// advertises metadata for /v1/models. +type Rule struct { + // Alias is the client-facing model ID that triggers this rule. + Alias string `yaml:"alias" json:"alias"` + // Target is the provider/model selector the request is rewritten to, e.g. + // "kimicode/kimi-for-coding". + Target string `yaml:"target" json:"target"` + // Thinking is the thinking policy to inject. One of "enabled", "disabled", + // or "passthrough" (default: passthrough). + Thinking ThinkingPolicy `yaml:"thinking,omitempty" json:"thinking,omitempty"` + // ContextWindow is the advertised context window for /v1/models metadata. + ContextWindow *int `yaml:"context_window,omitempty" json:"context_window,omitempty"` + // Modes declares the model's kinds for registry metadata, e.g. ["chat"]. + // Used to derive Categories via core.CategoriesForModes. + Modes []string `yaml:"modes,omitempty" json:"modes,omitempty"` +} + +// Normalizer applies rules to rewrite chat requests before provider dispatch. +// Rules are indexed by alias for O(1) lookup. +type Normalizer struct { + byAlias map[string]Rule +} + +// New creates a Normalizer from the given rules. Whitespace-only alias or +// target entries are skipped. Duplicate aliases are resolved last-wins so the +// env layer can override the YAML layer deterministically. +func New(rules []Rule) *Normalizer { + if len(rules) == 0 { + return nil + } + byAlias := make(map[string]Rule, len(rules)) + for _, r := range rules { + alias := strings.TrimSpace(r.Alias) + target := strings.TrimSpace(r.Target) + if alias == "" || target == "" { + continue + } + r.Alias = alias + r.Target = target + byAlias[alias] = r + } + if len(byAlias) == 0 { + return nil + } + return &Normalizer{byAlias: byAlias} +} + +// Lookup returns the rule registered under alias, if any. +func (n *Normalizer) Lookup(alias string) (Rule, bool) { + if n == nil { + return Rule{}, false + } + alias = strings.TrimSpace(alias) + if alias == "" { + return Rule{}, false + } + r, ok := n.byAlias[alias] + return r, ok +} + +// Aliases returns the registered aliases in deterministic order. +func (n *Normalizer) Aliases() []string { + if n == nil || len(n.byAlias) == 0 { + return nil + } + aliases := make([]string, 0, len(n.byAlias)) + for a := range n.byAlias { + aliases = append(aliases, a) + } + // sort.Strings is deterministic; no dedup needed since map keys are unique. + sort.Strings(aliases) + return aliases +} + +// AdaptChatRequest rewrites req.Model to the rule's target and injects the +// per-alias thinking policy into ExtraFields. It returns (rewrittenReq, true, +// nil) when a rule matched, or (req, false, nil) when no rule matched — the +// original request pointer is returned unchanged so the caller can use it +// directly without copying. +func (n *Normalizer) AdaptChatRequest(req *core.ChatRequest) (*core.ChatRequest, bool, error) { + if n == nil || req == nil { + return req, false, nil + } + r, ok := n.Lookup(req.Model) + if !ok { + return req, false, nil + } + + // Shallow-copy the request so the caller's original is never mutated. + adapted := *req + adapted.Model = r.Target + + // Inject the thinking policy when the rule pins one. + if r.Thinking != "" && r.Thinking != ThinkingPassthrough { + thinking, err := thinkingJSON(r.Thinking) + if err != nil { + return req, false, fmt.Errorf("modelnormalizer: rule %q: %w", r.Alias, err) + } + extra, err := core.MergeUnknownJSONFields(req.ExtraFields, map[string]json.RawMessage{ + "thinking": thinking, + }) + if err != nil { + return req, false, fmt.Errorf("modelnormalizer: rule %q: merge thinking: %w", r.Alias, err) + } + adapted.ExtraFields = extra + } + + return &adapted, true, nil +} + +// thinkingJSON serializes the thinking extension for the given policy. +// The shape matches what providers like Kimi Code and Xiaomi expect: +// {"type":"enabled"} or {"type":"disabled"}. +func thinkingJSON(policy ThinkingPolicy) (json.RawMessage, error) { + switch policy { + case ThinkingEnabled: + return json.RawMessage(`{"type":"enabled"}`), nil + case ThinkingDisabled: + return json.RawMessage(`{"type":"disabled"}`), nil + default: + return nil, fmt.Errorf("unknown thinking policy %q", policy) + } +} + +// MergeRules layers override rules over base rules, keyed by alias. Later +// entries in override replace matching base entries in place and append new +// ones. This mirrors the VIRTUAL_MODELS env layering pattern. +func MergeRules(base, override []Rule) []Rule { + if len(override) == 0 { + return base + } + merged := make([]Rule, len(base)) + copy(merged, base) + index := make(map[string]int, len(merged)) + for i, r := range merged { + index[strings.ToLower(strings.TrimSpace(r.Alias))] = i + } + for _, r := range override { + key := strings.ToLower(strings.TrimSpace(r.Alias)) + if pos, ok := index[key]; ok { + merged[pos] = r + continue + } + index[key] = len(merged) + merged = append(merged, r) + } + return merged +} + +// MergeMetadata synthesizes core.Model entries for each rule so /v1/models +// can advertise the canonical aliases alongside provider models. Rules with +// no modes or context_window are still listed (metadata is optional). +func MergeMetadata(rules []Rule) []core.Model { + if len(rules) == 0 { + return nil + } + models := make([]core.Model, 0, len(rules)) + seen := make(map[string]struct{}, len(rules)) + for _, r := range rules { + alias := strings.TrimSpace(r.Alias) + if alias == "" { + continue + } + if _, dup := seen[alias]; dup { + continue + } + seen[alias] = struct{}{} + + entry := core.Model{ + ID: alias, + Object: "model", + } + entry.Metadata = ruleMetadata(r) + models = append(models, entry) + } + return models +} + +func ruleMetadata(r Rule) *core.ModelMetadata { + var meta *core.ModelMetadata + if len(r.Modes) > 0 { + modes := make([]string, len(r.Modes)) + copy(modes, r.Modes) + if meta == nil { + meta = &core.ModelMetadata{} + } + meta.Modes = modes + meta.Categories = core.CategoriesForModes(modes) + } + if r.ContextWindow != nil { + if meta == nil { + meta = &core.ModelMetadata{} + } + meta.ContextWindow = r.ContextWindow + } + return meta +} + +// ExposedModels returns one core.Model entry per registered alias so the +// /v1/models handler can advertise canonical names alongside provider models. +// The normalizer satisfies the ExposedModelLister interface consumed by the +// handler's ListModels merge. Rules are deduplicated by alias (last-wins). +func (n *Normalizer) ExposedModels() []core.Model { + if n == nil { + return nil + } + rules := make([]Rule, 0, len(n.byAlias)) + for _, r := range n.byAlias { + rules = append(rules, r) + } + return MergeMetadata(rules) +} diff --git a/internal/modelnormalizer/normalizer_test.go b/internal/modelnormalizer/normalizer_test.go new file mode 100644 index 00000000..c06add96 --- /dev/null +++ b/internal/modelnormalizer/normalizer_test.go @@ -0,0 +1,399 @@ +package modelnormalizer + +import ( + "testing" + + "github.com/goccy/go-json" + + "github.com/enterpilot/gomodel/internal/core" +) + +func TestAdaptChatRequest_RewritesModel(t *testing.T) { + n := New([]Rule{ + {Alias: "kimi-k2.6", Target: "kimicode/kimi-for-coding", Thinking: ThinkingDisabled}, + }) + + req := &core.ChatRequest{Model: "kimi-k2.6"} + adapted, rewritten, err := n.AdaptChatRequest(req) + if err != nil { + t.Fatalf("AdaptChatRequest() error = %v", err) + } + if !rewritten { + t.Fatal("AdaptChatRequest() rewritten = false, want true") + } + if adapted.Model != "kimicode/kimi-for-coding" { + t.Fatalf("AdaptChatRequest() model = %q, want %q", adapted.Model, "kimicode/kimi-for-coding") + } + if adapted == req { + t.Fatal("AdaptChatRequest() returned the same pointer, want a copy") + } +} + +func TestAdaptChatRequest_InjectsThinkingDisabled(t *testing.T) { + n := New([]Rule{ + {Alias: "kimi-k2.6", Target: "kimicode/kimi-for-coding", Thinking: ThinkingDisabled}, + }) + + req := &core.ChatRequest{Model: "kimi-k2.6"} + adapted, rewritten, err := n.AdaptChatRequest(req) + if err != nil { + t.Fatalf("AdaptChatRequest() error = %v", err) + } + if !rewritten { + t.Fatal("AdaptChatRequest() rewritten = false, want true") + } + + thinking := adapted.ExtraFields.Lookup("thinking") + if thinking == nil { + t.Fatal("AdaptChatRequest() thinking field missing from ExtraFields") + } + var parsed struct { + Type string `json:"type"` + } + if err := json.Unmarshal(thinking, &parsed); err != nil { + t.Fatalf("unmarshal thinking: %v", err) + } + if parsed.Type != "disabled" { + t.Fatalf("thinking.type = %q, want %q", parsed.Type, "disabled") + } +} + +func TestAdaptChatRequest_InjectsThinkingEnabled(t *testing.T) { + n := New([]Rule{ + {Alias: "kimi-k2.7-code", Target: "kimicode/kimi-for-coding", Thinking: ThinkingEnabled}, + }) + + req := &core.ChatRequest{Model: "kimi-k2.7-code"} + adapted, rewritten, err := n.AdaptChatRequest(req) + if err != nil { + t.Fatalf("AdaptChatRequest() error = %v", err) + } + if !rewritten { + t.Fatal("AdaptChatRequest() rewritten = false, want true") + } + + thinking := adapted.ExtraFields.Lookup("thinking") + if thinking == nil { + t.Fatal("AdaptChatRequest() thinking field missing from ExtraFields") + } + var parsed struct { + Type string `json:"type"` + } + if err := json.Unmarshal(thinking, &parsed); err != nil { + t.Fatalf("unmarshal thinking: %v", err) + } + if parsed.Type != "enabled" { + t.Fatalf("thinking.type = %q, want %q", parsed.Type, "enabled") + } +} + +func TestAdaptChatRequest_PassthroughLeavesThinkingUntouched(t *testing.T) { + n := New([]Rule{ + {Alias: "kimi-k3", Target: "kimicode/k3", Thinking: ThinkingPassthrough}, + }) + + req := &core.ChatRequest{Model: "kimi-k3"} + adapted, rewritten, err := n.AdaptChatRequest(req) + if err != nil { + t.Fatalf("AdaptChatRequest() error = %v", err) + } + if !rewritten { + t.Fatal("AdaptChatRequest() rewritten = false, want true") + } + if adapted.Model != "kimicode/k3" { + t.Fatalf("AdaptChatRequest() model = %q, want %q", adapted.Model, "kimicode/k3") + } + // No thinking field should be injected. + if thinking := adapted.ExtraFields.Lookup("thinking"); thinking != nil { + t.Fatalf("thinking = %s, want nil for passthrough", thinking) + } +} + +func TestAdaptChatRequest_PreservesExistingExtraFields(t *testing.T) { + n := New([]Rule{ + {Alias: "kimi-k2.6", Target: "kimicode/kimi-for-coding", Thinking: ThinkingDisabled}, + }) + + req := &core.ChatRequest{ + Model: "kimi-k2.6", + ExtraFields: core.UnknownJSONFieldsFromMap(map[string]json.RawMessage{ + "custom": json.RawMessage(`"value"`), + }), + } + adapted, rewritten, err := n.AdaptChatRequest(req) + if err != nil { + t.Fatalf("AdaptChatRequest() error = %v", err) + } + if !rewritten { + t.Fatal("AdaptChatRequest() rewritten = false, want true") + } + // Both the original field and the injected thinking should be present. + if custom := adapted.ExtraFields.Lookup("custom"); custom == nil { + t.Fatal("custom field lost from ExtraFields") + } + if thinking := adapted.ExtraFields.Lookup("thinking"); thinking == nil { + t.Fatal("thinking field missing from ExtraFields") + } +} + +func TestAdaptChatRequest_UnknownModelPassesThrough(t *testing.T) { + n := New([]Rule{ + {Alias: "kimi-k2.6", Target: "kimicode/kimi-for-coding"}, + }) + + req := &core.ChatRequest{Model: "unknown-model"} + adapted, rewritten, err := n.AdaptChatRequest(req) + if err != nil { + t.Fatalf("AdaptChatRequest() error = %v", err) + } + if rewritten { + t.Fatal("AdaptChatRequest() rewritten = true, want false for unknown model") + } + if adapted != req { + t.Fatal("AdaptChatRequest() returned a different pointer for unknown model, want same") + } + if adapted.Model != "unknown-model" { + t.Fatalf("AdaptChatRequest() model = %q, want %q", adapted.Model, "unknown-model") + } +} + +func TestAdaptChatRequest_NilNormalizerPassesThrough(t *testing.T) { + var n *Normalizer + req := &core.ChatRequest{Model: "kimi-k2.6"} + adapted, rewritten, err := n.AdaptChatRequest(req) + if err != nil { + t.Fatalf("AdaptChatRequest() error = %v", err) + } + if rewritten { + t.Fatal("AdaptChatRequest() rewritten = true, want false for nil normalizer") + } + if adapted != req { + t.Fatal("AdaptChatRequest() returned a different pointer, want same") + } +} + +func TestAdaptChatRequest_NilRequestPassesThrough(t *testing.T) { + n := New([]Rule{{Alias: "a", Target: "b"}}) + adapted, rewritten, err := n.AdaptChatRequest(nil) + if err != nil { + t.Fatalf("AdaptChatRequest() error = %v", err) + } + if rewritten { + t.Fatal("AdaptChatRequest() rewritten = true, want false for nil request") + } + if adapted != nil { + t.Fatal("AdaptChatRequest() returned non-nil for nil request") + } +} + +func TestAdaptChatRequest_EmptyThinkingIsPassthrough(t *testing.T) { + n := New([]Rule{ + {Alias: "kimi-k3", Target: "kimicode/k3"}, // no thinking field + }) + + req := &core.ChatRequest{Model: "kimi-k3"} + adapted, rewritten, err := n.AdaptChatRequest(req) + if err != nil { + t.Fatalf("AdaptChatRequest() error = %v", err) + } + if !rewritten { + t.Fatal("AdaptChatRequest() rewritten = false, want true") + } + if thinking := adapted.ExtraFields.Lookup("thinking"); thinking != nil { + t.Fatalf("thinking = %s, want nil when no thinking policy set", thinking) + } +} + +func TestNew_SkipsEmptyAliasOrTarget(t *testing.T) { + n := New([]Rule{ + {Alias: "", Target: "kimicode/kimi-for-coding"}, + {Alias: "kimi-k2.6", Target: ""}, + {Alias: " ", Target: "kimicode/kimi-for-coding"}, + {Alias: "kimi-k2.6", Target: "kimicode/kimi-for-coding"}, + }) + if n == nil { + t.Fatal("New() returned nil, want non-nil with one valid rule") + } + aliases := n.Aliases() + if len(aliases) != 1 || aliases[0] != "kimi-k2.6" { + t.Fatalf("Aliases() = %v, want [kimi-k2.6]", aliases) + } +} + +func TestNew_LastWinsForDuplicateAlias(t *testing.T) { + n := New([]Rule{ + {Alias: "kimi-k2.6", Target: "kimicode/kimi-for-coding", Thinking: ThinkingDisabled}, + {Alias: "kimi-k2.6", Target: "kimicode/k3", Thinking: ThinkingPassthrough}, + }) + if n == nil { + t.Fatal("New() returned nil") + } + r, ok := n.Lookup("kimi-k2.6") + if !ok { + t.Fatal("Lookup() not found") + } + if r.Target != "kimicode/k3" { + t.Fatalf("Target = %q, want %q (last wins)", r.Target, "kimicode/k3") + } + if r.Thinking != ThinkingPassthrough { + t.Fatalf("Thinking = %q, want %q (last wins)", r.Thinking, ThinkingPassthrough) + } +} + +func TestNew_NilForNoValidRules(t *testing.T) { + n := New(nil) + if n != nil { + t.Fatal("New(nil) returned non-nil, want nil") + } + n = New([]Rule{{Alias: "", Target: ""}}) + if n != nil { + t.Fatal("New(empty rules) returned non-nil, want nil") + } +} + +func TestMergeRules_EnvOverridesConfig(t *testing.T) { + base := []Rule{ + {Alias: "kimi-k2.6", Target: "kimicode/kimi-for-coding", Thinking: ThinkingDisabled}, + {Alias: "kimi-k2.7-code", Target: "kimicode/kimi-for-coding"}, + } + override := []Rule{ + {Alias: "kimi-k2.6", Target: "kimicode/kimi-for-coding", Thinking: ThinkingPassthrough}, + {Alias: "kimi-k3", Target: "kimicode/k3"}, + } + + merged := MergeRules(base, override) + if len(merged) != 3 { + t.Fatalf("MergeRules() len = %d, want 3", len(merged)) + } + + // kimi-k2.6 is overridden in place (index 0). + if merged[0].Alias != "kimi-k2.6" || merged[0].Thinking != ThinkingPassthrough { + t.Fatalf("kimi-k2.6 not overridden: %+v", merged[0]) + } + // kimi-k2.7-code is untouched. + if merged[1].Alias != "kimi-k2.7-code" { + t.Fatalf("kimi-k2.7-code moved or changed: %+v", merged[1]) + } + // kimi-k3 is appended. + if merged[2].Alias != "kimi-k3" { + t.Fatalf("kimi-k3 not appended: %+v", merged[2]) + } +} + +func TestMergeRules_NilOverrideReturnsBase(t *testing.T) { + base := []Rule{{Alias: "a", Target: "b"}} + merged := MergeRules(base, nil) + if len(merged) != 1 || merged[0].Alias != "a" { + t.Fatalf("MergeRules() = %v, want base unchanged", merged) + } +} + +func TestMergeMetadata_SynthesizesModels(t *testing.T) { + cw := 262144 + rules := []Rule{ + {Alias: "kimi-k2.6", Target: "kimicode/kimi-for-coding", Modes: []string{"chat"}, ContextWindow: &cw}, + {Alias: "kimi-k3", Target: "kimicode/k3", Modes: []string{"chat"}}, + {Alias: "bge_m3_embed", Target: "kimicode/bge_m3_embed", Modes: []string{"embedding"}}, + } + + models := MergeMetadata(rules) + if len(models) != 3 { + t.Fatalf("MergeMetadata() len = %d, want 3", len(models)) + } + + byID := make(map[string]core.Model, len(models)) + for _, m := range models { + byID[m.ID] = m + } + + k26 := byID["kimi-k2.6"] + if k26.ID != "kimi-k2.6" || k26.Object != "model" { + t.Fatalf("kimi-k2.6 entry wrong: %+v", k26) + } + if k26.Metadata == nil { + t.Fatal("kimi-k2.6 metadata is nil") + } + if len(k26.Metadata.Modes) != 1 || k26.Metadata.Modes[0] != "chat" { + t.Fatalf("kimi-k2.6 modes = %v, want [chat]", k26.Metadata.Modes) + } + if k26.Metadata.ContextWindow == nil || *k26.Metadata.ContextWindow != 262144 { + t.Fatalf("kimi-k2.6 context_window = %v, want 262144", k26.Metadata.ContextWindow) + } + if len(k26.Metadata.Categories) != 1 || k26.Metadata.Categories[0] != core.CategoryTextGeneration { + t.Fatalf("kimi-k2.6 categories = %v, want [text_generation]", k26.Metadata.Categories) + } + + embed := byID["bge_m3_embed"] + if embed.Metadata == nil { + t.Fatal("bge_m3_embed metadata is nil") + } + if len(embed.Metadata.Modes) != 1 || embed.Metadata.Modes[0] != "embedding" { + t.Fatalf("bge_m3_embed modes = %v, want [embedding]", embed.Metadata.Modes) + } + if len(embed.Metadata.Categories) != 1 || embed.Metadata.Categories[0] != core.CategoryEmbedding { + t.Fatalf("bge_m3_embed categories = %v, want [embedding]", embed.Metadata.Categories) + } +} + +func TestMergeMetadata_NoMetadataFields(t *testing.T) { + rules := []Rule{ + {Alias: "kimi-k3", Target: "kimicode/k3"}, // no modes, no context_window + } + + models := MergeMetadata(rules) + if len(models) != 1 { + t.Fatalf("MergeMetadata() len = %d, want 1", len(models)) + } + if models[0].Metadata != nil { + t.Fatalf("Metadata = %+v, want nil when no metadata fields set", models[0].Metadata) + } +} + +func TestMergeMetadata_DedupesAliases(t *testing.T) { + rules := []Rule{ + {Alias: "kimi-k2.6", Target: "kimicode/kimi-for-coding"}, + {Alias: "kimi-k2.6", Target: "kimicode/kimi-for-coding"}, + } + models := MergeMetadata(rules) + if len(models) != 1 { + t.Fatalf("MergeMetadata() len = %d, want 1 (deduped)", len(models)) + } +} + +func TestMergeMetadata_NilForEmpty(t *testing.T) { + if models := MergeMetadata(nil); len(models) != 0 { + t.Fatalf("MergeMetadata(nil) = %v, want empty", models) + } + // Rules with blank aliases are silently skipped, so the result is empty + // (but not nil — a defensive caller's length check still works). + models := MergeMetadata([]Rule{{Alias: "", Target: "x"}}) + if len(models) != 0 { + t.Fatalf("MergeMetadata(blank alias) = %v, want empty", models) + } +} + +func TestThinkingPolicyValid(t *testing.T) { + valid := []ThinkingPolicy{"", ThinkingPassthrough, ThinkingEnabled, ThinkingDisabled} + for _, p := range valid { + if !p.Valid() { + t.Fatalf("ThinkingPolicy(%q).Valid() = false, want true", p) + } + } + invalid := []ThinkingPolicy{"bogus", "ENABLED", "off"} + for _, p := range invalid { + if p.Valid() { + t.Fatalf("ThinkingPolicy(%q).Valid() = true, want false", p) + } + } +} + +func TestLookup_TrimsWhitespace(t *testing.T) { + n := New([]Rule{{Alias: "kimi-k2.6", Target: "kimicode/kimi-for-coding"}}) + r, ok := n.Lookup(" kimi-k2.6 ") + if !ok { + t.Fatal("Lookup() not found with whitespace") + } + if r.Alias != "kimi-k2.6" { + t.Fatalf("Alias = %q, want %q", r.Alias, "kimi-k2.6") + } +} diff --git a/internal/server/http.go b/internal/server/http.go index e73aa32d..7e0727eb 100644 --- a/internal/server/http.go +++ b/internal/server/http.go @@ -26,6 +26,7 @@ import ( "github.com/enterpilot/gomodel/internal/core" "github.com/enterpilot/gomodel/internal/filestore" "github.com/enterpilot/gomodel/internal/mcpgateway" + "github.com/enterpilot/gomodel/internal/modelnormalizer" "github.com/enterpilot/gomodel/internal/responsecache" "github.com/enterpilot/gomodel/internal/responsestore" "github.com/enterpilot/gomodel/internal/session" @@ -85,6 +86,7 @@ type Config struct { TranslatedRequestPatcher TranslatedRequestPatcher // Optional: request patcher for translated routes after workflow resolution BatchRequestPreparer BatchRequestPreparer // Optional: batch request preparer before native provider submission ExposedModelLister ExposedModelLister // Optional: additional public models to merge into GET /v1/models + ModelNormalizer *modelnormalizer.Normalizer // Optional: rewrites chat model aliases + injects thinking policy before dispatch KeepOnlyAliasesAtModelsEndpoint bool // Whether GET /v1/models should hide concrete provider models PassthroughSemanticEnrichers []core.PassthroughSemanticEnricher // Optional: provider-owned passthrough semantic enrichers before workflow resolution BatchStore batchstore.Store // Optional: Batch lifecycle persistence store @@ -193,6 +195,20 @@ func New(provider core.RoutableProvider, cfg *Config) *Server { handler.storageProbe = cfg.StorageProbe handler.cacheProbe = cfg.CacheProbe } + // Synthesize /v1/models entries from normalizer rules when no lister is + // configured. When a lister is already set, layer the normalizer on top + // of it so canonical aliases are always advertised. + if cfg != nil && cfg.ModelNormalizer != nil { + if handler.exposedModelLister == nil { + handler.exposedModelLister = cfg.ModelNormalizer + } else { + primary := handler.exposedModelLister + handler.exposedModelLister = modelnormalizer.ChainedExposedModelLister{ + Primary: primary.ExposedModels, + Secondary: cfg.ModelNormalizer.ExposedModels, + } + } + } if cfg != nil && cfg.EnabledPassthroughProviders != nil { handler.setEnabledPassthroughProviders(cfg.EnabledPassthroughProviders) } @@ -375,6 +391,13 @@ func New(provider core.RoutableProvider, cfg *Config) *Server { e.Use(RequestRewriteMiddleware(cfg.RequestRewriters, auditLogger)) } + // Model normalization runs before workflow resolution so the rewritten + // target model is what resolution, failover, budgets, and caching operate + // on. A nil normalizer skips the middleware entirely. + if cfg != nil && cfg.ModelNormalizer != nil { + e.Use(ModelNormalizerMiddleware(cfg.ModelNormalizer, auditLogger)) + } + // Workflow resolution resolves the request-scoped workflow after auth so // managed auth key user-path overrides are visible to policy resolution while // still keeping workflow resolution failures loggable through the audit middleware. diff --git a/internal/server/model_normalizer_middleware.go b/internal/server/model_normalizer_middleware.go new file mode 100644 index 00000000..8c813500 --- /dev/null +++ b/internal/server/model_normalizer_middleware.go @@ -0,0 +1,93 @@ +package server + +import ( + "bytes" + "io" + "net/http" + + "github.com/goccy/go-json" + "github.com/labstack/echo/v5" + + "github.com/enterpilot/gomodel/internal/auditlog" + "github.com/enterpilot/gomodel/internal/core" + "github.com/enterpilot/gomodel/internal/modelnormalizer" +) + +// ModelNormalizerMiddleware rewrites the chat request body's `model` field +// and injects the per-alias thinking policy before workflow resolution sees +// the request. It runs after authentication (so rewriters only see +// authenticated traffic) and before WorkflowResolutionWithResolverAndPolicy +// (so the rewritten model is the one resolution operates on — otherwise +// ApplyResolvedSelector would re-stamp the original alias over the rewrite). +// +// The middleware is fail-closed: a rewrite error from the normalizer aborts +// the request with HTTP 400. The unchanged body falls through unchanged, +// keeping the seam invisible to clients when no rule matches. +func ModelNormalizerMiddleware(normalizer *modelnormalizer.Normalizer, auditLogger auditlog.LoggerInterface) echo.MiddlewareFunc { + return func(next echo.HandlerFunc) echo.HandlerFunc { + return func(c *echo.Context) error { + if normalizer == nil { + return next(c) + } + if c.Request().Method != http.MethodPost { + return next(c) + } + if c.Request().URL.Path != "/v1/chat/completions" { + return next(c) + } + + body, err := requestBodyBytes(c) + if err != nil { + return handleError(c, core.NewInvalidRequestError("failed to read request body", err)) + } + if len(body) == 0 { + return next(c) + } + + var req core.ChatRequest + if err := json.Unmarshal(body, &req); err != nil { + // Malformed body — let downstream handlers report it. + return next(c) + } + adapted, rewritten, err := normalizer.AdaptChatRequest(&req) + if err != nil { + return handleError(c, core.NewInvalidRequestError("model normalizer: "+err.Error(), err)) + } + if !rewritten { + return next(c) + } + + out, err := json.Marshal(adapted) + if err != nil { + return handleError(c, core.NewInvalidRequestError("model normalizer: marshal: "+err.Error(), err)) + } + pinNormalizerOriginalAuditBody(c, auditLogger) + req2 := c.Request() + req2.Body = io.NopCloser(bytes.NewReader(out)) + req2.ContentLength = int64(len(out)) + storeRequestBodySnapshot(c, out) + if auditLogger != nil && auditLogger.Config().Enabled { + auditlog.EnrichEntryWithRequestRevision(c, auditlog.RequestRevisionSnapshot{ + Rewriter: "model_normalizer", + BytesBefore: len(body), + BytesAfter: len(out), + }) + } + return next(c) + } + } +} + +// pinNormalizerOriginalAuditBody captures the pre-rewrite request body so +// the audit entry shows what the client actually sent (the canonical alias), +// not the rewritten upstream ID. +func pinNormalizerOriginalAuditBody(c *echo.Context, auditLogger auditlog.LoggerInterface) { + if auditLogger == nil || !auditLogger.Config().Enabled { + return + } + entry, ok := c.Get(string(auditlog.LogEntryKey)).(*auditlog.LogEntry) + if !ok || entry == nil { + return + } + auditlog.PopulateRequestData(entry, c.Request(), auditLogger.Config()) +} \ No newline at end of file diff --git a/internal/server/modelnormalizer_server_test.go b/internal/server/modelnormalizer_server_test.go new file mode 100644 index 00000000..1afc2a49 --- /dev/null +++ b/internal/server/modelnormalizer_server_test.go @@ -0,0 +1,171 @@ +package server + +import ( + "context" + "encoding/json" + "io" + "net/http" + "net/http/httptest" + "strings" + "testing" + + "github.com/stretchr/testify/require" + + "github.com/enterpilot/gomodel/internal/core" + "github.com/enterpilot/gomodel/internal/modelnormalizer" +) + +// chatCaptureProvider records the chat request it receives and returns a +// minimal response so the test can verify what the provider actually saw. +type chatCaptureProvider struct { + lastChat *core.ChatRequest +} + +func (p *chatCaptureProvider) ChatCompletion(_ context.Context, req *core.ChatRequest) (*core.ChatResponse, error) { + p.lastChat = req + return &core.ChatResponse{ID: "chatcmpl-1", Model: req.Model, Provider: "kimicode"}, nil +} +func (p *chatCaptureProvider) StreamChatCompletion(context.Context, *core.ChatRequest) (io.ReadCloser, error) { + return nil, nil +} +func (p *chatCaptureProvider) ListModels(context.Context) (*core.ModelsResponse, error) { + return &core.ModelsResponse{Object: "list", Data: []core.Model{ + {ID: "kimi-for-coding", Object: "model", OwnedBy: "kimicode"}, + {ID: "k3", Object: "model", OwnedBy: "kimicode"}, + }}, nil +} +func (p *chatCaptureProvider) Responses(context.Context, *core.ResponsesRequest) (*core.ResponsesResponse, error) { + return nil, nil +} +func (p *chatCaptureProvider) StreamResponses(context.Context, *core.ResponsesRequest) (io.ReadCloser, error) { + return nil, nil +} +func (p *chatCaptureProvider) Embeddings(context.Context, *core.EmbeddingRequest) (*core.EmbeddingResponse, error) { + return nil, nil +} +func (p *chatCaptureProvider) Supports(string) bool { return true } +func (p *chatCaptureProvider) GetProviderType(string) string { return "kimicode" } + +func TestHandler_ChatCompletionAppliesNormalizer(t *testing.T) { + provider := &chatCaptureProvider{} + normalizer := modelnormalizer.New([]modelnormalizer.Rule{ + {Alias: "kimi-k2.6", Target: "kimicode/kimi-for-coding", Thinking: modelnormalizer.ThinkingDisabled}, + }) + srv := New(provider, &Config{ModelNormalizer: normalizer}) + + body := `{"model":"kimi-k2.6","messages":[{"role":"user","content":"hi"}]}` + req := httptest.NewRequest(http.MethodPost, "/v1/chat/completions", strings.NewReader(body)) + req.Header.Set("Content-Type", "application/json") + rec := httptest.NewRecorder() + srv.echo.ServeHTTP(rec, req) + + require.Equal(t, http.StatusOK, rec.Code, "body: %s", rec.Body.String()) + + require.NotNil(t, provider.lastChat) + require.Equal(t, "kimi-for-coding", provider.lastChat.Model) + require.Equal(t, "kimicode", provider.lastChat.Provider) + thinking := provider.lastChat.ExtraFields.Lookup("thinking") + require.NotNil(t, thinking, "provider should receive thinking extension") + require.Contains(t, string(thinking), `"disabled"`) +} + +func TestHandler_ChatCompletionUnknownAliasPassesThrough(t *testing.T) { + provider := &chatCaptureProvider{} + normalizer := modelnormalizer.New([]modelnormalizer.Rule{ + {Alias: "kimi-k2.6", Target: "kimicode/kimi-for-coding"}, + }) + srv := New(provider, &Config{ModelNormalizer: normalizer}) + + body := `{"model":"kimicode/kimi-for-coding","messages":[{"role":"user","content":"hi"}]}` + req := httptest.NewRequest(http.MethodPost, "/v1/chat/completions", strings.NewReader(body)) + req.Header.Set("Content-Type", "application/json") + rec := httptest.NewRecorder() + srv.echo.ServeHTTP(rec, req) + + require.Equal(t, http.StatusOK, rec.Code, "body: %s", rec.Body.String()) + + require.NotNil(t, provider.lastChat) + require.Equal(t, "kimi-for-coding", provider.lastChat.Model) + require.Equal(t, "kimicode", provider.lastChat.Provider) + require.Nil(t, provider.lastChat.ExtraFields.Lookup("thinking")) +} + +func TestNew_WiresModelNormalizerIntoChatPath(t *testing.T) { + provider := &chatCaptureProvider{} + normalizer := modelnormalizer.New([]modelnormalizer.Rule{ + {Alias: "kimi-k2.6", Target: "kimicode/kimi-for-coding", Thinking: modelnormalizer.ThinkingDisabled}, + }) + + srv := New(provider, &Config{ModelNormalizer: normalizer}) + + // Send a chat completion request for the canonical alias; the provider + // should see the rewritten target model + injected thinking field. + body := `{"model":"kimi-k2.6","messages":[{"role":"user","content":"hi"}]}` + req := httptest.NewRequest(http.MethodPost, "/v1/chat/completions", strings.NewReader(body)) + req.Header.Set("Content-Type", "application/json") + rec := httptest.NewRecorder() + srv.echo.ServeHTTP(rec, req) + require.Equal(t, http.StatusOK, rec.Code, "body: %s", rec.Body.String()) + + require.NotNil(t, provider.lastChat) + require.Equal(t, "kimi-for-coding", provider.lastChat.Model) + require.Equal(t, "kimicode", provider.lastChat.Provider) + thinking := provider.lastChat.ExtraFields.Lookup("thinking") + require.NotNil(t, thinking, "provider should receive thinking extension") + require.Contains(t, string(thinking), `"disabled"`) +} + +func TestNew_ListModelsMergesNormalizerAliases(t *testing.T) { + provider := &chatCaptureProvider{} + cw := 262144 + normalizer := modelnormalizer.New([]modelnormalizer.Rule{ + {Alias: "kimi-k2.6", Target: "kimicode/kimi-for-coding", Modes: []string{"chat"}, ContextWindow: &cw}, + {Alias: "bge_m3_embed", Target: "kimicode/bge_m3_embed", Modes: []string{"embedding"}}, + }) + + srv := New(provider, &Config{ModelNormalizer: normalizer}) + + req := httptest.NewRequest(http.MethodGet, "/v1/models", nil) + rec := httptest.NewRecorder() + srv.echo.ServeHTTP(rec, req) + require.Equal(t, http.StatusOK, rec.Code) + + var resp core.ModelsResponse + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &resp)) + + byID := make(map[string]core.Model, len(resp.Data)) + for _, m := range resp.Data { + byID[m.ID] = m + } + + // Provider models are still listed. + require.Contains(t, byID, "kimi-for-coding") + require.Contains(t, byID, "k3") + + // Canonical aliases are synthesized with metadata. + require.Contains(t, byID, "kimi-k2.6") + require.Equal(t, "model", byID["kimi-k2.6"].Object) + require.NotNil(t, byID["kimi-k2.6"].Metadata) + require.Equal(t, []string{"chat"}, byID["kimi-k2.6"].Metadata.Modes) + require.NotNil(t, byID["kimi-k2.6"].Metadata.ContextWindow) + require.Equal(t, 262144, *byID["kimi-k2.6"].Metadata.ContextWindow) + + require.Contains(t, byID, "bge_m3_embed") + require.NotNil(t, byID["bge_m3_embed"].Metadata) + require.Equal(t, []string{"embedding"}, byID["bge_m3_embed"].Metadata.Modes) + require.Contains(t, byID["bge_m3_embed"].Metadata.Categories, core.CategoryEmbedding) +} + +func TestNew_ListModelsWithoutNormalizerLeavesHandlerUnchanged(t *testing.T) { + provider := &chatCaptureProvider{} + srv := New(provider, &Config{}) + + req := httptest.NewRequest(http.MethodGet, "/v1/models", nil) + rec := httptest.NewRecorder() + srv.echo.ServeHTTP(rec, req) + require.Equal(t, http.StatusOK, rec.Code) + + var resp core.ModelsResponse + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &resp)) + require.Len(t, resp.Data, 2) // only the provider's two models, no aliases +} \ No newline at end of file