diff --git a/README.md b/README.md index b62a393..6cea8c4 100644 --- a/README.md +++ b/README.md @@ -101,6 +101,13 @@ plugins: protocol: "messages" # "chat-completions" | "messages" | "responses" endpoint: "/v1/messages" # must start with / + # Optional per-model capability fixes (takes priority over catalog and models.dev) + model-metadata-overrides: + "gpt-5.6-luna": + thinking: + zero-allowed: true + levels: [none, low, medium, high, xhigh, max] + # Execution settings request-timeout: "5m" # upstream request timeout (default: "5m") max-response-bytes: 67108864 # max non-streaming response body size in bytes (default: 64 MiB) @@ -109,6 +116,8 @@ plugins: ### Configuration Options +The configured `catalog-url` alone determines which models are routable. On each successful discovery, the plugin fills missing limits, modalities, and unambiguous reasoning metadata from the `opencode-go` provider in [models.dev](https://models.dev/); if that request fails, the last good metadata is retained. Field priority is: `model-metadata-overrides` > catalog > models.dev > plugin defaults. `toggle` reasoning options are not inferred as effort levels. + | Option | Type | Default | Description | |---|---|---|---| | `api-keys` | `[]object` | *(Required)* | List of API keys (`- value: "..."`). Supports `${ENV_VAR}` expansion. Duplicates and empty values are rejected. | @@ -122,6 +131,7 @@ plugins: | `protocols.messages` | `bool` | `true` | Protocol switch for Messages endpoints. | | `protocols.responses` | `bool` | `true` | Protocol switch for Responses endpoints. | | `route-overrides` | `map` | `{}` | Map of model ID to `{ protocol: "...", endpoint: "..." }` overriding built-in family routing. Valid protocols: `chat-completions`, `messages`, `responses`. | +| `model-metadata-overrides` | `map` | `{}` | Per-model `thinking` (`min`, `max`, `zero-allowed`, `dynamic-allowed`, `levels`), `context-limit`, `output-limit`, `input-modes`, and `output-modes`. Each supplied field overrides the catalog and models.dev; unspecified fields retain their current value. | | `request-timeout` | `duration` | `5m` | Upstream HTTP request timeout. Must be positive. | | `max-response-bytes` | `int64` | `67108864` (64 MiB) | Maximum non-streaming response body size in bytes. | | `allow-http` | `bool` | `false` | When `true`, permits `http://` scheme in `base-url` / `catalog-url` for local testing. | diff --git a/internal/adapter/chatcompletions/request.go b/internal/adapter/chatcompletions/request.go index b084d3e..cec6886 100644 --- a/internal/adapter/chatcompletions/request.go +++ b/internal/adapter/chatcompletions/request.go @@ -38,7 +38,7 @@ func AuthHeaders(key string) http.Header { func BuildRequest(upstreamModel, sourceFormat string, sourceBody []byte, ts *pluginapi.ThinkingSupport) ([]byte, *errclass.Error) { switch sourceFormat { case "openai": - return buildOpenAIRequest(upstreamModel, sourceBody) + return buildOpenAIRequest(upstreamModel, sourceBody, ts) case "claude": return claudeToChat(upstreamModel, sourceBody, ts) case "openai-response": @@ -52,7 +52,7 @@ func BuildRequest(upstreamModel, sourceFormat string, sourceBody []byte, ts *plu // body to upstreamModel, normalizes role:"developer" messages to role:"system", // and strips any malformed top-level thinking object. DeepSeek models fail if // thinking lacks a valid string type field or if messages contain role:"developer". -func buildOpenAIRequest(upstreamModel string, body []byte) ([]byte, *errclass.Error) { +func buildOpenAIRequest(upstreamModel string, body []byte, ts *pluginapi.ThinkingSupport) ([]byte, *errclass.Error) { var req map[string]json.RawMessage if err := json.Unmarshal(body, &req); err != nil { return nil, errclass.Translation("malformed openai request JSON: " + err.Error()) @@ -60,6 +60,17 @@ func buildOpenAIRequest(upstreamModel string, body []byte) ([]byte, *errclass.Er if req == nil { return nil, errclass.Translation("malformed request body: JSON null is not a valid request") } + if raw := req["reasoning_effort"]; len(raw) > 0 && string(raw) != "null" { + var effort string + if json.Unmarshal(raw, &effort) != nil { + return nil, errclass.Translation("reasoning_effort must be a string") + } + if effort != "" { + if eErr := thinking.ValidateEffort(effort, ts); eErr != nil { + return nil, eErr + } + } + } rawThinking, hasThinking := req["thinking"] validThinking := hasThinking && isValidThinking(rawThinking) if hasThinking && !validThinking { diff --git a/internal/adapter/messages/request.go b/internal/adapter/messages/request.go index c35d358..8a2f0cc 100644 --- a/internal/adapter/messages/request.go +++ b/internal/adapter/messages/request.go @@ -56,7 +56,33 @@ type messagesRequest struct { func BuildRequest(upstreamModel string, sourceFormat string, sourceBody []byte, ts *pluginapi.ThinkingSupport) ([]byte, *errclass.Error) { switch sourceFormat { case "claude": - return shared.RewriteModelID(upstreamModel, sourceBody, "claude") + body, eErr := shared.RewriteModelID(upstreamModel, sourceBody, "claude") + if eErr != nil { + return nil, eErr + } + var native struct { + Thinking *shared.ClaudeThinking `json:"thinking"` + } + if json.Unmarshal(body, &native) != nil { + return nil, errclass.Translation("invalid thinking control") + } + if native.Thinking != nil { + switch native.Thinking.Type { + case "disabled": + if eErr := thinking.ValidateEffort("none", ts); eErr != nil { + return nil, eErr + } + case "enabled": + budget := native.Thinking.BudgetTokens + supported := thinking.SupportedLevels(ts) + if budget <= 0 || ts != nil && (ts.Min > 0 && budget < int64(ts.Min) || + ts.Max > 0 && budget > int64(ts.Max) || + len(supported) == 1 && supported[0] == "none") { + return nil, &errclass.Error{Class: errclass.ClassUnsupported, Message: "thinking budget_tokens is not supported for this model"} + } + } + } + return body, nil case "openai": return fromChatCompletions(upstreamModel, sourceBody, ts) case "openai-response": diff --git a/internal/adapter/messages/request_test.go b/internal/adapter/messages/request_test.go index 34da5b8..4885b81 100644 --- a/internal/adapter/messages/request_test.go +++ b/internal/adapter/messages/request_test.go @@ -68,6 +68,38 @@ func TestBuildRequestClaudeMalformed(t *testing.T) { } } +func TestNativeClaudeThinkingCapability(t *testing.T) { + bounded := &pluginapi.ThinkingSupport{Min: 1024, Max: 4096, Levels: []string{"low", "medium"}} + cases := []struct { + name string + thinking string + ts *pluginapi.ThinkingSupport + wantErr errclass.Class + }{ + {"within bounds", `{"type":"enabled","budget_tokens":2048}`, bounded, ""}, + {"below minimum", `{"type":"enabled","budget_tokens":512}`, bounded, errclass.ClassUnsupported}, + {"above maximum", `{"type":"enabled","budget_tokens":8192}`, bounded, errclass.ClassUnsupported}, + {"zero budget", `{"type":"enabled","budget_tokens":0}`, bounded, errclass.ClassUnsupported}, + {"missing budget", `{"type":"enabled"}`, bounded, errclass.ClassUnsupported}, + {"no enabled level", `{"type":"enabled","budget_tokens":2048}`, &pluginapi.ThinkingSupport{Levels: []string{"none"}}, errclass.ClassUnsupported}, + {"unknown capability fallback", `{"type":"enabled","budget_tokens":2048}`, nil, ""}, + {"malformed budget", `{"type":"enabled","budget_tokens":"bad"}`, bounded, errclass.ClassTranslation}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + body := []byte(`{"model":"m","messages":[],"max_tokens":10000,"thinking":` + tc.thinking + `}`) + out, eErr := BuildRequest("m", "claude", body, tc.ts) + if tc.wantErr == "" { + if eErr != nil || string(out) != string(body) { + t.Fatalf("native passthrough changed: %s, %v", out, eErr) + } + } else if eErr == nil || eErr.Class != tc.wantErr { + t.Fatalf("error = %v, want %s", eErr, tc.wantErr) + } + }) + } +} + func chatReq(t *testing.T, body string) (map[string]any, *errclass.Error) { return chatReqTS(t, nil, body) } diff --git a/internal/adapter/responses/request.go b/internal/adapter/responses/request.go index c5843d8..e6f0187 100644 --- a/internal/adapter/responses/request.go +++ b/internal/adapter/responses/request.go @@ -12,7 +12,6 @@ package responses import ( "encoding/json" "fmt" - "slices" "strings" "github.com/router-for-me/CLIProxyAPI/v7/sdk/pluginapi" @@ -38,7 +37,40 @@ var EndpointPath = catalog.RouteResponses.EndpointPath() func BuildRequest(upstreamModel string, sourceFormat string, sourceBody []byte, ts *pluginapi.ThinkingSupport) ([]byte, *errclass.Error) { switch sourceFormat { case "openai-response": - return shared.RewriteModelID(upstreamModel, sourceBody, "openai-response") + body, eErr := shared.RewriteModelID(upstreamModel, sourceBody, "openai-response") + if eErr != nil { + return nil, eErr + } + var req map[string]json.RawMessage + _ = json.Unmarshal(body, &req) + var reasoning map[string]json.RawMessage + if len(req["reasoning"]) == 0 || string(req["reasoning"]) == "null" { + return body, nil + } + if json.Unmarshal(req["reasoning"], &reasoning) != nil || reasoning == nil { + return nil, errclass.Translation("reasoning must be an object") + } + if raw := reasoning["effort"]; len(raw) > 0 && string(raw) != "null" { + var effort string + if json.Unmarshal(raw, &effort) != nil { + return nil, errclass.Translation("reasoning.effort must be a string") + } + if effort != "" { + if eErr := thinking.ValidateEffort(effort, ts); eErr != nil { + return nil, eErr + } + } + if strings.EqualFold(strings.TrimSpace(effort), "auto") { + delete(reasoning, "effort") + if len(reasoning) == 0 { + delete(req, "reasoning") + } else { + req["reasoning"], _ = json.Marshal(reasoning) + } + body, _ = json.Marshal(req) + } + } + return body, nil case "openai": return fromChatCompletions(upstreamModel, sourceBody, ts) case "claude": @@ -276,10 +308,7 @@ func reasoningEffortFor(effort string, ts *pluginapi.ThinkingSupport) (string, b case effort == "auto": return "", false case effort == "none": - if slices.Contains(thinking.SupportedLevels(ts), "none") { - return effort, true - } - return "", false + return effort, true // ValidateEffort admits it only when the model supports off. default: return effort, true } diff --git a/internal/adapter/responses/request_test.go b/internal/adapter/responses/request_test.go index dbea2fc..4d3b7ac 100644 --- a/internal/adapter/responses/request_test.go +++ b/internal/adapter/responses/request_test.go @@ -683,16 +683,14 @@ func TestFromChatCompletionsEffortCapability(t *testing.T) { } } -// Sentinels are capability-gated: a "none" the model does not declare and -// the dynamic "auto" sentinel are omitted — Responses has no off-switch, so -// omission is the no-forced-reasoning policy (matches the Messages-target -// leg); declared levels forward. -func TestFromChatCompletionsEffortSentinelsOmitted(t *testing.T) { +// ZeroAllowed admits an off state even when Levels omits "none"; auto has +// no Responses wire value and is omitted. +func TestFromChatCompletionsEffortSentinels(t *testing.T) { ts := &pluginapi.ThinkingSupport{ZeroAllowed: true, DynamicAllowed: true} m := decodeReq(t, mustBuild(t, "m", "openai", []byte(`{"messages":[],"reasoning_effort":"none"}`), ts)) - if _, has := m["reasoning"]; has { - t.Fatalf("none must omit reasoning: %v", m["reasoning"]) + if r := m["reasoning"].(map[string]any); r["effort"] != "none" { + t.Fatalf("none must forward: %v", r) } m = decodeReq(t, mustBuild(t, "m", "openai", diff --git a/internal/adapter/responses_parity_test.go b/internal/adapter/responses_parity_test.go index d3b1fd6..cbd89e5 100644 --- a/internal/adapter/responses_parity_test.go +++ b/internal/adapter/responses_parity_test.go @@ -8,12 +8,94 @@ import ( "strings" "testing" + "github.com/router-for-me/CLIProxyAPI/v7/sdk/pluginapi" + "opencode-go-cliproxyapi/internal/adapter/chatcompletions" "opencode-go-cliproxyapi/internal/adapter/messages" "opencode-go-cliproxyapi/internal/adapter/responses" "opencode-go-cliproxyapi/internal/errclass" ) +func TestReasoningEffortNativeAndTranslatedParity(t *testing.T) { + ts := &pluginapi.ThinkingSupport{Levels: []string{"minimal", "low", "medium", "high", "xhigh", "max"}, ZeroAllowed: true, DynamicAllowed: true} + for _, effort := range []string{"none", "minimal", "low", "medium", "high", "xhigh", "max", "auto"} { + t.Run(effort, func(t *testing.T) { + cc := []byte(`{"model":"m","messages":[],"reasoning_effort":"` + effort + `"}`) + resp := []byte(`{"model":"m","input":[],"reasoning":{"effort":"` + effort + `"}}`) + for _, tc := range []struct { + name string + build func(string, string, []byte, *pluginapi.ThinkingSupport) ([]byte, *errclass.Error) + format string + body []byte + field string + }{ + {"native chat", chatcompletions.BuildRequest, "openai", cc, "reasoning_effort"}, + {"translated chat", chatcompletions.BuildRequest, "openai-response", resp, "reasoning_effort"}, + {"native responses", responses.BuildRequest, "openai-response", resp, "reasoning"}, + {"translated responses", responses.BuildRequest, "openai", cc, "reasoning"}, + } { + out, eErr := tc.build("m", tc.format, tc.body, ts) + if eErr != nil { + t.Fatalf("%s: %v", tc.name, eErr) + } + var wire map[string]any + if err := json.Unmarshal(out, &wire); err != nil { + t.Fatal(err) + } + if tc.field == "reasoning_effort" && wire[tc.field] != effort { + t.Errorf("%s: %s = %v", tc.name, tc.field, wire[tc.field]) + } + if tc.field == "reasoning" { + if effort == "auto" && wire[tc.field] != nil { + t.Errorf("%s: auto must omit Responses reasoning: %v", tc.name, wire[tc.field]) + } else if effort != "auto" { + got, ok := wire[tc.field].(map[string]any) + if !ok || got["effort"] != effort { + t.Errorf("%s: reasoning = %v", tc.name, wire[tc.field]) + } + } + } + } + for _, source := range []struct { + format string + body []byte + }{{"openai", cc}, {"openai-response", resp}} { + out, eErr := messages.BuildRequest("m", source.format, source.body, ts) + if eErr != nil { + t.Fatalf("Messages from %s: %v", source.format, eErr) + } + var wire map[string]any + if err := json.Unmarshal(out, &wire); err != nil { + t.Fatal(err) + } + if (wire["thinking"] != nil) != (effort != "none" && effort != "auto") { + t.Errorf("Messages from %s: thinking = %v", source.format, wire["thinking"]) + } + } + }) + } + limited := &pluginapi.ThinkingSupport{Levels: []string{"high"}} + for _, body := range []struct { + format string + data []byte + build func(string, string, []byte, *pluginapi.ThinkingSupport) ([]byte, *errclass.Error) + }{ + {"openai", []byte(`{"model":"m","reasoning_effort":"max"}`), chatcompletions.BuildRequest}, + {"openai-response", []byte(`{"model":"m","reasoning":{"effort":"max"}}`), responses.BuildRequest}, + } { + if _, eErr := body.build("m", body.format, body.data, limited); eErr == nil || eErr.Class != errclass.ClassUnsupported { + t.Errorf("native %s accepted unsupported max: %v", body.format, eErr) + } + } + claudeOff := []byte(`{"model":"m","max_tokens":1024,"messages":[],"thinking":{"type":"disabled"}}`) + if _, eErr := messages.BuildRequest("m", "claude", claudeOff, limited); eErr == nil || eErr.Class != errclass.ClassUnsupported { + t.Errorf("native Messages accepted unsupported reasoning off: %v", eErr) + } + if _, eErr := messages.BuildRequest("m", "claude", claudeOff, ts); eErr != nil { + t.Errorf("native Messages rejected supported reasoning off: %v", eErr) + } +} + type feeder interface { Feed(chunk []byte) (events [][]byte, done bool, eErr *errclass.Error) } diff --git a/internal/catalog/catalog.go b/internal/catalog/catalog.go index 3118a98..3eea27e 100644 --- a/internal/catalog/catalog.go +++ b/internal/catalog/catalog.go @@ -17,6 +17,7 @@ import ( "net/url" "strings" "sync" + "time" "github.com/router-for-me/CLIProxyAPI/v7/sdk/pluginapi" @@ -28,6 +29,7 @@ import ( // completion bodies, so the fetch budget never drops below 16 MiB while // max-response-bytes stays meaningful for completions elsewhere. const catalogBudgetFloor = 16 << 20 +const modelsDevURL = "https://models.dev/api.json" // Route is the upstream protocol a model is served through (FR-004). type Route string @@ -114,6 +116,23 @@ type rawThinking struct { Levels []string `json:"levels"` } +type modelsDevModel struct { + Limit struct { + Context int64 `json:"context"` + Output int64 `json:"output"` + } `json:"limit"` + Modalities struct { + Input []string `json:"input"` + Output []string `json:"output"` + } `json:"modalities"` + ReasoningOptions []struct { + Type string `json:"type"` + Values []string `json:"values"` + Min *int `json:"min"` + Max *int `json:"max"` + } `json:"reasoning_options"` +} + // prefixRoutes covers unknown variants of known families; // longest prefix wins (list is ordered longest-first). var prefixRoutes = []struct { @@ -144,8 +163,9 @@ type Manager struct { // so SeedFrom can rebuild records against a NEW config (route overrides, // protocol flags, and prefix settings may all have changed since the // snapshot was built). - raw []rawModel - models []ModelRecord + raw []rawModel + metadata map[string]modelsDevModel + models []ModelRecord // index maps both PublicID and UpstreamID to their record so ID // resolution is O(1); rebuilt atomically with models on every swap. index map[string]ModelRecord @@ -175,8 +195,9 @@ func New(cfg config.Config, client HostClient) *Manager { func (m *Manager) SeedFrom(prev *Manager) { prev.mu.Lock() raw := prev.raw + metadata := prev.metadata prev.mu.Unlock() - m.swap(raw) + m.swap(raw, metadata) } // Refresh fetches and swaps the catalog snapshot. On failure it returns a @@ -225,10 +246,38 @@ func (m *Manager) Refresh(ctx context.Context, apiKey string) error { } else if err := json.Unmarshal(env.Data, &entries); err != nil { return m.fail("invalid json") } - m.swap(entries, warns...) + var metadata map[string]modelsDevModel + if len(entries) > 0 { + metadata = m.fetchMetadata(ctx) + } + m.swap(entries, metadata, warns...) return nil } +func (m *Manager) fetchMetadata(ctx context.Context) map[string]modelsDevModel { + // ponytail: fetch the full API until models.dev exposes a provider-only endpoint. + ctx, cancel := context.WithTimeout(ctx, 5*time.Second) + defer cancel() + resp, err := m.client.Do(ctx, pluginapi.HTTPRequest{ + Method: http.MethodGet, URL: modelsDevURL, + Headers: http.Header{"Accept": []string{"application/json"}}, + }) + if err != nil || resp.StatusCode != http.StatusOK || len(resp.Body) > catalogBudgetFloor { + return nil + } + var root map[string]json.RawMessage + if json.Unmarshal(resp.Body, &root) != nil { + return nil + } + var provider struct { + Models map[string]modelsDevModel `json:"models"` + } + if json.Unmarshal(root["opencode-go"], &provider) != nil || len(provider.Models) == 0 { + return nil + } + return provider.Models +} + // fail applies the stale policy and returns the classified error. func (m *Manager) fail(category string) error { m.mu.Lock() @@ -237,7 +286,7 @@ func (m *Manager) fail(category string) error { // Clear the index too — Lookup must stop resolving IDs whose // records are gone (FR-002). Raw entries go with them: a cleared // snapshot has nothing to seed. - m.raw, m.models, m.index, m.unsup, m.warns = nil, nil, nil, nil, nil + m.raw, m.metadata, m.models, m.index, m.unsup, m.warns = nil, nil, nil, nil, nil, nil } return fmt.Errorf("catalog refresh failed: %s", category) } @@ -246,7 +295,12 @@ func (m *Manager) fail(category string) error { // atomically under the mutex (FR-010 dedup, arch §5 route priority). // extraWarns are caller-supplied snapshot diagnostics (e.g. decode-level // shape-drift notices) recorded alongside the per-entry ones. -func (m *Manager) swap(entries []rawModel, extraWarns ...string) { +func (m *Manager) swap(entries []rawModel, metadata map[string]modelsDevModel, extraWarns ...string) { + m.mu.Lock() + defer m.mu.Unlock() + if metadata == nil { + metadata = m.metadata // select fallback and publish the snapshot under one lock. + } models := make([]ModelRecord, 0, len(entries)) index := make(map[string]ModelRecord, len(entries)*2) var unsup []UnsupportedModel @@ -315,6 +369,7 @@ func (m *Manager) swap(entries []rawModel, extraWarns ...string) { OutputModes: outputModes, Thinking: normalizeThinking(e.Thinking), } + applyMetadata(&rec, e, metadata[e.ID], m.cfg.ModelMetadataOverrides[e.ID]) // With a prefix enabled, one record's PublicID can equal another // record's UpstreamID (upstream "foo" and "opencode-go/foo" both // claim index key "/foo"); last-write-wins would silently @@ -339,9 +394,91 @@ func (m *Manager) swap(entries []rawModel, extraWarns ...string) { index[rec.PublicID] = rec index[rec.UpstreamID] = rec } - m.mu.Lock() - defer m.mu.Unlock() - m.raw, m.models, m.index, m.unsup, m.warns = entries, models, index, unsup, warns + m.raw, m.metadata, m.models, m.index, m.unsup, m.warns = entries, metadata, models, index, unsup, warns +} + +func applyMetadata(rec *ModelRecord, catalog rawModel, dev modelsDevModel, override config.ModelMetadataOverride) { + if rec.ContextLimit == 0 { + rec.ContextLimit = dev.Limit.Context + } + if rec.OutputLimit == 0 { + rec.OutputLimit = dev.Limit.Output + } + if len(rec.InputModes) == 0 { + rec.InputModes = dev.Modalities.Input + } + if len(rec.OutputModes) == 0 { + rec.OutputModes = dev.Modalities.Output + } + var derived *rawThinking + for _, option := range dev.ReasoningOptions { + switch option.Type { + case "effort": + if derived == nil { + derived = &rawThinking{} + } + derived.Levels = append(derived.Levels, option.Values...) + for _, level := range option.Values { + if strings.EqualFold(level, "none") { + b := true + derived.ZeroAllowed = &b + } + } + case "budget_tokens": + if derived == nil { + derived = &rawThinking{} + } + derived.Min, derived.Max = option.Min, option.Max + } + } + thinking := mergeThinking(derived, catalog.Thinking) + if o := override.Thinking; o != nil { + thinking = mergeThinking(thinking, &rawThinking{ + Min: o.Min, Max: o.Max, ZeroAllowed: o.ZeroAllowed, + DynamicAllowed: o.DynamicAllowed, Levels: o.Levels, + }) + } + rec.Thinking = normalizeThinking(thinking) + if override.ContextLimit != nil { + rec.ContextLimit = *override.ContextLimit + } + if override.OutputLimit != nil { + rec.OutputLimit = *override.OutputLimit + } + if override.InputModes != nil { + rec.InputModes = override.InputModes + } + if override.OutputModes != nil { + rec.OutputModes = override.OutputModes + } +} + +func mergeThinking(base, layer *rawThinking) *rawThinking { + if layer == nil { + return base + } + if base == nil { + base = &rawThinking{} + } else { + copy := *base + base = © + } + if layer.Min != nil { + base.Min = layer.Min + } + if layer.Max != nil { + base.Max = layer.Max + } + if layer.ZeroAllowed != nil { + base.ZeroAllowed = layer.ZeroAllowed + } + if layer.DynamicAllowed != nil { + base.DynamicAllowed = layer.DynamicAllowed + } + if layer.Levels != nil { + base.Levels = layer.Levels + } + return base } // protocolEnabled reports whether the resolved route's protocol flag is on diff --git a/internal/catalog/catalog_test.go b/internal/catalog/catalog_test.go index c666d1e..5f45dc8 100644 --- a/internal/catalog/catalog_test.go +++ b/internal/catalog/catalog_test.go @@ -19,15 +19,20 @@ const testKey = "sk-test-key-123" // fakeClient is a canned, thread-safe HostClient capturing the last request. type fakeClient struct { - mu sync.Mutex - resp pluginapi.HTTPResponse - err error - gotReq *pluginapi.HTTPRequest + mu sync.Mutex + resp pluginapi.HTTPResponse + err error + gotReq *pluginapi.HTTPRequest + metadataResp pluginapi.HTTPResponse + metadataErr error } func (f *fakeClient) Do(_ context.Context, req pluginapi.HTTPRequest) (pluginapi.HTTPResponse, error) { f.mu.Lock() defer f.mu.Unlock() + if req.URL == modelsDevURL { + return f.metadataResp, f.metadataErr + } f.gotReq = &req return f.resp, f.err } @@ -70,6 +75,59 @@ func findModel(t *testing.T, models []ModelRecord, upstreamID string) ModelRecor return ModelRecord{} } +func TestMetadataPrecedenceAndAvailability(t *testing.T) { + output := int64(900) + zero := false + cfg := testCfg() + cfg.ModelMetadataOverrides = map[string]config.ModelMetadataOverride{ + "gpt-5.6-luna": { + OutputLimit: &output, + Thinking: &config.ThinkingMetadata{ZeroAllowed: &zero, Levels: []string{"max"}}, + }, + } + fc := &fakeClient{ + resp: pluginapi.HTTPResponse{StatusCode: 200, Body: []byte(`{"data":[ + {"id":"gpt-5.6-luna","context_length":42,"thinking":{"levels":["high"]}}, + {"id":"glm-5.2"},{"id":"qwen3.8-max"}]}`)}, + metadataResp: pluginapi.HTTPResponse{StatusCode: 200, Body: []byte(`{"opencode-go":{"models":{ + "gpt-5.6-luna":{"limit":{"context":100,"output":200},"modalities":{"input":["text","image"],"output":["text"]},"reasoning_options":[{"type":"effort","values":["none","low","xhigh"]}]}, + "glm-5.2":{"limit":{"context":300,"output":400},"reasoning_options":[{"type":"effort","values":["high","max"]}]}, + "qwen3.8-max":{"reasoning_options":[{"type":"toggle"},{"type":"budget_tokens","min":1024,"max":32768},{"type":"effort","values":["low","xhigh"]}]}, + "gpt-undiscovered":{"reasoning_options":[{"type":"effort","values":["max"]}]} + }}}`)}, + } + m := New(cfg, fc) + mustRefresh(t, m) + if len(m.Models()) != 3 { + t.Fatalf("models.dev must not add routable models: %+v", m.Models()) + } + gpt := findModel(t, m.Models(), "gpt-5.6-luna") + if gpt.ContextLimit != 42 || gpt.OutputLimit != 900 || !reflect.DeepEqual(gpt.InputModes, []string{"text", "image"}) || + gpt.Thinking == nil || !reflect.DeepEqual(gpt.Thinking.Levels, []string{"max"}) || gpt.Thinking.ZeroAllowed { + t.Fatalf("metadata precedence failed: %+v", gpt) + } + glm := findModel(t, m.Models(), "glm-5.2") + if glm.ContextLimit != 300 || glm.OutputLimit != 400 || glm.Thinking == nil || + !reflect.DeepEqual(glm.Thinking.Levels, []string{"high", "max"}) { + t.Fatalf("models.dev fallback failed: %+v", glm) + } + qwen := findModel(t, m.Models(), "qwen3.8-max") + if qwen.Thinking == nil || qwen.Thinking.Min != 1024 || qwen.Thinking.Max != 32768 || + !reflect.DeepEqual(qwen.Thinking.Levels, []string{"low", "xhigh"}) { + t.Fatalf("reasoning options mapping failed: %+v", qwen) + } + fc.metadataResp = pluginapi.HTTPResponse{StatusCode: 503} + mustRefresh(t, m) + if got := findModel(t, m.Models(), "glm-5.2").Thinking; got == nil || !reflect.DeepEqual(got.Levels, []string{"high", "max"}) { + t.Fatalf("lost last good metadata during outage: %+v", got) + } + seeded := New(testCfg(), &fakeClient{}) + seeded.SeedFrom(m) + if got := findModel(t, seeded.Models(), "gpt-5.6-luna"); got.OutputLimit != 200 || !reflect.DeepEqual(got.Thinking.Levels, []string{"high"}) { + t.Fatalf("seed must reapply new config with cached metadata: %+v", got) + } +} + func TestRouteEndpointPath(t *testing.T) { cases := []struct { route Route diff --git a/internal/config/config.go b/internal/config/config.go index dec6629..31a31e2 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -47,33 +47,51 @@ type RouteOverride struct { Endpoint string `yaml:"endpoint"` } +type ThinkingMetadata struct { + Min *int `yaml:"min"` + Max *int `yaml:"max"` + ZeroAllowed *bool `yaml:"zero-allowed"` + DynamicAllowed *bool `yaml:"dynamic-allowed"` + Levels []string `yaml:"levels"` +} + +type ModelMetadataOverride struct { + ContextLimit *int64 `yaml:"context-limit"` + OutputLimit *int64 `yaml:"output-limit"` + InputModes []string `yaml:"input-modes"` + OutputModes []string `yaml:"output-modes"` + Thinking *ThinkingMetadata `yaml:"thinking"` +} + type Config struct { - BaseURL string - CatalogURL string - ModelPrefix ModelPrefix - APIKeys []APIKey - Catalog Catalog - Protocols Protocols - RouteOverrides map[string]RouteOverride - AllowHTTP bool - RequestTimeout time.Duration - MaxResponseBytes int64 + BaseURL string + CatalogURL string + ModelPrefix ModelPrefix + APIKeys []APIKey + Catalog Catalog + Protocols Protocols + RouteOverrides map[string]RouteOverride + ModelMetadataOverrides map[string]ModelMetadataOverride + AllowHTTP bool + RequestTimeout time.Duration + MaxResponseBytes int64 } // rawConfig mirrors the YAML shape; pointer fields distinguish "unset" // (apply default) from explicitly-set values including "" (validate as-is). // Unknown fields are ignored (host may pass extra keys). type rawConfig struct { - BaseURL *string `yaml:"base-url"` - CatalogURL *string `yaml:"catalog-url"` - ModelPrefix rawPrefix `yaml:"model-prefix"` - APIKeys []rawKey `yaml:"api-keys"` - Catalog rawCatalog `yaml:"catalog"` - Protocols rawProtocols `yaml:"protocols"` - RouteOverrides map[string]RouteOverride `yaml:"route-overrides"` - AllowHTTP bool `yaml:"allow-http"` - RequestTimeout *string `yaml:"request-timeout"` - MaxResponseBytes *int64 `yaml:"max-response-bytes"` + BaseURL *string `yaml:"base-url"` + CatalogURL *string `yaml:"catalog-url"` + ModelPrefix rawPrefix `yaml:"model-prefix"` + APIKeys []rawKey `yaml:"api-keys"` + Catalog rawCatalog `yaml:"catalog"` + Protocols rawProtocols `yaml:"protocols"` + RouteOverrides map[string]RouteOverride `yaml:"route-overrides"` + ModelMetadataOverrides map[string]ModelMetadataOverride `yaml:"model-metadata-overrides"` + AllowHTTP bool `yaml:"allow-http"` + RequestTimeout *string `yaml:"request-timeout"` + MaxResponseBytes *int64 `yaml:"max-response-bytes"` } type rawPrefix struct { @@ -143,10 +161,11 @@ func Load(yamlBytes []byte) (Config, error) { Messages: orDefault(raw.Protocols.Messages, true), Responses: orDefault(raw.Protocols.Responses, true), }, - RouteOverrides: raw.RouteOverrides, - AllowHTTP: raw.AllowHTTP, - RequestTimeout: requestTimeout, - MaxResponseBytes: orDefault(raw.MaxResponseBytes, DefaultMaxResponseBytes), + RouteOverrides: raw.RouteOverrides, + ModelMetadataOverrides: raw.ModelMetadataOverrides, + AllowHTTP: raw.AllowHTTP, + RequestTimeout: requestTimeout, + MaxResponseBytes: orDefault(raw.MaxResponseBytes, DefaultMaxResponseBytes), } if raw.CatalogURL != nil { // Mirror the derived-default trim so an explicit trailing-slash @@ -208,6 +227,19 @@ func (c Config) validate() error { return fmt.Errorf("route-overrides[%s].endpoint: must start with /", name) } } + for name, o := range c.ModelMetadataOverrides { + if o.ContextLimit != nil && *o.ContextLimit < 0 || o.OutputLimit != nil && *o.OutputLimit < 0 { + return fmt.Errorf("model-metadata-overrides[%s]: limits must not be negative", name) + } + if o.Thinking != nil { + if o.Thinking.Min != nil && *o.Thinking.Min < 0 || o.Thinking.Max != nil && *o.Thinking.Max < 0 { + return fmt.Errorf("model-metadata-overrides[%s].thinking: bounds must not be negative", name) + } + if o.Thinking.Min != nil && o.Thinking.Max != nil && *o.Thinking.Max > 0 && *o.Thinking.Min > *o.Thinking.Max { + return fmt.Errorf("model-metadata-overrides[%s].thinking: min exceeds max", name) + } + } + } if c.ModelPrefix.Enabled && !validPrefix(c.ModelPrefix.Value) { return fmt.Errorf("model-prefix.value: invalid provider-ID characters %q", c.ModelPrefix.Value) } diff --git a/internal/config/config_test.go b/internal/config/config_test.go index 30bfcdf..2d77c0c 100644 --- a/internal/config/config_test.go +++ b/internal/config/config_test.go @@ -20,6 +20,27 @@ func requireErrContains(t *testing.T, err error, want string) { // that target a later validation check. const withKey = "api-keys:\n - value: sk-dummy\n" +func TestModelMetadataOverrides(t *testing.T) { + c, err := Load([]byte(withKey + `model-metadata-overrides: + gpt-5.6-luna: + context-limit: 100 + input-modes: [text, image] + thinking: + zero-allowed: false + levels: [none, max] +`)) + if err != nil { + t.Fatal(err) + } + o := c.ModelMetadataOverrides["gpt-5.6-luna"] + if o.ContextLimit == nil || *o.ContextLimit != 100 || len(o.InputModes) != 2 || o.Thinking == nil || + o.Thinking.ZeroAllowed == nil || *o.Thinking.ZeroAllowed || len(o.Thinking.Levels) != 2 { + t.Fatalf("override decoded incorrectly: %+v", o) + } + _, err = Load([]byte(withKey + "model-metadata-overrides:\n gpt-5.6-luna:\n output-limit: -1\n")) + requireErrContains(t, err, "limits must not be negative") +} + func TestLoadMinimalAppliesAllDefaults(t *testing.T) { c, err := Load([]byte(withKey)) if err != nil { diff --git a/internal/plugin/executor_test.go b/internal/plugin/executor_test.go index cedaa5a..39c15f5 100644 --- a/internal/plugin/executor_test.go +++ b/internal/plugin/executor_test.go @@ -126,6 +126,9 @@ func upstreamRouter(t *testing.T, bodies map[string]string) func(string, []byte) var wire map[string]any _ = json.Unmarshal(payload, &wire) url, _ := wire["url"].(string) + if url == "https://models.dev/api.json" { + return hostOK(pluginapi.HTTPResponse{StatusCode: http.StatusServiceUnavailable}), nil + } for suffix, body := range bodies { if strings.HasSuffix(url, suffix) { return hostOK(pluginapi.HTTPResponse{ @@ -147,6 +150,9 @@ func wrapWithCatalog(catalogBody string, next func(string, []byte) ([]byte, erro var wire map[string]any _ = json.Unmarshal(payload, &wire) if method == pluginabi.MethodHostHTTPDo { + if url, _ := wire["url"].(string); url == "https://models.dev/api.json" { + return hostOK(pluginapi.HTTPResponse{StatusCode: http.StatusServiceUnavailable}), nil + } if url, _ := wire["url"].(string); strings.HasSuffix(url, "/models") { return hostOK(pluginapi.HTTPResponse{StatusCode: http.StatusOK, Body: []byte(catalogBody)}), nil } diff --git a/internal/plugin/integration_test.go b/internal/plugin/integration_test.go index c496b4d..0c9acfb 100644 --- a/internal/plugin/integration_test.go +++ b/internal/plugin/integration_test.go @@ -75,6 +75,9 @@ func (b *forwardingBridge) call(method string, payload []byte) ([]byte, error) { if err != nil { return hostErr("test", "undecodable host.http.do payload"), nil } + if wire.URL == "https://models.dev/api.json" { + return hostOK(pluginapi.HTTPResponse{StatusCode: http.StatusServiceUnavailable}), nil + } resp, err := b.perform(wire) if err != nil { return hostErr("transport", err.Error()), nil @@ -1116,8 +1119,15 @@ func TestNoKeyLeakageE2E(t *testing.T) { } headers, _ := wire["headers"].(map[string]any) auth, _ := headers["Authorization"].([]any) - authVal, _ := auth[0].(string) - foundInHeaders := strings.Contains(authVal, "sk-test-") + foundInHeaders := false + if len(auth) == 0 { + if wire["url"] != "https://models.dev/api.json" { + t.Fatalf("unexpected unauthenticated request: %v", wire["url"]) + } + } else { + authVal, _ := auth[0].(string) + foundInHeaders = strings.Contains(authVal, "sk-test-") + } delete(wire, "headers") rest, _ := json.Marshal(wire) for _, k := range keys { diff --git a/internal/plugin/plugin.go b/internal/plugin/plugin.go index 4f46faa..551f811 100644 --- a/internal/plugin/plugin.go +++ b/internal/plugin/plugin.go @@ -205,6 +205,11 @@ func pluginConfigFields() []pluginapi.ConfigField { Type: pluginapi.ConfigFieldTypeObject, Description: "Explicit route overrides per model (`: {protocol: string, endpoint: string}`).", }, + { + Name: "model-metadata-overrides", + Type: pluginapi.ConfigFieldTypeObject, + Description: "Per-model capability overrides (thinking, limits, input-modes, output-modes).", + }, { Name: "request-timeout", Type: pluginapi.ConfigFieldTypeString, diff --git a/internal/plugin/plugin_test.go b/internal/plugin/plugin_test.go index 24ad374..6db4ad1 100644 --- a/internal/plugin/plugin_test.go +++ b/internal/plugin/plugin_test.go @@ -426,7 +426,7 @@ func TestRegisterSuccessPublishesModels(t *testing.T) { t.Fatalf("schema_version = %d, want %d", reg.SchemaVersion, pluginabi.SchemaVersion) } if reg.Metadata.Name != "opencode-go-cliproxyapi" || reg.Metadata.Version != pluginVersion || - len(reg.Metadata.ConfigFields) != 10 { + len(reg.Metadata.ConfigFields) != 11 { t.Fatalf("metadata wrong: %+v", reg.Metadata) } if !reg.Capabilities.ModelProvider || !reg.Capabilities.AuthProvider { @@ -434,8 +434,8 @@ func TestRegisterSuccessPublishesModels(t *testing.T) { } calls := f.callsOf(pluginabi.MethodHostHTTPDo) - if len(calls) != 1 { - t.Fatalf("host.http.do calls = %d, want 1", len(calls)) + if len(calls) != 2 { + t.Fatalf("host.http.do calls = %d, want catalog and models.dev", len(calls)) } var wire map[string]any if err := json.Unmarshal(calls[0].payload, &wire); err != nil { @@ -494,6 +494,7 @@ func TestRegistrationConfigFields(t *testing.T) { {"catalog", pluginapi.ConfigFieldTypeObject}, {"protocols", pluginapi.ConfigFieldTypeObject}, {"route-overrides", pluginapi.ConfigFieldTypeObject}, + {"model-metadata-overrides", pluginapi.ConfigFieldTypeObject}, {"request-timeout", pluginapi.ConfigFieldTypeString}, {"max-response-bytes", pluginapi.ConfigFieldTypeInteger}, {"allow-http", pluginapi.ConfigFieldTypeBoolean}, @@ -766,8 +767,8 @@ func TestRegisterRefreshesWithConfiguredBearer(t *testing.T) { t.Fatalf("register rejected: %s", resp) } calls := f.callsOf(pluginabi.MethodHostHTTPDo) - if len(calls) != 1 { - t.Fatalf("host.http.do calls = %d, want 1", len(calls)) + if len(calls) != 2 { + t.Fatalf("host.http.do calls = %d, want catalog and models.dev", len(calls)) } var wire map[string]any if err := json.Unmarshal(calls[0].payload, &wire); err != nil { @@ -777,6 +778,12 @@ func TestRegisterRefreshesWithConfiguredBearer(t *testing.T) { if headers["Authorization"].([]any)[0].(string) != "Bearer "+dummyKey { t.Fatalf("expected configured bearer key, got %v", headers) } + if err := json.Unmarshal(calls[1].payload, &wire); err != nil { + t.Fatal(err) + } + if wire["url"] != "https://models.dev/api.json" || wire["headers"].(map[string]any)["Authorization"] != nil { + t.Fatalf("models.dev request must be public: %v", wire) + } } // TestRegisterWithBlockingCatalogReturnsQuickly pins the dedicated register