From 9c7e69923574ad68bca3bc70ec620280d2cb03a7 Mon Sep 17 00:00:00 2001 From: lyingbug <[email> Date: Thu, 13 Aug 2026 23:43:01 +0000 Subject: [PATCH 1/9] feat(models): add declarative model-plugin capability seam Introduce internal/models/llm as the plugin seam for model management: a vendor declares parameters, their domains, their wire encoding, and their form presentation in one descriptor, instead of spreading those facts across a provider adapter, a thinking-strategy enum, and a frontend copy of the same heuristics. - spi: Value/Param/Draft/Encoder/Constraint/Plan/Registry. Registration is reversible and validated at startup; resolution prefers model-specific descriptors over a vendor catch-all and can pin a protocol. - encoding: reusable encoders for every documented wire shape (thinking object, enable_thinking boolean, chat_template_kwargs, effort ladders, token budgets) plus the cross-parameter constraints vendors require. - vendors: built-in plugins for OpenAI, Azure, Anthropic, DeepSeek, Aliyun, Volcengine, Zhipu, Moonshot, LKEAP and self-hosted deployments, each linking the documentation it encodes. Thinking is not special-cased: it is three ordinary parameters, so the same machinery that pins Moonshot's temperature expresses Anthropic's budget. Covered by 20 golden wire-format cases asserting the exact outbound body. Co-authored-by: lyingbug --- internal/models/llm/encoding/constraint.go | 183 +++++++++ internal/models/llm/encoding/field.go | 119 ++++++ internal/models/llm/encoding/params.go | 208 ++++++++++ internal/models/llm/encoding/thinking.go | 212 +++++++++++ internal/models/llm/spi/descriptor.go | 275 ++++++++++++++ internal/models/llm/spi/draft.go | 132 +++++++ internal/models/llm/spi/encoder.go | 75 ++++ internal/models/llm/spi/kind.go | 66 ++++ internal/models/llm/spi/param.go | 231 ++++++++++++ internal/models/llm/spi/plan.go | 202 ++++++++++ internal/models/llm/spi/registry.go | 215 +++++++++++ internal/models/llm/spi/resolve.go | 119 ++++++ internal/models/llm/spi/value.go | 141 +++++++ internal/models/llm/vendors/anthropic.go | 149 ++++++++ internal/models/llm/vendors/cn_vendors.go | 317 ++++++++++++++++ internal/models/llm/vendors/openai.go | 154 ++++++++ internal/models/llm/vendors/params.go | 53 +++ internal/models/llm/vendors/vendors_test.go | 396 ++++++++++++++++++++ 18 files changed, 3247 insertions(+) create mode 100644 internal/models/llm/encoding/constraint.go create mode 100644 internal/models/llm/encoding/field.go create mode 100644 internal/models/llm/encoding/params.go create mode 100644 internal/models/llm/encoding/thinking.go create mode 100644 internal/models/llm/spi/descriptor.go create mode 100644 internal/models/llm/spi/draft.go create mode 100644 internal/models/llm/spi/encoder.go create mode 100644 internal/models/llm/spi/kind.go create mode 100644 internal/models/llm/spi/param.go create mode 100644 internal/models/llm/spi/plan.go create mode 100644 internal/models/llm/spi/registry.go create mode 100644 internal/models/llm/spi/resolve.go create mode 100644 internal/models/llm/spi/value.go create mode 100644 internal/models/llm/vendors/anthropic.go create mode 100644 internal/models/llm/vendors/cn_vendors.go create mode 100644 internal/models/llm/vendors/openai.go create mode 100644 internal/models/llm/vendors/params.go create mode 100644 internal/models/llm/vendors/vendors_test.go diff --git a/internal/models/llm/encoding/constraint.go b/internal/models/llm/encoding/constraint.go new file mode 100644 index 00000000000..1d1d4de4ee8 --- /dev/null +++ b/internal/models/llm/encoding/constraint.go @@ -0,0 +1,183 @@ +package encoding + +import ( + "fmt" + + "github.com/Tencent/WeKnora/internal/models/llm/spi" +) + +// The constraints below express the cross-parameter rules vendors document but +// no per-parameter domain can capture. They run before encoding and record a +// note for every adjustment, so a request that came back without thinking can +// be explained rather than guessed at. + +// ThinkingActive reports whether reasoning may happen, which is the predicate +// the depth constraints share. +// +// Only an explicit off makes it false. An unresolved mode means the caller had +// no opinion and the model's own default stands — and for a reasoning model +// that default is to think, so treating silence as "off" would discard a depth +// setting the request was going to honor. Auto is active too: the model may +// still choose to think, and a budget applies when it does. +func ThinkingActive(p *spi.Plan) bool { + v, ok := p.Value(spi.ParamThinkingMode) + if !ok { + return true + } + mode, _ := v.Enum() + return mode != spi.ThinkingOff +} + +// DependsOnThinking drops depth controls when thinking is off. +// +// It is what keeps a depth parameter from overwriting the mode on vendors +// where both live in the same field — DeepSeek's Responses-format control +// spells "off" as `reasoning.effort: none`, so a leftover effort value would +// silently re-enable reasoning. It also prevents the contradictory request of +// a token budget with thinking disabled, which vendors reject. +type DependsOnThinking struct { + // Params are the depth controls that require active thinking. + Params []spi.ParamID +} + +// ID reports the rule name. +func (d DependsOnThinking) ID() string { return "depends-on-thinking" } + +// Apply drops the dependent parameters when thinking is off. +func (d DependsOnThinking) Apply(p *spi.Plan) error { + if ThinkingActive(p) { + return nil + } + for _, id := range d.Params { + p.Drop(id, spi.NoteConstrained, d.ID(), + fmt.Sprintf("thinking is off, so %s does not apply", id)) + } + return nil +} + +// StreamOnlyThinking turns thinking off for non-streaming requests. +// +// Alibaba Cloud documents that several open-source Qwen3 models support +// thinking only in streaming mode and return an error otherwise. Forcing the +// mode off locally turns a vendor error into a working non-streaming answer, +// and the note records that the downgrade happened. +type StreamOnlyThinking struct{} + +// ID reports the rule name. +func (StreamOnlyThinking) ID() string { return "stream-only-thinking" } + +// Apply forces the mode off when the request does not stream. +func (s StreamOnlyThinking) Apply(p *spi.Plan) error { + if p.Stream || !ThinkingActive(p) { + return nil + } + if _, declared := p.Param(spi.ParamThinkingMode); !declared { + return nil + } + p.Adjust(spi.ParamThinkingMode, spi.EnumValue(spi.ThinkingOff), s.ID(), + "this model supports thinking only in streaming mode; thinking was turned off for this non-streaming request") + return nil +} + +// BudgetBelowMaxTokens keeps a thinking budget under the output ceiling. +// +// Anthropic requires budget_tokens to be smaller than max_tokens because +// thinking tokens count against the same ceiling, and rejects a request where +// it is not. Raising the ceiling rather than lowering the budget preserves +// what the caller asked to spend on reasoning while leaving room for an actual +// answer. +type BudgetBelowMaxTokens struct { + // Budget and MaxTokens are the parameters to relate. + Budget spi.ParamID + MaxTokens spi.ParamID + // Headroom is the number of tokens reserved for the answer itself. + Headroom int +} + +// ID reports the rule name. +func (b BudgetBelowMaxTokens) ID() string { return "budget-below-max-tokens" } + +// Apply raises the output ceiling so the budget fits beneath it. +func (b BudgetBelowMaxTokens) Apply(p *spi.Plan) error { + budgetValue, ok := p.Value(b.Budget) + if !ok { + return nil + } + budget, ok := budgetValue.Int() + if !ok { + return nil + } + headroom := b.Headroom + if headroom <= 0 { + headroom = 1 + } + required := budget + headroom + + maxValue, hasMax := p.Value(b.MaxTokens) + current, _ := maxValue.Int() + if hasMax && current >= required { + return nil + } + if _, declared := p.Param(b.MaxTokens); !declared { + return fmt.Errorf("%s needs %s to be declared so the budget can fit beneath it", b.ID(), b.MaxTokens) + } + p.Adjust(b.MaxTokens, spi.IntValue(required), b.ID(), + fmt.Sprintf("thinking budget %d must stay below %s; raised it to %d to leave room for the answer", + budget, b.MaxTokens, required)) + return nil +} + +// RequireMaxTokens supplies an output ceiling when the caller omits one. +// +// The Anthropic Messages API requires max_tokens on every request, unlike the +// OpenAI-shaped protocols where it is optional. Failing the request locally +// would be correct but unhelpful, so the plugin declares the fallback it wants +// and the note records that it was applied. +type RequireMaxTokens struct { + // Param is the ceiling parameter. + Param spi.ParamID + // Fallback is the value to send when the caller omits one. + Fallback int +} + +// ID reports the rule name. +func (r RequireMaxTokens) ID() string { return "require-max-tokens" } + +// Apply fills in the fallback ceiling. +func (r RequireMaxTokens) Apply(p *spi.Plan) error { + if p.Has(r.Param) { + return nil + } + if _, declared := p.Param(r.Param); !declared { + return fmt.Errorf("%s needs %s to be declared", r.ID(), r.Param) + } + p.Set(r.Param, spi.IntValue(r.Fallback)) + p.Note(r.Param, spi.NoteDefaulted, r.ID(), + fmt.Sprintf("this protocol requires %s; sending the plugin default of %d", r.Param, r.Fallback)) + return nil +} + +// ExclusiveWith drops one parameter when another is present. +// +// It expresses the vendor rules where two controls contradict each other, such +// as an effort ladder and a token budget that both claim to size the same +// reasoning pass. +type ExclusiveWith struct { + // Keep is the parameter that wins. + Keep spi.ParamID + // Drop is the parameter removed when Keep is present. + Drop spi.ParamID +} + +// ID reports the rule name. +func (e ExclusiveWith) ID() string { return "exclusive-with" } + +// Apply removes the losing parameter. +func (e ExclusiveWith) Apply(p *spi.Plan) error { + if !p.Has(e.Keep) || !p.Has(e.Drop) { + return nil + } + p.Drop(e.Drop, spi.NoteConstrained, e.ID(), + fmt.Sprintf("%s and %s cannot be sent together; %s takes precedence", e.Keep, e.Drop, e.Keep)) + return nil +} diff --git a/internal/models/llm/encoding/field.go b/internal/models/llm/encoding/field.go new file mode 100644 index 00000000000..7c2d8c28f27 --- /dev/null +++ b/internal/models/llm/encoding/field.go @@ -0,0 +1,119 @@ +// Package encoding provides the reusable encoders a vendor descriptor composes +// to express its wire format. +// +// Every encoder here corresponds to a shape that at least one vendor documents. +// Adding a vendor should normally mean composing these, not writing new code; +// when a vendor genuinely spells something none of them cover, the new encoder +// belongs here beside the others so the next vendor can reuse it. +package encoding + +import ( + "github.com/Tencent/WeKnora/internal/models/llm/spi" +) + +// Field encodes a value as a top-level body field. It covers the standard +// parameters every OpenAI-shaped protocol already understands, and the +// non-standard top-level fields several vendors add beside them. +type Field struct { + // Key is the JSON field name. + Key string +} + +// ID reports the encoding as the wire field it writes, which is the answer a +// user wants when asking where a knob goes. +func (f Field) ID() string { return f.Key } + +// Encode writes the value at the top level. +func (f Field) Encode(d *spi.Draft, v spi.Value) error { + d.Set(f.Key, v.JSON()) + return nil +} + +// Strip removes the field, for vendors that reject it. +func (f Field) Strip(d *spi.Draft) error { + d.Delete(f.Key) + return nil +} + +// NestedField encodes a value inside a nested object, creating the +// intermediate objects as needed. It covers `thinking.budget_tokens`, +// `reasoning.effort`, `output_config.effort`, and the other object-shaped +// vendor extensions. +type NestedField struct { + // Path is the object path, outermost first. + Path []string +} + +// Nested returns a NestedField for a path. +func Nested(path ...string) NestedField { return NestedField{Path: path} } + +// ID reports the dotted path. +func (n NestedField) ID() string { return joinPath(n.Path) } + +// Encode writes the value at the nested path. +func (n NestedField) Encode(d *spi.Draft, v spi.Value) error { + d.SetNested(v.JSON(), n.Path...) + return nil +} + +// Strip removes the leaf key, leaving any sibling fields in the same object +// intact — another parameter may legitimately own them. +func (n NestedField) Strip(d *spi.Draft) error { + if len(n.Path) == 0 { + return nil + } + if len(n.Path) == 1 { + d.Delete(n.Path[0]) + return nil + } + parent, ok := d.GetNested(n.Path[:len(n.Path)-1]...) + if !ok { + return nil + } + if obj, ok := parent.(map[string]any); ok { + delete(obj, n.Path[len(n.Path)-1]) + } + return nil +} + +// AliasField encodes a value under a vendor-specific name while removing the +// protocol's canonical spelling of the same parameter. +// +// It exists because renaming is not the same as writing a second field: the +// protocol driver has already written the canonical key from the caller's +// generic options, and OpenAI's reasoning models reject a request that carries +// both max_tokens and max_completion_tokens. +type AliasField struct { + // Canonical is the protocol's own spelling, removed on encode. + Canonical string + // Wire is the vendor's spelling, written on encode. + Wire string +} + +// ID reports the vendor spelling this alias writes. +func (a AliasField) ID() string { return a.Wire } + +// Encode moves the value from the canonical key to the vendor key. +func (a AliasField) Encode(d *spi.Draft, v spi.Value) error { + d.Delete(a.Canonical) + d.Set(a.Wire, v.JSON()) + return nil +} + +// Strip removes both spellings. +func (a AliasField) Strip(d *spi.Draft) error { + d.Delete(a.Canonical) + d.Delete(a.Wire) + return nil +} + +func joinPath(path []string) string { + out := "" + for i, seg := range path { + if i > 0 { + out += "." + } + out += seg + } + return out +} diff --git a/internal/models/llm/encoding/params.go b/internal/models/llm/encoding/params.go new file mode 100644 index 00000000000..5f4d2144a29 --- /dev/null +++ b/internal/models/llm/encoding/params.go @@ -0,0 +1,208 @@ +package encoding + +import ( + "fmt" + + "github.com/Tencent/WeKnora/internal/models/llm/spi" +) + +// The constructors below build the parameter declarations vendors share, so a +// descriptor reads as a list of claims about one vendor rather than a wall of +// struct literals. A vendor that differs overrides the field it differs in; +// everything else stays the documented baseline. + +// UI groups. The frontend renders one section per group, in this order. +const ( + // GroupSampling holds temperature and the other generation controls. + GroupSampling = "sampling" + // GroupThinking holds the reasoning controls. + GroupThinking = "thinking" + // GroupLimits holds output ceilings. + GroupLimits = "limits" +) + +// labelKey and helpKey build the i18n keys the frontend resolves. The backend +// never carries display text: it declares structure and identity, the frontend +// owns language. +func labelKey(id spi.ParamID) string { return fmt.Sprintf("model.param.%s.label", id) } +func helpKey(id spi.ParamID) string { return fmt.Sprintf("model.param.%s.help", id) } + +// optionKey builds the i18n key for one enum option. +func optionKey(id spi.ParamID, value string) string { + return fmt.Sprintf("model.param.%s.option.%s", id, value) +} + +// Options turns a vendor's documented vocabulary into enum options with their +// i18n keys attached. +func Options(id spi.ParamID, values ...string) []spi.EnumOption { + out := make([]spi.EnumOption, 0, len(values)) + for _, v := range values { + out = append(out, spi.EnumOption{Value: v, LabelKey: optionKey(id, v)}) + } + return out +} + +// Temperature declares the sampling temperature over a vendor's documented +// range. Vendors do differ here: the OpenAI-shaped range is 0–2 while several +// Chinese vendors document 0–1. +func Temperature(min, max float64) spi.Param { + return numeric(spi.ParamTemperature, spi.KindFloat, min, max, GroupSampling, 10, + Field{Key: "temperature"}) +} + +// TopP declares nucleus sampling mass. +func TopP() spi.Param { + return numeric(spi.ParamTopP, spi.KindFloat, 0, 1, GroupSampling, 20, + Field{Key: "top_p"}) +} + +// FrequencyPenalty declares the frequency penalty. +func FrequencyPenalty() spi.Param { + return numeric(spi.ParamFrequencyPenalty, spi.KindFloat, -2, 2, GroupSampling, 30, + Field{Key: "frequency_penalty"}) +} + +// PresencePenalty declares the presence penalty. +func PresencePenalty() spi.Param { + return numeric(spi.ParamPresencePenalty, spi.KindFloat, -2, 2, GroupSampling, 40, + Field{Key: "presence_penalty"}) +} + +// Seed declares a reproducible-sampling seed. +func Seed() spi.Param { + p := numeric(spi.ParamSeed, spi.KindInt, 0, 0, GroupSampling, 50, Field{Key: "seed"}) + p.Min, p.Max = nil, nil + return p +} + +// MaxTokens declares the output ceiling under its standard name. +func MaxTokens() spi.Param { + return MaxTokensAs("max_tokens") +} + +// MaxTokensAs declares the output ceiling under a vendor-specific name, +// removing the protocol's canonical spelling. OpenAI's reasoning models +// require max_completion_tokens and the Responses protocol uses +// max_output_tokens, both in place of max_tokens rather than beside it. +func MaxTokensAs(wire string) spi.Param { + var encoder spi.Encoder = Field{Key: wire} + if wire != "max_tokens" { + encoder = AliasField{Canonical: "max_tokens", Wire: wire} + } + p := numeric(spi.ParamMaxTokens, spi.KindInt, 1, 0, GroupLimits, 10, encoder) + p.Max = nil + p.UI.Widget = spi.WidgetNumber + return p +} + +// ThinkingMode declares the reasoning toggle with the modes a vendor supports. +// Passing only spi.ThinkingOn and spi.ThinkingOff yields a two-state control; +// adding spi.ThinkingAuto exposes the vendor's model-decides mode. +func ThinkingMode(encoder spi.Encoder, modes ...string) spi.Param { + if len(modes) == 0 { + modes = []string{spi.ThinkingOn, spi.ThinkingOff} + } + return spi.Param{ + ID: spi.ParamThinkingMode, + Kind: spi.KindEnum, + Enum: Options(spi.ParamThinkingMode, modes...), + Encode: encoder, + UI: spi.ParamUI{ + Group: GroupThinking, + LabelKey: labelKey(spi.ParamThinkingMode), + HelpKey: helpKey(spi.ParamThinkingMode), + Order: 10, + }, + } +} + +// ThinkingEffort declares a reasoning-depth ladder over the vendor's own +// vocabulary. The ladders are genuinely different — OpenAI, Zhipu, DeepSeek, +// and Alibaba each publish their own rungs — so this takes the values rather +// than normalizing them into a shared scale that would send rungs the vendor +// does not accept. +func ThinkingEffort(encoder spi.Encoder, values ...string) spi.Param { + return spi.Param{ + ID: spi.ParamThinkingEffort, + Kind: spi.KindEnum, + Enum: Options(spi.ParamThinkingEffort, values...), + Encode: encoder, + UI: spi.ParamUI{ + Group: GroupThinking, + LabelKey: labelKey(spi.ParamThinkingEffort), + HelpKey: helpKey(spi.ParamThinkingEffort), + Order: 20, + }, + } +} + +// ThinkingBudget declares a reasoning-token cap over the vendor's documented +// range. +func ThinkingBudget(encoder spi.Encoder, min, max int) spi.Param { + p := numeric(spi.ParamThinkingBudget, spi.KindInt, float64(min), float64(max), + GroupThinking, 30, encoder) + if max <= 0 { + p.Max = nil + } + p.UI.Widget = spi.WidgetNumber + return p +} + +// Forbidden declares a parameter the vendor rejects, so the field is removed +// from the request and the form hides the control. +// +// Declaring it is not the same as omitting it: the protocol driver writes a +// canonical body from the caller's generic options, so a forbidden parameter +// must be actively removed, and saying so here is what tells the form not to +// offer a knob that cannot work. +func Forbidden(id spi.ParamID, kind spi.ValueKind, wireKey string) spi.Param { + return spi.Param{ + ID: id, + Kind: kind, + Support: spi.SupportForbidden, + Encode: Field{Key: wireKey}, + UI: spi.ParamUI{Hidden: true}, + } +} + +// Pinned declares a parameter the vendor accepts at exactly one value. +func Pinned(id spi.ParamID, kind spi.ValueKind, value spi.Value, encoder spi.Encoder) spi.Param { + return spi.Param{ + ID: id, + Kind: kind, + Support: spi.SupportPinned, + Pin: spi.Ptr(value), + Encode: encoder, + UI: spi.ParamUI{Hidden: true}, + } +} + +// numeric builds a bounded numeric parameter with its UI metadata. +func numeric(id spi.ParamID, kind spi.ValueKind, min, max float64, group string, order int, encoder spi.Encoder) spi.Param { + return spi.Param{ + ID: id, + Kind: kind, + Min: spi.Float(min), + Max: spi.Float(max), + Encode: encoder, + UI: spi.ParamUI{ + Group: group, + LabelKey: labelKey(id), + HelpKey: helpKey(id), + Order: order, + }, + } +} + +// SamplingSet returns the sampling parameters an OpenAI-shaped vendor supports +// without deviation, which is the majority case. +func SamplingSet(maxTemperature float64) []spi.Param { + return []spi.Param{ + Temperature(0, maxTemperature), + TopP(), + FrequencyPenalty(), + PresencePenalty(), + Seed(), + MaxTokens(), + } +} diff --git a/internal/models/llm/encoding/thinking.go b/internal/models/llm/encoding/thinking.go new file mode 100644 index 00000000000..a6374db921d --- /dev/null +++ b/internal/models/llm/encoding/thinking.go @@ -0,0 +1,212 @@ +package encoding + +import ( + "fmt" + + "github.com/Tencent/WeKnora/internal/models/llm/spi" +) + +// The thinking encoders below map the neutral thinking mode onto each +// documented wire shape. Reasoning depth needs no dedicated encoder: an effort +// ladder is a plain field (`reasoning_effort`) or a nested one +// (`reasoning.effort`, `output_config.effort`), and a token budget likewise +// (`thinking_budget`, `thinking.budget_tokens`), so Field and NestedField +// already cover them. Only the mode needs translating, because "on" is spelled +// as a boolean by one vendor, an enum by another, and a nested template +// argument by a third. + +// ThinkingObject encodes the mode as an object with a type discriminator: +// +// {"thinking": {"type": "enabled"}} +// +// It is the most widely adopted shape, used by DeepSeek, Zhipu GLM, Volcengine +// Ark, and Tencent LKEAP. The vendor spellings are parameters because they do +// differ: Ark documents a third `auto` value that lets the model decide, and +// Anthropic reuses the same object with `adaptive` in that role. +type ThinkingObject struct { + // Key is the body field holding the object, `thinking` for every + // documented user of this shape. + Key string + // On, Off, and Auto are the vendor's spellings. An empty spelling means + // the vendor does not document that mode, and encoding it is an error + // rather than a silent omission. + On string + Off string + Auto string +} + +// ID reports the wire field. +func (t ThinkingObject) ID() string { return t.Key + ".type" } + +// Encode writes the vendor's spelling of the requested mode. +func (t ThinkingObject) Encode(d *spi.Draft, v spi.Value) error { + spelling, err := t.spell(v) + if err != nil { + return err + } + d.SetNested(spelling, t.Key, "type") + return nil +} + +// Strip removes the whole object, including any budget a sibling parameter +// wrote: a disabled thinking configuration carrying a budget is contradictory, +// and vendors reject it. +func (t ThinkingObject) Strip(d *spi.Draft) error { + d.Delete(t.Key) + return nil +} + +func (t ThinkingObject) spell(v spi.Value) (string, error) { + mode, ok := v.Enum() + if !ok { + return "", fmt.Errorf("thinking mode must be an enum, got %s", v.Kind) + } + var spelling string + switch mode { + case spi.ThinkingOn: + spelling = t.On + case spi.ThinkingOff: + spelling = t.Off + case spi.ThinkingAuto: + spelling = t.Auto + default: + return "", fmt.Errorf("unknown thinking mode %q", mode) + } + if spelling == "" { + return "", fmt.Errorf("thinking mode %q has no wire spelling for %s", mode, t.Key) + } + return spelling, nil +} + +// EnableThinkingBool encodes the mode as a top-level boolean: +// +// {"enable_thinking": true} +// +// This is Alibaba Cloud Model Studio's shape for the Qwen hybrid-thinking +// models. The OpenAI Python SDK reaches it through `extra_body`, which is only +// that SDK's mechanism for passing a non-standard field: on the wire it is an +// ordinary top-level key, and Model Studio's own batch-inference documentation +// says so explicitly by requiring it beside `model`. +// +// The shape has no third state, so a descriptor using it must not offer auto. +type EnableThinkingBool struct { + // Key is the body field, `enable_thinking` as documented. + Key string +} + +// ID reports the wire field. +func (e EnableThinkingBool) ID() string { return e.Key } + +// Encode writes the boolean form of the requested mode. +func (e EnableThinkingBool) Encode(d *spi.Draft, v spi.Value) error { + mode, ok := v.Enum() + if !ok { + return fmt.Errorf("thinking mode must be an enum, got %s", v.Kind) + } + switch mode { + case spi.ThinkingOn: + d.Set(e.Key, true) + case spi.ThinkingOff: + d.Set(e.Key, false) + default: + return fmt.Errorf("%s cannot express thinking mode %q", e.Key, mode) + } + return nil +} + +// Strip removes the field. +func (e EnableThinkingBool) Strip(d *spi.Draft) error { + d.Delete(e.Key) + return nil +} + +// ChatTemplateKwargs encodes the mode as a chat-template argument: +// +// {"chat_template_kwargs": {"enable_thinking": true}} +// +// This is how a self-hosted inference server passes the flag through to the +// model's chat template, so it is the right default for vLLM-style generic +// deployments and NVIDIA NIM rather than for any hosted vendor API. +type ChatTemplateKwargs struct { + // Key is the outer field, `chat_template_kwargs` by convention. + Key string + // Arg is the template argument name, `enable_thinking` by convention. + Arg string +} + +// ID reports the dotted path. +func (c ChatTemplateKwargs) ID() string { return c.Key + "." + c.Arg } + +// Encode writes the boolean template argument. +func (c ChatTemplateKwargs) Encode(d *spi.Draft, v spi.Value) error { + mode, ok := v.Enum() + if !ok { + return fmt.Errorf("thinking mode must be an enum, got %s", v.Kind) + } + switch mode { + case spi.ThinkingOn: + d.SetNested(true, c.Key, c.Arg) + case spi.ThinkingOff: + d.SetNested(false, c.Key, c.Arg) + default: + return fmt.Errorf("%s cannot express thinking mode %q", c.ID(), mode) + } + return nil +} + +// Strip removes the template argument, leaving other arguments in place. +func (c ChatTemplateKwargs) Strip(d *spi.Draft) error { + return Nested(c.Key, c.Arg).Strip(d) +} + +// EffortNone encodes the mode by writing a sentinel onto an effort ladder whose +// lowest rung means "do not think": +// +// {"reasoning": {"effort": "none"}} +// +// DeepSeek's Responses-format control and Alibaba's Responses API both +// document `none` this way, so on those surfaces the mode toggle and the depth +// control are the same field. Keeping that as its own encoder means the +// descriptor still declares two parameters — a toggle a user understands and a +// depth ladder — while the wire carries the one field the vendor defines. +type EffortNone struct { + // Encoder writes the effort value; it is the same encoder the depth + // parameter uses, so both agree on where the field lives. + Encoder spi.Encoder + // Off is the ladder rung meaning "no thinking", `none` as documented. + Off string + // On is the rung to write when thinking is requested but no explicit depth + // was given. Empty leaves the vendor's own default in place, which is the + // safer choice when the vendor documents one. + On string +} + +// ID reports the underlying field. +func (e EffortNone) ID() string { return e.Encoder.ID() } + +// Encode writes the ladder rung standing for the requested mode. +func (e EffortNone) Encode(d *spi.Draft, v spi.Value) error { + mode, ok := v.Enum() + if !ok { + return fmt.Errorf("thinking mode must be an enum, got %s", v.Kind) + } + switch mode { + case spi.ThinkingOff: + return e.Encoder.Encode(d, spi.EnumValue(e.Off)) + case spi.ThinkingOn: + if e.On == "" { + return nil + } + return e.Encoder.Encode(d, spi.EnumValue(e.On)) + default: + return fmt.Errorf("%s cannot express thinking mode %q", e.ID(), mode) + } +} + +// Strip removes the effort field. +func (e EffortNone) Strip(d *spi.Draft) error { + if stripper, ok := e.Encoder.(spi.Stripper); ok { + return stripper.Strip(d) + } + return nil +} diff --git a/internal/models/llm/spi/descriptor.go b/internal/models/llm/spi/descriptor.go new file mode 100644 index 00000000000..6a0c4e3b02e --- /dev/null +++ b/internal/models/llm/spi/descriptor.go @@ -0,0 +1,275 @@ +package spi + +import ( + "fmt" + "regexp" + "strings" +) + +// AuthKind is how a plugin authenticates a request. +type AuthKind string + +const ( + // AuthBearer sends `Authorization: Bearer `, the common case. + AuthBearer AuthKind = "bearer" + // AuthHeader sends the key in a named header, such as Azure's `api-key` + // or Anthropic's `x-api-key`. + AuthHeader AuthKind = "header" + // AuthSigned delegates to a Signer, for vendors that sign request bytes. + AuthSigned AuthKind = "signed" +) + +// Signer produces authentication headers over the exact outbound bytes. It is +// an interface rather than a declaration because request signing is genuinely +// code, and pretending otherwise would push a signing algorithm into config. +type Signer interface { + // Sign returns the headers authenticating body for this request. + Sign(body []byte, creds Credentials) (map[string]string, error) +} + +// Credentials carries the secrets a plugin may need. Bearer and header schemes +// use APIKey alone; signing schemes also need the application pair. +type Credentials struct { + APIKey string + AppID string + AppSecret string +} + +// Auth declares how requests are authenticated. +type Auth struct { + // Kind selects the scheme; the zero value is AuthBearer. + Kind AuthKind `json:"kind,omitempty"` + // Header names the header for AuthHeader. + Header string `json:"header,omitempty"` + // Signer implements AuthSigned. + Signer Signer `json:"-"` + // Static are headers sent on every request, such as Anthropic's required + // `anthropic-version`. + Static map[string]string `json:"static,omitempty"` +} + +// EffectiveKind reports the scheme, treating the zero value as AuthBearer. +func (a Auth) EffectiveKind() AuthKind { + if a.Kind == "" { + return AuthBearer + } + return a.Kind +} + +// ModelMatcher selects the models a descriptor applies to, so one vendor can +// publish several descriptors — a reasoning-model descriptor and a general one, +// for instance — without a predicate function that only its author can audit. +// An empty matcher matches every model, which is the vendor's catch-all. +type ModelMatcher struct { + // Prefixes matches a case-insensitive model-name prefix. + Prefixes []string `json:"prefixes,omitempty"` + // Contains matches a case-insensitive substring. + Contains []string `json:"contains,omitempty"` + // Pattern matches a regular expression against the lowercased name. + Pattern string `json:"pattern,omitempty"` + + compiled *regexp.Regexp +} + +// IsCatchAll reports whether the matcher accepts every model. +func (m ModelMatcher) IsCatchAll() bool { + return len(m.Prefixes) == 0 && len(m.Contains) == 0 && m.Pattern == "" +} + +// Matches reports whether the matcher accepts a model name. +func (m ModelMatcher) Matches(model string) bool { + if m.IsCatchAll() { + return true + } + name := strings.ToLower(strings.TrimSpace(model)) + for _, prefix := range m.Prefixes { + if strings.HasPrefix(name, strings.ToLower(prefix)) { + return true + } + } + for _, sub := range m.Contains { + if strings.Contains(name, strings.ToLower(sub)) { + return true + } + } + if m.compiled != nil { + return m.compiled.MatchString(name) + } + return false +} + +// compile prepares the regular expression, reporting a malformed pattern at +// registration time rather than on the first request that uses it. +func (m *ModelMatcher) compile() error { + if m.Pattern == "" { + return nil + } + re, err := regexp.Compile(m.Pattern) + if err != nil { + return fmt.Errorf("compile model pattern %q: %w", m.Pattern, err) + } + m.compiled = re + return nil +} + +// ReasoningReplay declares whether prior-turn reasoning must be sent back to +// the vendor. It is a wire requirement, not a preference: DeepSeek returns 400 +// when a tool-calling turn's reasoning_content is dropped, and Anthropic +// requires thinking blocks to return unmodified. +type ReasoningReplay string + +const ( + // ReplayNever means prior reasoning is ignored and need not be sent. + ReplayNever ReasoningReplay = "never" + // ReplayWithTools means reasoning must be replayed on turns that carry + // tool calls, and may be dropped otherwise. + ReplayWithTools ReasoningReplay = "with_tools" + // ReplayAlways means every prior reasoning block must be replayed verbatim. + ReplayAlways ReasoningReplay = "always" +) + +// Descriptor is a vendor plugin: everything that distinguishes one vendor's +// model API from the protocol baseline, declared rather than coded. +// +// A descriptor is registered per (vendor, kind), and a vendor may register +// several for one kind when its models genuinely differ — an OpenAI reasoning +// descriptor beside the general one, for example. Resolution prefers the +// descriptor whose matcher is specific. +type Descriptor struct { + // Vendor is the provider identity, matching the stored model's provider + // field, e.g. "deepseek". + Vendor string + // Kind is the capability this descriptor serves. + Kind ModelKind + // Protocol is the wire protocol the vendor speaks. + Protocol ProtocolID + // Models restricts the descriptor to matching model names. An empty + // matcher makes it the vendor's catch-all. + Models ModelMatcher + // DisplayName is the human-readable vendor name for diagnostics; the UI + // resolves its own label from the provider catalog. + DisplayName string + + // DefaultBaseURL is used when the model configuration leaves it empty. + DefaultBaseURL string + // EndpointPath overrides the protocol's standard path, for vendors serving + // a compatible protocol from a different route. + EndpointPath string + // Auth declares the authentication scheme. + Auth Auth + + // Params declares every request parameter this plugin exposes, in the + // order a form should present them. + Params []Param + // Constraints are cross-parameter rules applied after resolution and + // before encoding. + Constraints []Constraint + + // ReasoningReplay declares whether prior-turn reasoning must be sent back. + ReasoningReplay ReasoningReplay + // DocURL points at the vendor documentation this descriptor follows. + DocURL string +} + +// Param reports a declared parameter by id. +func (d Descriptor) Param(id ParamID) (Param, bool) { + for _, p := range d.Params { + if p.ID == id { + return p, true + } + } + return Param{}, false +} + +// Supports reports whether the descriptor declares a parameter at all, which +// is what the UI asks before rendering a control and what the debug endpoint +// asks before offering a toggle. +func (d Descriptor) Supports(id ParamID) bool { + p, ok := d.Param(id) + return ok && p.EffectiveSupport() != SupportForbidden +} + +// EffectiveReplay reports the reasoning-replay rule, treating the zero value +// as ReplayNever. +func (d Descriptor) EffectiveReplay() ReasoningReplay { + if d.ReasoningReplay == "" { + return ReplayNever + } + return d.ReasoningReplay +} + +// validate checks a descriptor's internal consistency at registration time. +// These are the mistakes a plugin author actually makes, and catching them at +// startup beats discovering them through a vendor's 400 in production. +func (d *Descriptor) validate() error { + if strings.TrimSpace(d.Vendor) == "" { + return fmt.Errorf("descriptor: vendor is required") + } + if d.Kind == "" { + return fmt.Errorf("descriptor %s: kind is required", d.Vendor) + } + if d.Protocol == "" { + return fmt.Errorf("descriptor %s/%s: protocol is required", d.Vendor, d.Kind) + } + if err := d.Models.compile(); err != nil { + return fmt.Errorf("descriptor %s/%s: %w", d.Vendor, d.Kind, err) + } + if d.Auth.EffectiveKind() == AuthHeader && strings.TrimSpace(d.Auth.Header) == "" { + return fmt.Errorf("descriptor %s/%s: header auth requires a header name", d.Vendor, d.Kind) + } + if d.Auth.EffectiveKind() == AuthSigned && d.Auth.Signer == nil { + return fmt.Errorf("descriptor %s/%s: signed auth requires a signer", d.Vendor, d.Kind) + } + seen := make(map[ParamID]struct{}, len(d.Params)) + for _, p := range d.Params { + if _, dup := seen[p.ID]; dup { + return fmt.Errorf("descriptor %s/%s: duplicate parameter %s", d.Vendor, d.Kind, p.ID) + } + seen[p.ID] = struct{}{} + if err := validateParam(d, p); err != nil { + return err + } + } + return nil +} + +func validateParam(d *Descriptor, p Param) error { + where := fmt.Sprintf("descriptor %s/%s parameter %s", d.Vendor, d.Kind, p.ID) + if p.Kind == "" { + return fmt.Errorf("%s: kind is required", where) + } + if p.Encode == nil { + return fmt.Errorf("%s: every parameter needs an encoder", where) + } + + // A forbidden parameter is never sent, so its domain is irrelevant; what it + // does need is the ability to remove itself, because the protocol driver + // has already written the canonical field. + if p.EffectiveSupport() == SupportForbidden { + if _, ok := p.Encode.(Stripper); !ok { + return fmt.Errorf("%s: a forbidden parameter needs an encoder that can strip its field, "+ + "otherwise the protocol's own value would still reach the vendor", where) + } + return nil + } + + if p.Kind == KindEnum && len(p.Enum) == 0 { + return fmt.Errorf("%s: enum parameter needs at least one option", where) + } + if p.Kind != KindEnum && len(p.Enum) > 0 { + return fmt.Errorf("%s: only enum parameters may declare options", where) + } + if p.EffectiveSupport() == SupportPinned && p.Pin == nil { + return fmt.Errorf("%s: pinned parameter needs a pin value", where) + } + if p.Pin != nil && !p.AllowsValue(*p.Pin) { + return fmt.Errorf("%s: pin value %s is outside the declared domain", where, p.Pin) + } + if p.Default != nil && !p.AllowsValue(*p.Default) { + return fmt.Errorf("%s: default value %s is outside the declared domain", where, p.Default) + } + if p.Min != nil && p.Max != nil && *p.Min > *p.Max { + return fmt.Errorf("%s: min %v exceeds max %v", where, *p.Min, *p.Max) + } + return nil +} diff --git a/internal/models/llm/spi/draft.go b/internal/models/llm/spi/draft.go new file mode 100644 index 00000000000..aba5d729fb4 --- /dev/null +++ b/internal/models/llm/spi/draft.go @@ -0,0 +1,132 @@ +package spi + +// Draft is the mutable outbound request an encoder edits. The protocol driver +// builds the canonical body; encoders then write the vendor's non-standard +// fields onto it. +// +// Body is a map rather than a typed struct on purpose. Vendor extensions are +// exactly the fields no shared struct can enumerate — enable_thinking, +// thinking.budget_tokens, chat_template_kwargs, extra_content.google — and a +// typed body would force every one of them through an embedded-struct hack +// plus a separate "must send raw" flag, which is how the previous design +// accumulated four parallel request types. A map also makes the outbound +// bytes directly assertable in tests, so a vendor claim can be checked +// against its documentation. +type Draft struct { + // Protocol is the wire protocol this draft targets. An encoder written for + // one protocol must not be attached to a descriptor using another. + Protocol ProtocolID + // Model is the model identifier as the vendor expects it. + Model string + // Stream reports whether this is a streaming request, which several + // vendors treat as a hard constraint rather than a preference. + Stream bool + // Body is the outbound JSON body. + Body map[string]any + // Header carries protocol or vendor headers beyond authentication, such as + // Anthropic's beta opt-ins. + Header map[string]string + // Endpoint overrides the URL the consumer would otherwise derive from the + // base URL. Empty keeps the protocol's standard path. + Endpoint string +} + +// NewDraft returns a draft with initialized maps. +func NewDraft(protocol ProtocolID, model string, stream bool) *Draft { + return &Draft{ + Protocol: protocol, + Model: model, + Stream: stream, + Body: map[string]any{}, + Header: map[string]string{}, + } +} + +// Set writes a top-level body field. +func (d *Draft) Set(key string, value any) { + if d.Body == nil { + d.Body = map[string]any{} + } + d.Body[key] = value +} + +// Get reads a top-level body field. +func (d *Draft) Get(key string) (any, bool) { + if d.Body == nil { + return nil, false + } + v, ok := d.Body[key] + return v, ok +} + +// Delete removes a top-level body field. Used by forbidden parameters, whose +// contract is that the field must not reach the wire at all. +func (d *Draft) Delete(key string) { + delete(d.Body, key) +} + +// SetNested writes value at a nested object path, creating intermediate +// objects as needed. It is how `thinking.type`, `chat_template_kwargs +// .enable_thinking`, and `output_config.effort` are expressed without each +// encoder hand-rolling the same map bookkeeping. +// +// An intermediate key holding a non-object is replaced, because the encoder +// declaring the path owns that subtree. +func (d *Draft) SetNested(value any, path ...string) { + if len(path) == 0 { + return + } + if d.Body == nil { + d.Body = map[string]any{} + } + node := d.Body + for _, key := range path[:len(path)-1] { + child, ok := node[key].(map[string]any) + if !ok { + child = map[string]any{} + node[key] = child + } + node = child + } + node[path[len(path)-1]] = value +} + +// GetNested reads the value at a nested object path. +func (d *Draft) GetNested(path ...string) (any, bool) { + if len(path) == 0 || d.Body == nil { + return nil, false + } + var node any = d.Body + for _, key := range path { + obj, ok := node.(map[string]any) + if !ok { + return nil, false + } + node, ok = obj[key] + if !ok { + return nil, false + } + } + return node, true +} + +// SetHeader records an outbound header. +func (d *Draft) SetHeader(key, value string) { + if d.Header == nil { + d.Header = map[string]string{} + } + d.Header[key] = value +} + +// Rename moves a body field to a different key, preserving its value. It +// expresses the common case of a vendor spelling a standard parameter +// differently, such as OpenAI reasoning models requiring max_completion_tokens +// in place of max_tokens. +func (d *Draft) Rename(from, to string) { + v, ok := d.Get(from) + if !ok { + return + } + d.Delete(from) + d.Set(to, v) +} diff --git a/internal/models/llm/spi/encoder.go b/internal/models/llm/spi/encoder.go new file mode 100644 index 00000000000..0f2568f3b93 --- /dev/null +++ b/internal/models/llm/spi/encoder.go @@ -0,0 +1,75 @@ +package spi + +// Encoder writes one resolved parameter value onto the outbound draft. It is +// the whole of "how this vendor spells this knob". +// +// Encoders are small and composable so a new vendor is usually a descriptor +// that reuses existing encoders, not new code. When a vendor genuinely spells +// something no existing encoder covers, the new encoder lives beside the +// others in internal/models/llm/encoding and becomes reusable in turn. +type Encoder interface { + // ID names the encoding. It is reported in diagnostics and surfaced to the + // model editor so a user can see which wire field a knob maps to, which is + // the question the previous free-text thinking_control setting was really + // trying to answer. + ID() string + // Encode writes v onto the draft. It is invoked only with a value the + // parameter's declared domain already accepted. + Encode(d *Draft, v Value) error +} + +// EncoderFunc adapts a function to Encoder, for the many encoders that need no +// state beyond their id. +type EncoderFunc struct { + Name string + Fn func(d *Draft, v Value) error +} + +// ID reports the encoder name. +func (e EncoderFunc) ID() string { return e.Name } + +// Encode applies the wrapped function. +func (e EncoderFunc) Encode(d *Draft, v Value) error { return e.Fn(d, v) } + +// Stripper is an encoder that can also remove its field from a draft. +// +// It exists for forbidden parameters, and the distinction matters: the +// protocol driver builds a canonical body that may already carry the field +// from the caller's generic options, so a vendor that rejects the field needs +// it actively removed, not merely "not set". OpenAI's reasoning models reject +// temperature outright, so leaving the protocol's value in place would fail +// the request rather than degrade politely. +// +// It is an optional interface so ordinary encoders stay a single method. +type Stripper interface { + // Strip removes the encoder's field from the draft. + Strip(d *Draft) error +} + +// Constraint is a cross-parameter rule the plugin enforces before anything is +// encoded: a relationship the per-parameter domains cannot express. +// +// Constraints run against the resolved Plan, so they can drop a value, adjust +// it, or reject the request outright — and every adjustment they make is +// recorded as a note rather than applied silently. A user who asked for +// thinking and did not get it should be able to find out why. +type Constraint interface { + // ID names the rule for diagnostics. + ID() string + // Apply inspects and may adjust the plan. Returning an error rejects the + // request, which is correct when the vendor would reject it anyway and a + // local error is clearer than a remote 400. + Apply(p *Plan) error +} + +// ConstraintFunc adapts a function to Constraint. +type ConstraintFunc struct { + Name string + Fn func(p *Plan) error +} + +// ID reports the constraint name. +func (c ConstraintFunc) ID() string { return c.Name } + +// Apply runs the wrapped function. +func (c ConstraintFunc) Apply(p *Plan) error { return c.Fn(p) } diff --git a/internal/models/llm/spi/kind.go b/internal/models/llm/spi/kind.go new file mode 100644 index 00000000000..c05d84e7a72 --- /dev/null +++ b/internal/models/llm/spi/kind.go @@ -0,0 +1,66 @@ +// Package spi defines the model-plugin capability seam: the declarative +// descriptors a vendor plugin publishes, and the registry that resolves one +// for a configured model. +// +// The seam has three roles, and adding a capability means designing all three: +// +// - Definition — the descriptors and interfaces in this package. +// - Provider — a protocol driver (internal/models/llm/protocol) plus a vendor +// descriptor (internal/models/llm/vendors) declaring how it differs. +// - Consumer — the model factories in internal/models, plus the capability +// API that renders the model editor form. +// +// One rule keeps the seam honest: a vendor plugin DECLARES facts, it never +// branches on a provider name. Anything that would otherwise need a central +// `switch providerName` belongs in a descriptor field instead. That is what +// lets a single declaration drive request building, validation, the frontend +// form, and diagnostics without those four surfaces drifting apart. +package spi + +// ModelKind is the capability a model serves. A plugin declares the kinds it +// handles, so one descriptor can cover chat while another covers embeddings +// for the same vendor. +type ModelKind string + +const ( + // KindChat is a conversational completion model. + KindChat ModelKind = "chat" + // KindEmbedding is a text or multimodal embedding model. + KindEmbedding ModelKind = "embedding" + // KindRerank is a query/document relevance scoring model. + KindRerank ModelKind = "rerank" + // KindVision is a vision-language model driven through a dedicated API. + KindVision ModelKind = "vision" + // KindASR is a speech-to-text model. + KindASR ModelKind = "asr" +) + +// ProtocolID names a wire protocol. A protocol owns request serialization, +// response parsing, and stream decoding; a vendor plugin picks one and +// declares only its deltas. +// +// The three standard protocols are the supported baseline. A vendor needing +// something none of them express implements its own driver and names it here. +type ProtocolID string + +const ( + // ProtocolOpenAIChat is the OpenAI Chat Completions protocol + // (POST /chat/completions). The de-facto industry baseline: most + // vendors are "OpenAI-compatible plus a few non-standard fields". + ProtocolOpenAIChat ProtocolID = "openai-chat" + // ProtocolOpenAIResponses is the OpenAI Responses protocol + // (POST /responses). Item-based input and output with first-class + // reasoning items and a fully event-typed stream. + ProtocolOpenAIResponses ProtocolID = "openai-responses" + // ProtocolAnthropicMessages is the Anthropic Messages protocol + // (POST /v1/messages). Content-block based, with thinking blocks + // that must be replayed verbatim across turns. + ProtocolAnthropicMessages ProtocolID = "anthropic-messages" + // ProtocolOllama is the local Ollama protocol, named here so local models + // resolve through the same registry as remote ones. + ProtocolOllama ProtocolID = "ollama" +) + +// String reports the protocol id so it can be logged and surfaced in +// diagnostics without a conversion at every call site. +func (p ProtocolID) String() string { return string(p) } diff --git a/internal/models/llm/spi/param.go b/internal/models/llm/spi/param.go new file mode 100644 index 00000000000..242a6b93dbf --- /dev/null +++ b/internal/models/llm/spi/param.go @@ -0,0 +1,231 @@ +package spi + +// ParamID is the neutral identity of a request parameter, shared across every +// vendor. Callers, stored configuration, and the UI all speak ParamIDs; each +// plugin declares how its vendor spells them on the wire. +// +// Thinking is not privileged here. It is three ordinary parameters, which is +// exactly why the same machinery that pins Moonshot's temperature can also +// express Anthropic's thinking budget. +type ParamID string + +const ( + // ParamTemperature is the sampling temperature. + ParamTemperature ParamID = "temperature" + // ParamTopP is nucleus sampling mass. + ParamTopP ParamID = "top_p" + // ParamMaxTokens is the ceiling on generated tokens. Vendors disagree on + // the wire name (max_tokens / max_completion_tokens / max_output_tokens), + // which is an encoder concern, not a separate parameter. + ParamMaxTokens ParamID = "max_tokens" + // ParamFrequencyPenalty penalizes token frequency. + ParamFrequencyPenalty ParamID = "frequency_penalty" + // ParamPresencePenalty penalizes token presence. + ParamPresencePenalty ParamID = "presence_penalty" + // ParamSeed requests reproducible sampling. + ParamSeed ParamID = "seed" + // ParamToolChoice steers tool selection. + ParamToolChoice ParamID = "tool_choice" + // ParamParallelToolCalls allows several tool calls per assistant turn. + ParamParallelToolCalls ParamID = "parallel_tool_calls" + + // ParamThinkingMode turns reasoning on, off, or hands the decision to the + // model. Its wire spelling is the single largest source of vendor + // divergence: enable_thinking, thinking.type, chat_template_kwargs, + // reasoning.effort=none, or nothing at all for thinking-only models. + ParamThinkingMode ParamID = "thinking.mode" + // ParamThinkingEffort selects reasoning depth from a vendor's own ladder. + // The vocabularies genuinely differ, so this is an enum over the vendor's + // documented values rather than a normalized scale. + ParamThinkingEffort ParamID = "thinking.effort" + // ParamThinkingBudget caps reasoning tokens. + ParamThinkingBudget ParamID = "thinking.budget" + // ParamThinkingSummary requests a readable summary of the reasoning, for + // the protocols that expose one. It is separate from the mode because a + // model can reason without surfacing any of it: OpenAI's Responses API + // omits the summary entirely unless it is asked for. + ParamThinkingSummary ParamID = "thinking.summary" +) + +// String reports the parameter id for logging and diagnostics. +func (p ParamID) String() string { return string(p) } + +// Thinking mode is an enum rather than a bool because a third state is real: +// several vendors let the model decide per request (Ark `auto`, Anthropic +// adaptive, Zhipu dynamic thinking). Collapsing that into a bool would make +// those modes unreachable. +const ( + // ThinkingOff disables reasoning. + ThinkingOff = "off" + // ThinkingOn forces reasoning before the answer. + ThinkingOn = "on" + // ThinkingAuto lets the model decide whether to reason on each request. + ThinkingAuto = "auto" +) + +// Support is how a plugin dispositions a parameter. +type Support string + +const ( + // SupportUser means the caller may set the parameter, within the declared + // domain, and it reaches the wire. + SupportUser Support = "user" + // SupportPinned means the plugin always sends its Pin value and the caller + // cannot override it — Moonshot's temperature=1, for instance. + SupportPinned Support = "pinned" + // SupportForbidden means the parameter must never reach the wire. The + // OpenAI reasoning models reject temperature outright, so sending a + // caller's value would fail the request rather than degrade politely. + SupportForbidden Support = "forbidden" +) + +// EnumOption is one member of a vendor's closed vocabulary, carrying the wire +// value together with the i18n key a form uses to label it. +type EnumOption struct { + // Value is the exact string the vendor's API accepts. + Value string `json:"value"` + // LabelKey is an i18n key the frontend resolves; the backend never + // hardcodes display text so language stays a frontend concern. + LabelKey string `json:"label_key,omitempty"` + // HelpKey optionally explains the option in the form. + HelpKey string `json:"help_key,omitempty"` +} + +// Widget is the control a form should render for a parameter. It is a hint +// derived from the parameter's shape, kept explicit so a plugin can ask for a +// slider where the generic mapping would pick a number box. +type Widget string + +const ( + // WidgetSwitch renders a boolean toggle. + WidgetSwitch Widget = "switch" + // WidgetSelect renders a closed vocabulary. + WidgetSelect Widget = "select" + // WidgetNumber renders a numeric input. + WidgetNumber Widget = "number" + // WidgetSlider renders a bounded numeric range. + WidgetSlider Widget = "slider" +) + +// ParamUI is how a form should present a parameter. It is part of the same +// declaration as the wire encoding on purpose: a vendor that adds a knob gets +// the form field for free, and a knob that disappears cannot linger in the UI. +type ParamUI struct { + // Hidden keeps a parameter off the form while still honoring it on the + // wire (pinned and forbidden parameters are usually hidden). + Hidden bool `json:"hidden,omitempty"` + // Group buckets the field in the form, e.g. "thinking" or "sampling". + Group string `json:"group,omitempty"` + // LabelKey and HelpKey are i18n keys resolved by the frontend. + LabelKey string `json:"label_key,omitempty"` + HelpKey string `json:"help_key,omitempty"` + // Widget overrides the control inferred from the parameter's kind. + Widget Widget `json:"widget,omitempty"` + // Order sorts fields within a group; equal values keep declaration order. + Order int `json:"order,omitempty"` +} + +// Param declares one request parameter: its neutral identity, its accepted +// domain, how it reaches the wire, and how a form presents it. +// +// This single declaration is the source of truth for four surfaces — request +// building, server-side validation, the model editor form, and the debug +// endpoint's report — so they cannot drift apart the way a hand-maintained +// backend table and frontend copy always do. +type Param struct { + // ID is the neutral parameter identity. + ID ParamID `json:"id"` + // Kind is the value domain. + Kind ValueKind `json:"kind"` + // Support dispositions the parameter; the zero value is SupportUser. + Support Support `json:"support,omitempty"` + // Enum is the vendor's closed vocabulary, required when Kind is KindEnum. + Enum []EnumOption `json:"enum,omitempty"` + // Min and Max bound a numeric parameter. Nil means unbounded on that side. + Min *float64 `json:"min,omitempty"` + Max *float64 `json:"max,omitempty"` + // Default is what the plugin sends when the caller says nothing. Nil means + // "send nothing and let the model's own default stand", which is different + // from sending the vendor's documented default explicitly: it keeps the + // request minimal and survives the vendor changing that default. A + // non-nil default is therefore a deliberate claim that this vendor needs + // the field on every request — Aliyun's thinking models are the case that + // forces it, since omitting enable_thinking is not neutral there. + Default *Value `json:"default,omitempty"` + // Pin is the forced value for a SupportPinned parameter. + Pin *Value `json:"pin,omitempty"` + // Encode writes the resolved value onto the outbound draft. Nil means the + // parameter is declared for the UI and validation but the protocol driver + // already handles it (temperature on an OpenAI-shaped body, for instance). + Encode Encoder `json:"-"` + // UI is the form presentation. + UI ParamUI `json:"ui"` + // DocURL points at the vendor documentation this declaration follows, so a + // reviewer can check the claim rather than trust it. + DocURL string `json:"doc_url,omitempty"` +} + +// EffectiveSupport reports the parameter's disposition, treating the zero +// value as SupportUser. +func (p Param) EffectiveSupport() Support { + if p.Support == "" { + return SupportUser + } + return p.Support +} + +// EffectiveWidget reports the control to render, inferring one from the +// parameter's kind when the declaration does not override it. +func (p Param) EffectiveWidget() Widget { + if p.UI.Widget != "" { + return p.UI.Widget + } + switch p.Kind { + case KindBool: + return WidgetSwitch + case KindEnum: + return WidgetSelect + case KindFloat: + if p.Min != nil && p.Max != nil { + return WidgetSlider + } + return WidgetNumber + default: + return WidgetNumber + } +} + +// AllowsValue reports whether v is inside the parameter's declared domain. +// Validation lives with the declaration so the API, the form, and the request +// path all reject the same values. +func (p Param) AllowsValue(v Value) bool { + if v.Kind != p.Kind { + return false + } + switch p.Kind { + case KindEnum: + for _, opt := range p.Enum { + if opt.Value == v.Str { + return true + } + } + return false + case KindInt, KindFloat: + if p.Min != nil && v.Num < *p.Min { + return false + } + if p.Max != nil && v.Num > *p.Max { + return false + } + return true + default: + return true + } +} + +// Float returns a pointer to f, for populating Min, Max, and numeric bounds +// without a temporary at every call site. +func Float(f float64) *float64 { return &f } + +// Ptr returns a pointer to v, for populating Default and Pin inline. +func Ptr(v Value) *Value { return &v } diff --git a/internal/models/llm/spi/plan.go b/internal/models/llm/spi/plan.go new file mode 100644 index 00000000000..98caafe7eb0 --- /dev/null +++ b/internal/models/llm/spi/plan.go @@ -0,0 +1,202 @@ +package spi + +import "fmt" + +// NoteReason classifies why a plan differs from what the caller asked for. +type NoteReason string + +const ( + // NoteDefaulted means the plugin supplied a value the caller omitted. + NoteDefaulted NoteReason = "defaulted" + // NotePinned means the plugin overrode the caller because the vendor + // accepts only one value. + NotePinned NoteReason = "pinned" + // NoteForbidden means the parameter was dropped because the vendor rejects + // it for this model. + NoteForbidden NoteReason = "forbidden" + // NoteUnsupported means the parameter is not declared by this plugin, so + // the caller's value has nowhere to go. + NoteUnsupported NoteReason = "unsupported" + // NoteOutOfDomain means the caller's value fell outside the declared + // domain and was dropped rather than sent for the vendor to reject. + NoteOutOfDomain NoteReason = "out_of_domain" + // NoteConstrained means a constraint adjusted or dropped the value. + NoteConstrained NoteReason = "constrained" +) + +// Note records one difference between the request as asked and the request as +// it will be sent. Notes are the plan's honesty mechanism: every silent +// adjustment a vendor quirk forces becomes an inspectable fact instead of +// surprising behavior, and the model debug endpoint reports them verbatim. +type Note struct { + // Param is the affected parameter. + Param ParamID `json:"param"` + // Reason classifies the difference. + Reason NoteReason `json:"reason"` + // Detail explains it in one sentence, for a human reading diagnostics. + Detail string `json:"detail,omitempty"` + // Source names the constraint responsible, when one is. + Source string `json:"source,omitempty"` +} + +// Plan is the resolved outcome of applying a descriptor to one request: the +// parameters that will be sent, their values, and every adjustment made along +// the way. +// +// It exists as a separate step from encoding so the same resolution can be +// asserted in a test, reported by the debug endpoint, and rendered by the UI +// without issuing a request. A plan is what makes "which field will carry my +// thinking toggle" answerable rather than guessable. +type Plan struct { + // Vendor is the descriptor that produced this plan. + Vendor string `json:"vendor"` + // Protocol is the wire protocol the request will use. + Protocol ProtocolID `json:"protocol"` + // Model is the resolved model identifier. + Model string `json:"model"` + // Stream reports whether this is a streaming request. + Stream bool `json:"stream"` + // Notes records every adjustment, in the order it was made. + Notes []Note `json:"notes,omitempty"` + + values map[ParamID]Value + order []ParamID + params map[ParamID]Param + encoded map[ParamID]string + strip []ParamID +} + +// newPlan returns an empty plan for a descriptor and request shape. +func newPlan(vendor string, protocol ProtocolID, model string, stream bool) *Plan { + return &Plan{ + Vendor: vendor, + Protocol: protocol, + Model: model, + Stream: stream, + values: map[ParamID]Value{}, + params: map[ParamID]Param{}, + encoded: map[ParamID]string{}, + } +} + +// Value reports the resolved value of a parameter and whether it is present. +func (p *Plan) Value(id ParamID) (Value, bool) { + v, ok := p.values[id] + return v, ok +} + +// Has reports whether a parameter will be sent. +func (p *Plan) Has(id ParamID) bool { + _, ok := p.values[id] + return ok +} + +// Param reports the declaration behind a resolved parameter, so a constraint +// can consult the domain it must respect. +func (p *Plan) Param(id ParamID) (Param, bool) { + decl, ok := p.params[id] + return decl, ok +} + +// Params reports the resolved parameters in declaration order. +func (p *Plan) Params() []ParamID { + out := make([]ParamID, 0, len(p.order)) + for _, id := range p.order { + if _, ok := p.values[id]; ok { + out = append(out, id) + } + } + return out +} + +// Set records a resolved value, appending it to the order the first time. +func (p *Plan) Set(id ParamID, v Value) { + if _, exists := p.values[id]; !exists { + p.order = append(p.order, id) + } + p.values[id] = v +} + +// Drop removes a resolved value and records why. +func (p *Plan) Drop(id ParamID, reason NoteReason, source, detail string) { + if _, ok := p.values[id]; !ok { + return + } + delete(p.values, id) + p.Note(id, reason, source, detail) +} + +// Adjust replaces a resolved value and records why. +func (p *Plan) Adjust(id ParamID, v Value, source, detail string) { + p.Set(id, v) + p.Note(id, NoteConstrained, source, detail) +} + +// Note records an adjustment without changing any value. +func (p *Plan) Note(id ParamID, reason NoteReason, source, detail string) { + p.Notes = append(p.Notes, Note{Param: id, Reason: reason, Source: source, Detail: detail}) +} + +// EncodedBy reports the encoder id that carried a parameter onto the wire. +// This is the fact the model editor shows in place of the old free-text +// thinking_control setting, and it is derived from the request path rather +// than duplicated from it. +func (p *Plan) EncodedBy(id ParamID) string { + return p.encoded[id] +} + +// Encodings reports every parameter that reached the wire and the encoder that +// carried it, keyed by parameter id. +func (p *Plan) Encodings() map[ParamID]string { + out := make(map[ParamID]string, len(p.encoded)) + for id, enc := range p.encoded { + out[id] = enc + } + return out +} + +// Stripped reports the parameters removed from the wire because the vendor +// forbids them. +func (p *Plan) Stripped() []ParamID { + out := make([]ParamID, len(p.strip)) + copy(out, p.strip) + return out +} + +// Apply writes the plan onto a draft: forbidden fields are removed first, then +// each resolved parameter's encoder runs in declaration order so a later +// encoder can build on an earlier one's output. +// +// Stripping precedes encoding because the protocol driver has already written +// its canonical body, and a vendor that forbids a field needs it gone whether +// or not the caller mentioned it. +func (p *Plan) Apply(d *Draft) error { + for _, id := range p.strip { + decl, ok := p.params[id] + if !ok || decl.Encode == nil { + continue + } + stripper, ok := decl.Encode.(Stripper) + if !ok { + continue + } + if err := stripper.Strip(d); err != nil { + return fmt.Errorf("strip %s via %s: %w", id, decl.Encode.ID(), err) + } + } + for _, id := range p.order { + v, ok := p.values[id] + if !ok { + continue + } + decl, ok := p.params[id] + if !ok || decl.Encode == nil { + continue + } + if err := decl.Encode.Encode(d, v); err != nil { + return fmt.Errorf("encode %s via %s: %w", id, decl.Encode.ID(), err) + } + p.encoded[id] = decl.Encode.ID() + } + return nil +} diff --git a/internal/models/llm/spi/registry.go b/internal/models/llm/spi/registry.go new file mode 100644 index 00000000000..1b81bd37244 --- /dev/null +++ b/internal/models/llm/spi/registry.go @@ -0,0 +1,215 @@ +package spi + +import ( + "fmt" + "sort" + "strings" + "sync" +) + +// Registry holds the registered vendor plugins and resolves one for a +// configured model. +// +// Registration is reversible: Register returns the function that removes the +// descriptor again. That keeps tests from leaking state into one another and +// leaves room for plugins whose lifetime is shorter than the process, which is +// the direction this seam is meant to grow. +type Registry struct { + mu sync.RWMutex + entries []entry + nextSeq int +} + +type entry struct { + desc Descriptor + seq int +} + +// NewRegistry returns an empty registry. +func NewRegistry() *Registry { return &Registry{} } + +// Register validates and adds a descriptor, returning the function that +// removes it. A descriptor that fails validation is rejected outright rather +// than half-registered, so a malformed plugin fails at startup instead of +// producing a request the vendor will reject. +func (r *Registry) Register(desc Descriptor) (func(), error) { + if err := desc.validate(); err != nil { + return nil, err + } + r.mu.Lock() + defer r.mu.Unlock() + seq := r.nextSeq + r.nextSeq++ + r.entries = append(r.entries, entry{desc: desc, seq: seq}) + return func() { r.remove(seq) }, nil +} + +// MustRegister adds a descriptor and panics if it is malformed. Plugin +// packages call it from init, where a descriptor is a compile-time constant of +// the program and a mistake in one is a programming error, not a runtime +// condition worth degrading around. +func (r *Registry) MustRegister(desc Descriptor) func() { + undo, err := r.Register(desc) + if err != nil { + panic(fmt.Sprintf("llm/spi: %v", err)) + } + return undo +} + +func (r *Registry) remove(seq int) { + r.mu.Lock() + defer r.mu.Unlock() + for i, e := range r.entries { + if e.seq == seq { + r.entries = append(r.entries[:i], r.entries[i+1:]...) + return + } + } +} + +// Query selects a descriptor. It is a struct rather than a parameter list +// because the selectors grow: model-specific descriptors and protocol choice +// arrived after the first two, and a caller that ignores one should not have +// to name it. +type Query struct { + // Vendor is the provider identity, matched case-insensitively. + Vendor string + // Kind is the capability required. + Kind ModelKind + // Model is the model name, matched against each descriptor's matcher. + Model string + // Protocol optionally pins the wire protocol, for vendors offering more + // than one. Empty accepts whichever the vendor registered first, which is + // its documented default. + Protocol ProtocolID +} + +// Resolve returns the descriptor handling a query. +// +// A vendor may register several descriptors for one kind; the one whose model +// matcher is specific wins over the vendor's catch-all, and among equally +// specific matchers the earliest registration wins. That ordering is what lets +// "OpenAI reasoning models" be a separate declaration from "OpenAI models" +// without either one knowing about the other. +func (r *Registry) Resolve(q Query) (Descriptor, bool) { + r.mu.RLock() + defer r.mu.RUnlock() + + name := strings.ToLower(strings.TrimSpace(q.Vendor)) + var best *entry + for i := range r.entries { + e := &r.entries[i] + if strings.ToLower(e.desc.Vendor) != name || e.desc.Kind != q.Kind { + continue + } + if q.Protocol != "" && e.desc.Protocol != q.Protocol { + continue + } + if !e.desc.Models.Matches(q.Model) { + continue + } + if best == nil || better(e, best) { + best = e + } + } + if best == nil { + return Descriptor{}, false + } + return best.desc, true +} + +// Protocols reports the distinct protocols a vendor offers for a kind, in +// registration order. The model editor uses it to decide whether to show a +// protocol choice at all: one entry means there is nothing to choose. +func (r *Registry) Protocols(vendor string, kind ModelKind) []ProtocolID { + r.mu.RLock() + defer r.mu.RUnlock() + + name := strings.ToLower(strings.TrimSpace(vendor)) + seen := map[ProtocolID]struct{}{} + var out []ProtocolID + for _, e := range r.entries { + if strings.ToLower(e.desc.Vendor) != name || e.desc.Kind != kind { + continue + } + if _, dup := seen[e.desc.Protocol]; dup { + continue + } + seen[e.desc.Protocol] = struct{}{} + out = append(out, e.desc.Protocol) + } + return out +} + +// better reports whether candidate outranks incumbent: specific matchers beat +// catch-alls, then earlier registration wins. +func better(candidate, incumbent *entry) bool { + candidateSpecific := !candidate.desc.Models.IsCatchAll() + incumbentSpecific := !incumbent.desc.Models.IsCatchAll() + if candidateSpecific != incumbentSpecific { + return candidateSpecific + } + return candidate.seq < incumbent.seq +} + +// List reports every descriptor registered for a kind, ordered by vendor and +// then registration. It backs the capability catalog the model editor renders. +func (r *Registry) List(kind ModelKind) []Descriptor { + r.mu.RLock() + defer r.mu.RUnlock() + + out := make([]entry, 0, len(r.entries)) + for _, e := range r.entries { + if e.desc.Kind == kind { + out = append(out, e) + } + } + sort.SliceStable(out, func(i, j int) bool { + if out[i].desc.Vendor != out[j].desc.Vendor { + return out[i].desc.Vendor < out[j].desc.Vendor + } + return out[i].seq < out[j].seq + }) + descs := make([]Descriptor, 0, len(out)) + for _, e := range out { + descs = append(descs, e.desc) + } + return descs +} + +// Vendors reports the distinct vendors registered for a kind. +func (r *Registry) Vendors(kind ModelKind) []string { + seen := map[string]struct{}{} + var out []string + for _, desc := range r.List(kind) { + if _, dup := seen[desc.Vendor]; dup { + continue + } + seen[desc.Vendor] = struct{}{} + out = append(out, desc.Vendor) + } + return out +} + +// Default is the process-wide registry that vendor packages register into from +// their init functions. +var Default = NewRegistry() + +// Register adds a descriptor to the default registry. +func Register(desc Descriptor) (func(), error) { return Default.Register(desc) } + +// MustRegister adds a descriptor to the default registry, panicking if it is +// malformed. +func MustRegister(desc Descriptor) func() { return Default.MustRegister(desc) } + +// Resolve looks up a descriptor in the default registry. +func Resolve(q Query) (Descriptor, bool) { return Default.Resolve(q) } + +// List reports the default registry's descriptors for a kind. +func List(kind ModelKind) []Descriptor { return Default.List(kind) } + +// Protocols reports the protocols a vendor offers for a kind in the default +// registry. +func Protocols(vendor string, kind ModelKind) []ProtocolID { + return Default.Protocols(vendor, kind) +} diff --git a/internal/models/llm/spi/resolve.go b/internal/models/llm/spi/resolve.go new file mode 100644 index 00000000000..588d0cd304a --- /dev/null +++ b/internal/models/llm/spi/resolve.go @@ -0,0 +1,119 @@ +package spi + +import "fmt" + +// Request is one call's neutral intent: which model, whether it streams, and +// the parameter values the caller wants. Callers never name wire fields, which +// is what lets the same request run against any vendor. +type Request struct { + // Model is the model identifier to send. + Model string + // Stream reports whether the response is streamed. + Stream bool + // Values are the caller's parameter values, keyed by neutral id. A missing + // key means "no opinion", which is distinct from an explicit zero. + Values map[ParamID]Value +} + +// Value reports a caller value and whether it was supplied. +func (r Request) Value(id ParamID) (Value, bool) { + v, ok := r.Values[id] + return v, ok +} + +// WithValue returns a copy of the request carrying an additional value. +func (r Request) WithValue(id ParamID, v Value) Request { + values := make(map[ParamID]Value, len(r.Values)+1) + for k, existing := range r.Values { + values[k] = existing + } + values[id] = v + r.Values = values + return r +} + +// Plan resolves a request against the descriptor, producing the inspectable +// outcome without touching the wire. +// +// Resolution is deliberately a separate step from encoding: the same plan +// answers "what will actually be sent" for a test, for the debug endpoint, and +// for the model editor, so none of them has to reimplement the precedence +// rules and then drift from the request path. +// +// Precedence per parameter is pin, then caller, then plugin default, then +// omit. Every departure from the caller's request is recorded as a note. +func (d Descriptor) Plan(req Request) (*Plan, error) { + model := req.Model + if model == "" { + model = d.Vendor + } + plan := newPlan(d.Vendor, d.Protocol, model, req.Stream) + + declared := make(map[ParamID]struct{}, len(d.Params)) + for _, p := range d.Params { + declared[p.ID] = struct{}{} + plan.params[p.ID] = p + resolveParam(plan, p, req) + } + + // A caller value with nowhere to go is worth reporting rather than + // dropping in silence: it is usually a UI offering a control the resolved + // vendor does not actually have. + for id := range req.Values { + if _, ok := declared[id]; !ok { + plan.Note(id, NoteUnsupported, d.Vendor, + fmt.Sprintf("%s does not declare %s for %s models", d.Vendor, id, d.Kind)) + } + } + + for _, constraint := range d.Constraints { + if err := constraint.Apply(plan); err != nil { + return nil, fmt.Errorf("constraint %s: %w", constraint.ID(), err) + } + } + return plan, nil +} + +// resolveParam applies the precedence rules for one declared parameter. +func resolveParam(plan *Plan, p Param, req Request) { + given, supplied := req.Value(p.ID) + + switch p.EffectiveSupport() { + case SupportForbidden: + plan.strip = append(plan.strip, p.ID) + if supplied { + plan.Note(p.ID, NoteForbidden, plan.Vendor, + fmt.Sprintf("%s rejects %s for model %s; the field is removed from the request", + plan.Vendor, p.ID, plan.Model)) + } + return + + case SupportPinned: + plan.Set(p.ID, *p.Pin) + if supplied && given != *p.Pin { + plan.Note(p.ID, NotePinned, plan.Vendor, + fmt.Sprintf("%s accepts only %s=%s for model %s; the requested %s was replaced", + plan.Vendor, p.ID, p.Pin, plan.Model, given)) + } + return + } + + if supplied { + if p.AllowsValue(given) { + plan.Set(p.ID, given) + return + } + plan.Note(p.ID, NoteOutOfDomain, plan.Vendor, + fmt.Sprintf("%s is outside the values %s accepts for %s", given, plan.Vendor, p.ID)) + // Fall through: a rejected value should not silently become the + // plugin default, so only an explicitly declared default applies. + } + + if p.Default != nil { + plan.Set(p.ID, *p.Default) + if !supplied { + plan.Note(p.ID, NoteDefaulted, plan.Vendor, + fmt.Sprintf("%s requires %s on every request; sending %s", plan.Vendor, p.ID, p.Default)) + } + } +} diff --git a/internal/models/llm/spi/value.go b/internal/models/llm/spi/value.go new file mode 100644 index 00000000000..45aa8519dd3 --- /dev/null +++ b/internal/models/llm/spi/value.go @@ -0,0 +1,141 @@ +package spi + +import ( + "fmt" + "math" + "strconv" + "strings" +) + +// ValueKind is the domain of a parameter value. Every request parameter a +// vendor exposes is one of these four shapes, which is what lets one +// declaration serve the wire, the validator, and the form widget. +type ValueKind string + +const ( + // KindBool is an on/off parameter rendered as a switch. + KindBool ValueKind = "bool" + // KindInt is an integer parameter rendered as a number input. + KindInt ValueKind = "int" + // KindFloat is a fractional parameter rendered as a slider or number input. + KindFloat ValueKind = "float" + // KindEnum is a closed vocabulary rendered as a select. The vocabulary is + // the VENDOR's, not a normalized one: OpenAI's reasoning effort ladder and + // Zhipu's are different sets, and pretending otherwise would silently send + // a value the vendor rejects. + KindEnum ValueKind = "enum" +) + +// Value is a parameter value tagged with its kind. It is deliberately a small +// concrete struct rather than `any`: encoders, validators, and the UI schema +// all need to know the shape without a type switch on every access. +type Value struct { + Kind ValueKind `json:"kind"` + Bool bool `json:"bool,omitempty"` + Num float64 `json:"num,omitempty"` + Str string `json:"str,omitempty"` +} + +// Bool returns a boolean value. +func BoolValue(v bool) Value { return Value{Kind: KindBool, Bool: v} } + +// IntValue returns an integer value. +func IntValue(v int) Value { return Value{Kind: KindInt, Num: float64(v)} } + +// FloatValue returns a fractional value. +func FloatValue(v float64) Value { return Value{Kind: KindFloat, Num: v} } + +// EnumValue returns a value drawn from a vendor's closed vocabulary. +func EnumValue(v string) Value { return Value{Kind: KindEnum, Str: v} } + +// Int reports the value as an integer, rounding a float toward zero. The +// second result is false when the value is not numeric. +func (v Value) Int() (int, bool) { + if v.Kind != KindInt && v.Kind != KindFloat { + return 0, false + } + return int(math.Trunc(v.Num)), true +} + +// Float reports the value as a float. The second result is false when the +// value is not numeric. +func (v Value) Float() (float64, bool) { + if v.Kind != KindInt && v.Kind != KindFloat { + return 0, false + } + return v.Num, true +} + +// Enum reports the value as a vocabulary string. The second result is false +// when the value is not an enum. +func (v Value) Enum() (string, bool) { + if v.Kind != KindEnum { + return "", false + } + return v.Str, true +} + +// JSON reports the value as it should appear in a JSON request body. +func (v Value) JSON() any { + switch v.Kind { + case KindBool: + return v.Bool + case KindInt: + return int(math.Trunc(v.Num)) + case KindFloat: + return v.Num + case KindEnum: + return v.Str + default: + return nil + } +} + +// String renders the value for logs and diagnostics. +func (v Value) String() string { + switch v.Kind { + case KindBool: + return strconv.FormatBool(v.Bool) + case KindInt: + return strconv.Itoa(int(math.Trunc(v.Num))) + case KindFloat: + return strconv.FormatFloat(v.Num, 'g', -1, 64) + case KindEnum: + return v.Str + default: + return "" + } +} + +// ParseValue reads a value of the given kind from its string form, which is +// how values arrive from stored model configuration and from HTTP requests. +func ParseValue(kind ValueKind, raw string) (Value, error) { + raw = strings.TrimSpace(raw) + switch kind { + case KindBool: + b, err := strconv.ParseBool(raw) + if err != nil { + return Value{}, fmt.Errorf("parse bool %q: %w", raw, err) + } + return BoolValue(b), nil + case KindInt: + n, err := strconv.Atoi(raw) + if err != nil { + return Value{}, fmt.Errorf("parse int %q: %w", raw, err) + } + return IntValue(n), nil + case KindFloat: + f, err := strconv.ParseFloat(raw, 64) + if err != nil { + return Value{}, fmt.Errorf("parse float %q: %w", raw, err) + } + return FloatValue(f), nil + case KindEnum: + if raw == "" { + return Value{}, fmt.Errorf("parse enum: empty value") + } + return EnumValue(raw), nil + default: + return Value{}, fmt.Errorf("unknown value kind %q", kind) + } +} diff --git a/internal/models/llm/vendors/anthropic.go b/internal/models/llm/vendors/anthropic.go new file mode 100644 index 00000000000..4bad1b79384 --- /dev/null +++ b/internal/models/llm/vendors/anthropic.go @@ -0,0 +1,149 @@ +package vendors + +import ( + "github.com/Tencent/WeKnora/internal/models/llm/encoding" + "github.com/Tencent/WeKnora/internal/models/llm/spi" +) + +// Anthropic Claude on the Messages protocol. +// +// Claude has two thinking modes and they are not interchangeable, which is why +// this file registers two descriptors instead of one with a flag: +// +// - Manual (extended) thinking, `{"thinking": {"type": "enabled", +// "budget_tokens": N}}`, is the only mode on Claude 4.5 and earlier. It is +// deprecated on the 4.6 generation and returns 400 on 4.7 and later. +// - Adaptive thinking, `{"thinking": {"type": "adaptive"}}` with depth set by +// `output_config.effort`, is available from 4.6 onward and rejects +// budget_tokens. +// +// Sending the wrong one is a hard error rather than a degraded response, so the +// model matcher is doing real work here. +// +// Docs: https://platform.claude.com/docs/en/build-with-claude/thinking +// +// https://platform.claude.com/docs/en/build-with-claude/extended-thinking +// https://platform.claude.com/docs/en/build-with-claude/effort +const ( + anthropicBaseURL = "https://api.anthropic.com" + anthropicVersion = "2023-06-01" + anthropicThinkingDoc = "https://platform.claude.com/docs/en/build-with-claude/thinking" + anthropicExtendedDoc = "https://platform.claude.com/docs/en/build-with-claude/extended-thinking" + anthropicEffortDoc = "https://platform.claude.com/docs/en/build-with-claude/effort" + + // anthropicMinBudget is the documented floor; the API rejects less. + anthropicMinBudget = 1024 + // anthropicAnswerHeadroom is the room left for the answer when a budget + // forces the output ceiling up. Thinking tokens count against max_tokens, + // so a ceiling equal to the budget would leave nothing to answer with. + anthropicAnswerHeadroom = 4096 + // anthropicDefaultMaxTokens is the fallback ceiling, since the Messages + // API requires max_tokens on every request unlike the OpenAI-shaped + // protocols where it is optional. + anthropicDefaultMaxTokens = 1024 +) + +// anthropicAdaptiveModels matches the generations that support adaptive +// thinking: 4.6 and later, and the 5 series. Claude 4.5 and earlier, and the +// older claude-3-x names, fall through to the manual descriptor. +var anthropicAdaptiveModels = spi.ModelMatcher{ + Pattern: `^claude-[a-z]+-(4-[6-9]|[5-9])`, +} + +// anthropicAuth is shared by both descriptors: the key travels in x-api-key, +// and the version header is mandatory on every request. +var anthropicAuth = spi.Auth{ + Kind: spi.AuthHeader, + Header: "x-api-key", + Static: map[string]string{"anthropic-version": anthropicVersion}, +} + +func init() { + // Adaptive thinking: Claude decides whether and how deeply to think, and + // effort is the only depth control. Both "on" and "auto" map to adaptive + // because these models have no always-think mode — asking for more + // thinking means raising the effort, not pinning a budget. + spi.MustRegister(spi.Descriptor{ + Vendor: "anthropic", + Kind: spi.KindChat, + Protocol: spi.ProtocolAnthropicMessages, + DisplayName: "Anthropic Claude (adaptive thinking)", + Models: anthropicAdaptiveModels, + + DefaultBaseURL: anthropicBaseURL, + Auth: anthropicAuth, + Params: []spi.Param{ + encoding.ThinkingMode( + encoding.ThinkingObject{Key: "thinking", On: "adaptive", Off: "disabled", Auto: "adaptive"}, + thinkingModes(true)...), + // Effort lives in output_config, not in the thinking object. The + // documentation calls this out explicitly, and "adaptive" is a + // thinking mode rather than an effort rung. + anthropicEffortParam(), + encoding.Temperature(0, 1), + encoding.TopP(), + encoding.MaxTokens(), + }, + Constraints: []spi.Constraint{ + encoding.RequireMaxTokens{Param: spi.ParamMaxTokens, Fallback: anthropicDefaultMaxTokens}, + }, + ReasoningReplay: spi.ReplayAlways, + DocURL: anthropicThinkingDoc, + }) + + // Manual extended thinking: a fixed token budget, with the constraints the + // documentation states — a 1,024 floor and a budget strictly below the + // output ceiling. + spi.MustRegister(spi.Descriptor{ + Vendor: "anthropic", + Kind: spi.KindChat, + Protocol: spi.ProtocolAnthropicMessages, + DisplayName: "Anthropic Claude (extended thinking)", + + DefaultBaseURL: anthropicBaseURL, + Auth: anthropicAuth, + Params: []spi.Param{ + encoding.ThinkingMode( + encoding.ThinkingObject{Key: "thinking", On: "enabled", Off: "disabled"}, + thinkingModes(false)...), + anthropicBudgetParam(), + encoding.Temperature(0, 1), + encoding.TopP(), + encoding.MaxTokens(), + }, + Constraints: []spi.Constraint{ + // Order matters: drop a budget that does not apply, then supply the + // required ceiling, then raise it if the surviving budget needs + // more room. + encoding.DependsOnThinking{Params: []spi.ParamID{spi.ParamThinkingBudget}}, + encoding.RequireMaxTokens{Param: spi.ParamMaxTokens, Fallback: anthropicDefaultMaxTokens}, + encoding.BudgetBelowMaxTokens{ + Budget: spi.ParamThinkingBudget, + MaxTokens: spi.ParamMaxTokens, + Headroom: anthropicAnswerHeadroom, + }, + }, + ReasoningReplay: spi.ReplayAlways, + DocURL: anthropicExtendedDoc, + }) +} + +// anthropicEffortParam declares the effort ladder at output_config.effort. +// It is declared only on the adaptive descriptor: the effort documentation +// says the parameter is broadly available, while the extended-thinking page +// names Opus 4.5 as the only manual-mode model supporting it. Where the two +// disagree, sending nothing is the safe reading — an unsupported field is a +// 400, whereas omitting it costs only the vendor's own default. +func anthropicEffortParam() spi.Param { + p := encoding.ThinkingEffort(encoding.Nested("output_config", "effort"), + "low", "medium", "high", "xhigh", "max") + p.DocURL = anthropicEffortDoc + return p +} + +// anthropicBudgetParam declares the manual thinking budget. +func anthropicBudgetParam() spi.Param { + p := encoding.ThinkingBudget(encoding.Nested("thinking", "budget_tokens"), anthropicMinBudget, 0) + p.DocURL = anthropicExtendedDoc + return p +} diff --git a/internal/models/llm/vendors/cn_vendors.go b/internal/models/llm/vendors/cn_vendors.go new file mode 100644 index 00000000000..4685883e0f9 --- /dev/null +++ b/internal/models/llm/vendors/cn_vendors.go @@ -0,0 +1,317 @@ +package vendors + +import ( + "github.com/Tencent/WeKnora/internal/models/llm/encoding" + "github.com/Tencent/WeKnora/internal/models/llm/spi" +) + +// The vendors below all serve an OpenAI-compatible Chat Completions endpoint +// and all support reasoning, yet no two spell it the same way. Collected here, +// the differences are easy to compare against the documentation each entry +// links; scattered across if-else branches, they were the reason this seam +// exists. + +const ( + deepSeekBaseURL = "https://api.deepseek.com" + deepSeekThinkDoc = "https://api-docs.deepseek.com/guides/thinking_mode" + aliyunBaseURL = "https://dashscope.aliyuncs.com/compatible-mode/v1" + aliyunThinkDoc = "https://help.aliyun.com/zh/model-studio/deep-thinking" + volcengineBaseURL = "https://ark.cn-beijing.volces.com/api/v3" + volcengineDoc = "https://www.volcengine.com/docs/82379/1494384" + zhipuBaseURL = "https://open.bigmodel.cn/api/paas/v4" + zhipuThinkDoc = "https://docs.bigmodel.cn/cn/guide/capabilities/thinking" + lkeapDoc = "https://cloud.tencent.com/document/product/1772/115963" +) + +func init() { + registerDeepSeek() + registerAliyun() + registerVolcengine() + registerZhipu() + registerMoonshot() + registerLKEAP() + registerSelfHosted() +} + +// DeepSeek toggles thinking with the `thinking` object and sizes it with a +// three-rung reasoning_effort ladder. Its distinguishing requirement is on the +// response side: when a turn carries tool calls, that turn's reasoning_content +// must be replayed in every later request or the API answers 400. +// +// Docs: https://api-docs.deepseek.com/guides/thinking_mode +func registerDeepSeek() { + spi.MustRegister(spi.Descriptor{ + Vendor: "deepseek", + Kind: spi.KindChat, + Protocol: spi.ProtocolOpenAIChat, + DisplayName: "DeepSeek", + + DefaultBaseURL: deepSeekBaseURL, + Auth: spi.Auth{Kind: spi.AuthBearer}, + Params: append([]spi.Param{ + encoding.ThinkingMode(thinkingObject(""), thinkingModes(false)...), + effortParam(encoding.Field{Key: "reasoning_effort"}, deepSeekThinkDoc, "low", "high", "max"), + // Preserved from the previous implementation, which stripped + // tool_choice for this vendor. Declaring it forbidden keeps that + // behavior while making it visible: the request path removes the + // field and the plan records that it did. + encoding.Forbidden(spi.ParamToolChoice, spi.KindEnum, "tool_choice"), + }, encoding.SamplingSet(2)...), + Constraints: []spi.Constraint{ + encoding.DependsOnThinking{Params: []spi.ParamID{spi.ParamThinkingEffort}}, + }, + ReasoningReplay: spi.ReplayWithTools, + DocURL: deepSeekThinkDoc, + }) +} + +// Alibaba Cloud Model Studio uses a plain boolean plus a token budget, both as +// top-level fields. Two properties make it the awkward one: +// +// - The boolean must be sent on every request for hybrid-thinking models, +// because omitting it is not neutral — several Qwen generations default to +// thinking on, so silence means the opposite of what a caller expects. +// - Some open-source Qwen3 models accept thinking only in streaming mode and +// error otherwise, which the stream-only constraint downgrades locally. +// +// Docs: https://help.aliyun.com/zh/model-studio/deep-thinking +func registerAliyun() { + thinkingModeParam := encoding.ThinkingMode( + encoding.EnableThinkingBool{Key: "enable_thinking"}, thinkingModes(false)...) + thinkingModeParam.Default = spi.Ptr(spi.EnumValue(spi.ThinkingOff)) + thinkingModeParam.DocURL = aliyunThinkDoc + + budget := encoding.ThinkingBudget(encoding.Field{Key: "thinking_budget"}, 1, 0) + budget.DocURL = aliyunThinkDoc + + // Qwen thinking-capable families. Non-thinking Aliyun models fall through + // to the catch-all descriptor, which sends no thinking fields at all. + spi.MustRegister(spi.Descriptor{ + Vendor: "aliyun", + Kind: spi.KindChat, + Protocol: spi.ProtocolOpenAIChat, + DisplayName: "Alibaba Cloud Model Studio (thinking)", + Models: spi.ModelMatcher{ + Prefixes: []string{"qwen3", "qwen-plus", "qwen-max", "qwen-turbo", "qwq", "qvq"}, + }, + + DefaultBaseURL: aliyunBaseURL, + Auth: spi.Auth{Kind: spi.AuthBearer}, + Params: append([]spi.Param{thinkingModeParam, budget}, encoding.SamplingSet(2)...), + Constraints: []spi.Constraint{ + encoding.StreamOnlyThinking{}, + encoding.DependsOnThinking{Params: []spi.ParamID{spi.ParamThinkingBudget}}, + }, + ReasoningReplay: spi.ReplayWithTools, + DocURL: aliyunThinkDoc, + }) + + spi.MustRegister(spi.Descriptor{ + Vendor: "aliyun", + Kind: spi.KindChat, + Protocol: spi.ProtocolOpenAIChat, + DisplayName: "Alibaba Cloud Model Studio", + DefaultBaseURL: aliyunBaseURL, + Auth: spi.Auth{Kind: spi.AuthBearer}, + Params: encoding.SamplingSet(2), + DocURL: "https://help.aliyun.com/zh/model-studio/compatibility-of-openai-with-dashscope", + }) +} + +// Volcengine Ark is the one vendor documenting a third thinking state: `auto` +// lets the model skip reasoning on simple questions. A boolean toggle cannot +// express that, which is why the neutral mode is an enum. +// +// Docs: https://www.volcengine.com/docs/82379/1494384 +func registerVolcengine() { + spi.MustRegister(spi.Descriptor{ + Vendor: "volcengine", + Kind: spi.KindChat, + Protocol: spi.ProtocolOpenAIChat, + DisplayName: "Volcengine Ark", + + DefaultBaseURL: volcengineBaseURL, + Auth: spi.Auth{Kind: spi.AuthBearer}, + Params: append([]spi.Param{ + encoding.ThinkingMode(thinkingObject("auto"), thinkingModes(true)...), + budgetParam(encoding.Nested("thinking", "budget_tokens"), volcengineDoc), + }, encoding.SamplingSet(2)...), + Constraints: []spi.Constraint{ + encoding.DependsOnThinking{Params: []spi.ParamID{spi.ParamThinkingBudget}}, + }, + ReasoningReplay: spi.ReplayWithTools, + DocURL: volcengineDoc, + }) +} + +// Zhipu GLM shares the `thinking` object with DeepSeek but publishes a +// seven-rung effort ladder, and only from GLM-5.2 onward. Declaring the ladder +// on a matcher rather than on every GLM keeps it off the models that would +// reject it. +// +// Docs: https://docs.bigmodel.cn/cn/guide/capabilities/thinking +func registerZhipu() { + zhipuLadder := []string{"none", "minimal", "low", "medium", "high", "xhigh", "max"} + + spi.MustRegister(spi.Descriptor{ + Vendor: "zhipu", + Kind: spi.KindChat, + Protocol: spi.ProtocolOpenAIChat, + DisplayName: "Zhipu GLM (effort)", + // GLM-5.2 and later support reasoning_effort; earlier thinking models + // support only the on/off object. + Models: spi.ModelMatcher{Pattern: `^glm-(5\.[2-9]|[6-9])`}, + + DefaultBaseURL: zhipuBaseURL, + Auth: spi.Auth{Kind: spi.AuthBearer}, + Params: append([]spi.Param{ + encoding.ThinkingMode(thinkingObject(""), thinkingModes(false)...), + effortParam(encoding.Field{Key: "reasoning_effort"}, zhipuThinkDoc, zhipuLadder...), + }, encoding.SamplingSet(1)...), + Constraints: []spi.Constraint{ + encoding.DependsOnThinking{Params: []spi.ParamID{spi.ParamThinkingEffort}}, + }, + ReasoningReplay: spi.ReplayWithTools, + DocURL: zhipuThinkDoc, + }) + + spi.MustRegister(spi.Descriptor{ + Vendor: "zhipu", + Kind: spi.KindChat, + Protocol: spi.ProtocolOpenAIChat, + DisplayName: "Zhipu GLM", + Models: spi.ModelMatcher{Prefixes: []string{"glm-4.5", "glm-4.6", "glm-4.7", "glm-5"}}, + + DefaultBaseURL: zhipuBaseURL, + Auth: spi.Auth{Kind: spi.AuthBearer}, + Params: append([]spi.Param{ + encoding.ThinkingMode(thinkingObject(""), thinkingModes(false)...), + }, encoding.SamplingSet(1)...), + ReasoningReplay: spi.ReplayWithTools, + DocURL: zhipuThinkDoc, + }) + + spi.MustRegister(spi.Descriptor{ + Vendor: "zhipu", + Kind: spi.KindChat, + Protocol: spi.ProtocolOpenAIChat, + DisplayName: "Zhipu", + DefaultBaseURL: zhipuBaseURL, + Auth: spi.Auth{Kind: spi.AuthBearer}, + Params: encoding.SamplingSet(1), + DocURL: "https://docs.bigmodel.cn/cn/guide/develop/openai/introduction", + }) +} + +// Moonshot's v1 line accepts only temperature=1. Pinning it is a parameter +// disposition rather than request-shaping code, and the plan records the +// override so a user who set 0.2 can see what happened to it. +func registerMoonshot() { + pinnedTemp := encoding.Pinned(spi.ParamTemperature, spi.KindFloat, + spi.FloatValue(1), encoding.Field{Key: "temperature"}) + + spi.MustRegister(spi.Descriptor{ + Vendor: "moonshot", + Kind: spi.KindChat, + Protocol: spi.ProtocolOpenAIChat, + DisplayName: "Moonshot Kimi (fixed temperature)", + Models: spi.ModelMatcher{Prefixes: []string{"kimi-k2", "kimi-latest", "moonshot-v1"}}, + + DefaultBaseURL: "https://api.moonshot.cn/v1", + Auth: spi.Auth{Kind: spi.AuthBearer}, + Params: []spi.Param{ + pinnedTemp, + encoding.MaxTokens(), + }, + DocURL: "https://platform.moonshot.cn/docs/api/chat", + }) + + spi.MustRegister(spi.Descriptor{ + Vendor: "moonshot", + Kind: spi.KindChat, + Protocol: spi.ProtocolOpenAIChat, + DisplayName: "Moonshot Kimi", + DefaultBaseURL: "https://api.moonshot.cn/v1", + Auth: spi.Auth{Kind: spi.AuthBearer}, + Params: encoding.SamplingSet(1), + DocURL: "https://platform.moonshot.cn/docs/api/chat", + }) +} + +// Tencent LKEAP serves DeepSeek models. The V3 line takes the `thinking` +// object; the R1 line reasons unconditionally and rejects the toggle, so it +// falls through to a descriptor that sends none. +// +// Docs: https://cloud.tencent.com/document/product/1772/115963 +func registerLKEAP() { + spi.MustRegister(spi.Descriptor{ + Vendor: "lkeap", + Kind: spi.KindChat, + Protocol: spi.ProtocolOpenAIChat, + DisplayName: "Tencent LKEAP (DeepSeek V3)", + Models: spi.ModelMatcher{Contains: []string{"deepseek-v3"}}, + + Auth: spi.Auth{Kind: spi.AuthBearer}, + Params: append([]spi.Param{ + encoding.ThinkingMode(thinkingObject(""), thinkingModes(false)...), + }, encoding.SamplingSet(2)...), + ReasoningReplay: spi.ReplayWithTools, + DocURL: lkeapDoc, + }) + + spi.MustRegister(spi.Descriptor{ + Vendor: "lkeap", + Kind: spi.KindChat, + Protocol: spi.ProtocolOpenAIChat, + DisplayName: "Tencent LKEAP", + + Auth: spi.Auth{Kind: spi.AuthBearer}, + Params: encoding.SamplingSet(2), + ReasoningReplay: spi.ReplayWithTools, + DocURL: lkeapDoc, + }) +} + +// Self-hosted inference servers pass the flag through to the model's chat +// template rather than interpreting it themselves, so the field lives in +// chat_template_kwargs. This is the right default for a vLLM or SGLang +// deployment and for NVIDIA NIM, and the wrong one for any hosted vendor API — +// which is exactly the distinction a per-vendor declaration can make and a +// single global default could not. +func registerSelfHosted() { + for _, v := range []struct{ vendor, display, doc string }{ + {"generic", "OpenAI-compatible deployment", "https://docs.vllm.ai/en/latest/features/reasoning_outputs.html"}, + {"nvidia", "NVIDIA NIM", "https://docs.nvidia.com/nim/large-language-models/latest/reasoning.html"}, + {"gpustack", "GPUStack", "https://docs.gpustack.ai"}, + } { + spi.MustRegister(spi.Descriptor{ + Vendor: v.vendor, + Kind: spi.KindChat, + Protocol: spi.ProtocolOpenAIChat, + DisplayName: v.display, + + Auth: spi.Auth{Kind: spi.AuthBearer}, + Params: append([]spi.Param{ + encoding.ThinkingMode( + encoding.ChatTemplateKwargs{Key: "chat_template_kwargs", Arg: "enable_thinking"}, + thinkingModes(false)...), + }, encoding.SamplingSet(2)...), + ReasoningReplay: spi.ReplayWithTools, + DocURL: v.doc, + }) + } +} + +// effortParam declares a reasoning ladder with its documentation link. +func effortParam(encoder spi.Encoder, doc string, values ...string) spi.Param { + p := encoding.ThinkingEffort(encoder, values...) + p.DocURL = doc + return p +} + +// budgetParam declares a reasoning-token cap with its documentation link. +func budgetParam(encoder spi.Encoder, doc string) spi.Param { + p := encoding.ThinkingBudget(encoder, 1, 0) + p.DocURL = doc + return p +} diff --git a/internal/models/llm/vendors/openai.go b/internal/models/llm/vendors/openai.go new file mode 100644 index 00000000000..bf10ceaa075 --- /dev/null +++ b/internal/models/llm/vendors/openai.go @@ -0,0 +1,154 @@ +package vendors + +import ( + "github.com/Tencent/WeKnora/internal/models/llm/encoding" + "github.com/Tencent/WeKnora/internal/models/llm/spi" +) + +// OpenAI, and the Azure deployment of it. +// +// Two facts drive these declarations. First, the reasoning models reject the +// sampling parameters the general models accept and require +// max_completion_tokens in place of max_tokens, so they are a separate +// descriptor rather than a branch inside one. Second, reasoning depth is a +// model-dependent ladder: the documentation lists none, minimal, low, medium, +// high, xhigh, and max, and says which rungs exist varies by model, so the +// declaration offers the full published ladder and lets the vendor reject a +// rung its model does not implement. +// +// Docs: https://developers.openai.com/api/docs/guides/reasoning + +// openAIEffortLadder is the published reasoning-effort vocabulary. Depth and +// the on/off toggle are the same field here, which is why `none` doubles as +// the off switch through encoding.EffortNone. +var openAIEffortLadder = []string{"minimal", "low", "medium", "high", "xhigh", "max"} + +const ( + openAIBaseURL = "https://api.openai.com/v1" + openAIReasoningDoc = "https://developers.openai.com/api/docs/guides/reasoning" +) + +// openAIReasoningModels matches the o-series and GPT-5 and later families, +// which is the boundary the API itself draws for the restricted parameter set. +var openAIReasoningModels = spi.ModelMatcher{ + Pattern: `^(o[1-9]|gpt-[5-9]|gpt-1[0-9])`, +} + +// reasoningEffortParam declares the depth ladder on the Chat Completions +// spelling of the field. +func reasoningEffortParam(values ...string) spi.Param { + p := encoding.ThinkingEffort(encoding.Field{Key: "reasoning_effort"}, values...) + p.DocURL = openAIReasoningDoc + return p +} + +// openAIReasoningParams is the parameter set shared by OpenAI and Azure +// reasoning deployments: no sampling knobs, a renamed output ceiling, and the +// effort ladder doubling as the thinking toggle. +func openAIReasoningParams() []spi.Param { + effort := encoding.Field{Key: "reasoning_effort"} + return []spi.Param{ + encoding.ThinkingMode(encoding.EffortNone{Encoder: effort, Off: "none"}, + spi.ThinkingOn, spi.ThinkingOff), + reasoningEffortParam(openAIEffortLadder...), + // The reasoning models reject these outright rather than ignoring + // them, so they must leave the body entirely. + encoding.Forbidden(spi.ParamTemperature, spi.KindFloat, "temperature"), + encoding.Forbidden(spi.ParamTopP, spi.KindFloat, "top_p"), + encoding.Forbidden(spi.ParamFrequencyPenalty, spi.KindFloat, "frequency_penalty"), + encoding.Forbidden(spi.ParamPresencePenalty, spi.KindFloat, "presence_penalty"), + encoding.MaxTokensAs("max_completion_tokens"), + } +} + +// openAIReasoningConstraints keeps the toggle and the ladder from fighting +// over the single field they share. +func openAIReasoningConstraints() []spi.Constraint { + return []spi.Constraint{ + encoding.DependsOnThinking{Params: []spi.ParamID{spi.ParamThinkingEffort}}, + } +} + +func init() { + // OpenAI reasoning models on Chat Completions. + spi.MustRegister(spi.Descriptor{ + Vendor: "openai", + Kind: spi.KindChat, + Protocol: spi.ProtocolOpenAIChat, + DisplayName: "OpenAI (reasoning)", + Models: openAIReasoningModels, + + DefaultBaseURL: openAIBaseURL, + Auth: spi.Auth{Kind: spi.AuthBearer}, + Params: openAIReasoningParams(), + Constraints: openAIReasoningConstraints(), + DocURL: openAIReasoningDoc, + }) + + // OpenAI general models on Chat Completions. + spi.MustRegister(spi.Descriptor{ + Vendor: "openai", + Kind: spi.KindChat, + Protocol: spi.ProtocolOpenAIChat, + DisplayName: "OpenAI", + + DefaultBaseURL: openAIBaseURL, + Auth: spi.Auth{Kind: spi.AuthBearer}, + Params: encoding.SamplingSet(2), + DocURL: "https://platform.openai.com/docs/api-reference/chat", + }) + + // OpenAI on the Responses protocol, where reasoning is a nested object and + // the output ceiling is spelled max_output_tokens. Offering it as a second + // protocol for the same vendor is what lets a model choose between them + // without a second provider entry. + spi.MustRegister(spi.Descriptor{ + Vendor: "openai", + Kind: spi.KindChat, + Protocol: spi.ProtocolOpenAIResponses, + DisplayName: "OpenAI (Responses)", + + DefaultBaseURL: openAIBaseURL, + Auth: spi.Auth{Kind: spi.AuthBearer}, + Params: []spi.Param{ + encoding.ThinkingMode( + encoding.EffortNone{Encoder: encoding.Nested("reasoning", "effort"), Off: "none"}, + spi.ThinkingOn, spi.ThinkingOff), + responsesEffortParam(openAIEffortLadder...), + responsesSummaryParam(), + encoding.Temperature(0, 2), + encoding.TopP(), + encoding.MaxTokensAs("max_output_tokens"), + }, + Constraints: []spi.Constraint{ + encoding.DependsOnThinking{Params: []spi.ParamID{spi.ParamThinkingEffort}}, + }, + DocURL: openAIReasoningDoc, + }) + + // Azure OpenAI mirrors the model families but authenticates with the + // api-key header instead of a bearer token. + spi.MustRegister(spi.Descriptor{ + Vendor: "azure_openai", + Kind: spi.KindChat, + Protocol: spi.ProtocolOpenAIChat, + DisplayName: "Azure OpenAI (reasoning)", + Models: openAIReasoningModels, + + Auth: spi.Auth{Kind: spi.AuthHeader, Header: "api-key"}, + Params: openAIReasoningParams(), + Constraints: openAIReasoningConstraints(), + DocURL: openAIReasoningDoc, + }) + + spi.MustRegister(spi.Descriptor{ + Vendor: "azure_openai", + Kind: spi.KindChat, + Protocol: spi.ProtocolOpenAIChat, + DisplayName: "Azure OpenAI", + + Auth: spi.Auth{Kind: spi.AuthHeader, Header: "api-key"}, + Params: encoding.SamplingSet(2), + DocURL: "https://learn.microsoft.com/azure/ai-services/openai/reference", + }) +} diff --git a/internal/models/llm/vendors/params.go b/internal/models/llm/vendors/params.go new file mode 100644 index 00000000000..1ff0102dfce --- /dev/null +++ b/internal/models/llm/vendors/params.go @@ -0,0 +1,53 @@ +// Package vendors holds the built-in model plugins: one declaration per vendor +// describing how it differs from the protocol baseline. +// +// Each declaration follows the vendor's published documentation and links to +// it, so a reviewer can check a claim instead of trusting it. A vendor that +// changes its API is a change to one declaration, and the request path, the +// validator, the model editor form, and the debug report all follow from it. +package vendors + +import ( + "github.com/Tencent/WeKnora/internal/models/llm/encoding" + "github.com/Tencent/WeKnora/internal/models/llm/spi" +) + +// responsesEffortParam declares the reasoning ladder on the Responses +// protocol, where effort is nested inside the reasoning object. +func responsesEffortParam(values ...string) spi.Param { + return encoding.ThinkingEffort(encoding.Nested("reasoning", "effort"), values...) +} + +// responsesSummaryParam declares the reasoning-summary control the Responses +// protocol exposes. The documentation is explicit that no summary is returned +// unless the request opts in, so this is a real knob rather than a display +// preference the client could apply afterwards. +// +// Docs: https://developers.openai.com/api/docs/guides/reasoning +func responsesSummaryParam() spi.Param { + p := encoding.ThinkingEffort(encoding.Nested("reasoning", "summary"), + "auto", "concise", "detailed") + p.ID = spi.ParamThinkingSummary + p.Enum = encoding.Options(spi.ParamThinkingSummary, "auto", "concise", "detailed") + p.UI.LabelKey = "model.param.thinking.summary.label" + p.UI.HelpKey = "model.param.thinking.summary.help" + p.UI.Order = 25 + return p +} + +// thinkingObject builds the `{"thinking": {"type": ...}}` toggle shared by +// DeepSeek, Zhipu, Volcengine, and Tencent LKEAP, with the vendor's own +// spellings. Only Volcengine documents the third `auto` value, so the others +// pass an empty string and the encoder refuses to invent one. +func thinkingObject(auto string) encoding.ThinkingObject { + return encoding.ThinkingObject{Key: "thinking", On: "enabled", Off: "disabled", Auto: auto} +} + +// thinkingModes returns the mode vocabulary for a vendor, including auto only +// when the vendor documents it. +func thinkingModes(withAuto bool) []string { + if withAuto { + return []string{spi.ThinkingOn, spi.ThinkingOff, spi.ThinkingAuto} + } + return []string{spi.ThinkingOn, spi.ThinkingOff} +} diff --git a/internal/models/llm/vendors/vendors_test.go b/internal/models/llm/vendors/vendors_test.go new file mode 100644 index 00000000000..2302bd23321 --- /dev/null +++ b/internal/models/llm/vendors/vendors_test.go @@ -0,0 +1,396 @@ +package vendors + +import ( + "encoding/json" + "testing" + + "github.com/Tencent/WeKnora/internal/models/llm/spi" +) + +// The table below is the executable form of every wire claim these plugins +// make. Each case names the vendor documentation it encodes, so a reviewer can +// check the expected body against the vendor's own examples rather than +// against the implementation that produced it. +// +// It is also the regression net for the behavior this seam replaced: the +// previous design expressed these differences as provider branches and a +// four-value thinking_control enum, and a change to one vendor could silently +// alter another. + +// canonicalBody is what an OpenAI-shaped protocol driver writes before any +// vendor encoder runs. Cases assert on the body after the plan is applied, so +// a parameter a vendor forbids must actually disappear from this. +func canonicalBody() map[string]any { + return map[string]any{ + "model": "placeholder", + "messages": []any{}, + "temperature": 0.7, + "max_tokens": float64(2048), + "tool_choice": "auto", + } +} + +type wireCase struct { + name string + vendor string + model string + // protocol pins the descriptor when a vendor offers more than one. + protocol spi.ProtocolID + stream bool + values map[spi.ParamID]spi.Value + + // want lists body paths that must hold a value, as dotted paths. + want map[string]any + // absent lists body paths that must not be present. + absent []string + // notes lists parameters that must carry an adjustment note, so a silent + // behavior change cannot pass as a deliberate one. + notes []spi.ParamID +} + +func TestVendorWireFormats(t *testing.T) { + cases := []wireCase{ + // --- DeepSeek: thinking object plus a three-rung effort ladder. + // https://api-docs.deepseek.com/guides/thinking_mode + { + name: "deepseek thinking on", + vendor: "deepseek", model: "deepseek-chat", + values: modeOnly(spi.ThinkingOn), + want: map[string]any{"thinking.type": "enabled"}, + absent: []string{"tool_choice"}, + }, + { + name: "deepseek thinking on with effort", + vendor: "deepseek", model: "deepseek-chat", + values: map[spi.ParamID]spi.Value{ + spi.ParamThinkingMode: spi.EnumValue(spi.ThinkingOn), + spi.ParamThinkingEffort: spi.EnumValue("high"), + }, + want: map[string]any{"thinking.type": "enabled", "reasoning_effort": "high"}, + }, + { + // Effort must not survive a disabled toggle: on the Responses + // spelling the two share one field, and a leftover rung would turn + // thinking back on. + name: "deepseek effort dropped when thinking off", + vendor: "deepseek", model: "deepseek-chat", + values: map[spi.ParamID]spi.Value{ + spi.ParamThinkingMode: spi.EnumValue(spi.ThinkingOff), + spi.ParamThinkingEffort: spi.EnumValue("high"), + }, + want: map[string]any{"thinking.type": "disabled"}, + absent: []string{"reasoning_effort"}, + notes: []spi.ParamID{spi.ParamThinkingEffort}, + }, + { + name: "deepseek rejects an effort rung it does not publish", + vendor: "deepseek", model: "deepseek-chat", + values: map[spi.ParamID]spi.Value{ + spi.ParamThinkingMode: spi.EnumValue(spi.ThinkingOn), + spi.ParamThinkingEffort: spi.EnumValue("medium"), + }, + absent: []string{"reasoning_effort"}, + notes: []spi.ParamID{spi.ParamThinkingEffort}, + }, + + // --- Alibaba Cloud Model Studio: a top-level boolean plus a budget. + // https://help.aliyun.com/zh/model-studio/deep-thinking + { + // The boolean is sent even when the caller says nothing, because + // several Qwen generations default to thinking on: omitting it + // would mean the opposite of the caller's silence. + name: "aliyun pins the boolean when the caller is silent", + vendor: "aliyun", model: "qwen3-max", stream: true, + want: map[string]any{"enable_thinking": false}, + notes: []spi.ParamID{spi.ParamThinkingMode}, + }, + { + name: "aliyun thinking on with a budget", + vendor: "aliyun", model: "qwen3-max", stream: true, + values: map[spi.ParamID]spi.Value{ + spi.ParamThinkingMode: spi.EnumValue(spi.ThinkingOn), + spi.ParamThinkingBudget: spi.IntValue(500), + }, + want: map[string]any{"enable_thinking": true, "thinking_budget": 500}, + }, + { + // Documented: several Qwen3 models accept thinking only when + // streaming. Downgrading locally beats a vendor error. + name: "aliyun downgrades thinking off stream", + vendor: "aliyun", model: "qwen3-max", stream: false, + values: modeOnly(spi.ThinkingOn), + want: map[string]any{"enable_thinking": false}, + notes: []spi.ParamID{spi.ParamThinkingMode}, + }, + { + name: "aliyun non-thinking model sends no thinking fields", + vendor: "aliyun", model: "qwen2.5-14b-instruct", stream: true, + values: modeOnly(spi.ThinkingOn), + absent: []string{"enable_thinking", "thinking"}, + notes: []spi.ParamID{spi.ParamThinkingMode}, + }, + + // --- Volcengine Ark: the only vendor documenting a third state. + { + name: "volcengine auto", + vendor: "volcengine", model: "doubao-seed-1-6", + values: modeOnly(spi.ThinkingAuto), + want: map[string]any{"thinking.type": "auto"}, + }, + { + name: "volcengine budget rides inside the thinking object", + vendor: "volcengine", model: "doubao-seed-1-6", + values: map[spi.ParamID]spi.Value{ + spi.ParamThinkingMode: spi.EnumValue(spi.ThinkingOn), + spi.ParamThinkingBudget: spi.IntValue(32000), + }, + want: map[string]any{"thinking.type": "enabled", "thinking.budget_tokens": 32000}, + }, + + // --- Zhipu GLM: the effort ladder exists only from GLM-5.2. + // https://docs.bigmodel.cn/cn/guide/capabilities/thinking + { + name: "zhipu 5.2 accepts effort", + vendor: "zhipu", model: "glm-5.2", + values: map[spi.ParamID]spi.Value{ + spi.ParamThinkingMode: spi.EnumValue(spi.ThinkingOn), + spi.ParamThinkingEffort: spi.EnumValue("xhigh"), + }, + want: map[string]any{"thinking.type": "enabled", "reasoning_effort": "xhigh"}, + }, + { + name: "zhipu 4.6 has no effort ladder", + vendor: "zhipu", model: "glm-4.6", + values: map[spi.ParamID]spi.Value{ + spi.ParamThinkingMode: spi.EnumValue(spi.ThinkingOn), + spi.ParamThinkingEffort: spi.EnumValue("xhigh"), + }, + want: map[string]any{"thinking.type": "enabled"}, + absent: []string{"reasoning_effort"}, + notes: []spi.ParamID{spi.ParamThinkingEffort}, + }, + + // --- OpenAI: the reasoning families reject the sampling knobs and + // rename the ceiling. https://developers.openai.com/api/docs/guides/reasoning + { + name: "openai reasoning model strips sampling and renames the ceiling", + vendor: "openai", model: "o3-mini", + values: map[spi.ParamID]spi.Value{ + spi.ParamThinkingEffort: spi.EnumValue("high"), + spi.ParamMaxTokens: spi.IntValue(4096), + }, + want: map[string]any{"reasoning_effort": "high", "max_completion_tokens": 4096}, + absent: []string{"temperature", "top_p", "max_tokens"}, + }, + { + name: "openai general model keeps sampling", + vendor: "openai", model: "gpt-4o", + values: map[spi.ParamID]spi.Value{spi.ParamTemperature: spi.FloatValue(0.3)}, + want: map[string]any{"temperature": 0.3, "max_tokens": 2048}, + }, + { + name: "openai thinking off writes the none rung", + vendor: "openai", model: "gpt-5", + values: modeOnly(spi.ThinkingOff), + want: map[string]any{"reasoning_effort": "none"}, + }, + { + name: "openai responses nests reasoning and renames the ceiling", + vendor: "openai", model: "gpt-5", protocol: spi.ProtocolOpenAIResponses, + values: map[spi.ParamID]spi.Value{ + spi.ParamThinkingMode: spi.EnumValue(spi.ThinkingOn), + spi.ParamThinkingEffort: spi.EnumValue("medium"), + spi.ParamThinkingSummary: spi.EnumValue("auto"), + spi.ParamMaxTokens: spi.IntValue(300), + }, + want: map[string]any{ + "reasoning.effort": "medium", + "reasoning.summary": "auto", + "max_output_tokens": 300, + }, + absent: []string{"max_tokens"}, + }, + + // --- Anthropic: two incompatible thinking modes chosen by model. + // https://platform.claude.com/docs/en/build-with-claude/thinking + { + name: "anthropic 4.6 uses adaptive thinking and effort", + vendor: "anthropic", model: "claude-sonnet-4-6", + values: map[spi.ParamID]spi.Value{ + spi.ParamThinkingMode: spi.EnumValue(spi.ThinkingOn), + spi.ParamThinkingEffort: spi.EnumValue("high"), + }, + want: map[string]any{ + "thinking.type": "adaptive", + "output_config.effort": "high", + "max_tokens": 1024, + }, + absent: []string{"thinking.budget_tokens"}, + }, + { + // budget_tokens must stay below max_tokens, and thinking tokens + // count against the same ceiling, so the ceiling is raised rather + // than the budget cut. + name: "anthropic 4.5 uses a budget and raises the ceiling", + vendor: "anthropic", model: "claude-sonnet-4-5", + values: map[spi.ParamID]spi.Value{ + spi.ParamThinkingMode: spi.EnumValue(spi.ThinkingOn), + spi.ParamThinkingBudget: spi.IntValue(10000), + }, + want: map[string]any{ + "thinking.type": "enabled", + "thinking.budget_tokens": 10000, + "max_tokens": 14096, + }, + notes: []spi.ParamID{spi.ParamMaxTokens}, + }, + { + name: "anthropic drops a budget when thinking is off", + vendor: "anthropic", model: "claude-sonnet-4-5", + values: map[spi.ParamID]spi.Value{ + spi.ParamThinkingMode: spi.EnumValue(spi.ThinkingOff), + spi.ParamThinkingBudget: spi.IntValue(10000), + }, + want: map[string]any{"thinking.type": "disabled"}, + absent: []string{"thinking.budget_tokens"}, + notes: []spi.ParamID{spi.ParamThinkingBudget}, + }, + { + name: "anthropic rejects a budget below the documented floor", + vendor: "anthropic", model: "claude-sonnet-4-5", + values: map[spi.ParamID]spi.Value{ + spi.ParamThinkingMode: spi.EnumValue(spi.ThinkingOn), + spi.ParamThinkingBudget: spi.IntValue(512), + }, + absent: []string{"thinking.budget_tokens"}, + notes: []spi.ParamID{spi.ParamThinkingBudget}, + }, + + // --- Self-hosted deployments pass the flag to the chat template. + { + name: "generic deployment uses chat_template_kwargs", + vendor: "generic", model: "qwen3-32b", + values: modeOnly(spi.ThinkingOn), + want: map[string]any{"chat_template_kwargs.enable_thinking": true}, + absent: []string{"enable_thinking", "thinking"}, + }, + + // --- Moonshot pins temperature for its fixed-temperature line. + { + name: "moonshot pins temperature and reports the override", + vendor: "moonshot", model: "moonshot-v1-8k", + values: map[spi.ParamID]spi.Value{spi.ParamTemperature: spi.FloatValue(0.2)}, + want: map[string]any{"temperature": float64(1)}, + notes: []spi.ParamID{spi.ParamTemperature}, + }, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + desc, ok := spi.Resolve(spi.Query{ + Vendor: tc.vendor, + Kind: spi.KindChat, + Model: tc.model, + Protocol: tc.protocol, + }) + if !ok { + t.Fatalf("no descriptor registered for %s/%s", tc.vendor, tc.model) + } + + plan, err := desc.Plan(spi.Request{Model: tc.model, Stream: tc.stream, Values: tc.values}) + if err != nil { + t.Fatalf("plan: %v", err) + } + + draft := spi.NewDraft(desc.Protocol, tc.model, tc.stream) + draft.Body = canonicalBody() + if err := plan.Apply(draft); err != nil { + t.Fatalf("apply: %v", err) + } + + for path, want := range tc.want { + got, ok := lookup(draft, path) + if !ok { + t.Errorf("body is missing %s\nbody: %s", path, dump(draft)) + continue + } + if !numericEqual(got, want) { + t.Errorf("body[%s] = %#v, want %#v\nbody: %s", path, got, want, dump(draft)) + } + } + for _, path := range tc.absent { + if got, ok := lookup(draft, path); ok { + t.Errorf("body[%s] should be absent, got %#v\nbody: %s", path, got, dump(draft)) + } + } + for _, id := range tc.notes { + if !hasNote(plan, id) { + t.Errorf("expected an adjustment note for %s, got %+v", id, plan.Notes) + } + } + }) + } +} + +func modeOnly(mode string) map[spi.ParamID]spi.Value { + return map[spi.ParamID]spi.Value{spi.ParamThinkingMode: spi.EnumValue(mode)} +} + +// lookup reads a dotted path out of the draft body. +func lookup(d *spi.Draft, path string) (any, bool) { + return d.GetNested(splitPath(path)...) +} + +func splitPath(path string) []string { + var out []string + start := 0 + for i := 0; i < len(path); i++ { + if path[i] == '.' { + out = append(out, path[start:i]) + start = i + 1 + } + } + return append(out, path[start:]) +} + +// numericEqual compares values while tolerating the int/float distinction that +// JSON round-tripping erases; the assertion is about the wire value, not the +// Go type that produced it. +func numericEqual(got, want any) bool { + gf, gok := asFloat(got) + wf, wok := asFloat(want) + if gok && wok { + return gf == wf + } + return got == want +} + +func asFloat(v any) (float64, bool) { + switch n := v.(type) { + case int: + return float64(n), true + case int64: + return float64(n), true + case float64: + return n, true + default: + return 0, false + } +} + +func hasNote(plan *spi.Plan, id spi.ParamID) bool { + for _, note := range plan.Notes { + if note.Param == id { + return true + } + } + return false +} + +func dump(d *spi.Draft) string { + out, err := json.Marshal(d.Body) + if err != nil { + return err.Error() + } + return string(out) +} From 4c2648f648bb232d89797b2255e95b966da530dd Mon Sep 17 00:00:00 2001 From: lyingbug <[email> Date: Thu, 13 Aug 2026 23:50:14 +0000 Subject: [PATCH 2/9] feat(models): implement the three standard protocol drivers Add OpenAI Chat Completions, OpenAI Responses, and Anthropic Messages as protocol drivers behind one seam. A protocol owns the canonical body, the endpoint, response decoding, and stream decoding; everything vendor-specific stays in the descriptor that selects it. The three learn about reasoning in three different ways - a delta field, a typed event, and a content block - and the shared emit helpers make the decoded stream identical downstream, including for models that inline their reasoning in tags across chunk boundaries. - spi/message.go: neutral request vocabulary shared by every protocol. - sse: event-aware reader; the typed protocols need the event name, and reasoning payloads exceed the default scanner limit. - protocol/{openaichat,responses,anthropic}: the drivers. - protocol/all: side-effect registration. Covered by stream fixtures following each vendor's documented event shape, asserting the full decoded transcript rather than individual fields. Co-authored-by: lyingbug --- internal/models/llm/protocol/all/all.go | 12 + .../models/llm/protocol/all/stream_test.go | 232 ++++++++ .../models/llm/protocol/anthropic/driver.go | 546 ++++++++++++++++++ .../models/llm/protocol/internal/emit/emit.go | 168 ++++++ .../models/llm/protocol/openaichat/driver.go | 453 +++++++++++++++ internal/models/llm/protocol/protocol.go | 123 ++++ .../models/llm/protocol/responses/driver.go | 438 ++++++++++++++ internal/models/llm/spi/message.go | 190 ++++++ internal/models/llm/sse/reader.go | 109 ++++ 9 files changed, 2271 insertions(+) create mode 100644 internal/models/llm/protocol/all/all.go create mode 100644 internal/models/llm/protocol/all/stream_test.go create mode 100644 internal/models/llm/protocol/anthropic/driver.go create mode 100644 internal/models/llm/protocol/internal/emit/emit.go create mode 100644 internal/models/llm/protocol/openaichat/driver.go create mode 100644 internal/models/llm/protocol/protocol.go create mode 100644 internal/models/llm/protocol/responses/driver.go create mode 100644 internal/models/llm/spi/message.go create mode 100644 internal/models/llm/sse/reader.go diff --git a/internal/models/llm/protocol/all/all.go b/internal/models/llm/protocol/all/all.go new file mode 100644 index 00000000000..4577a425b2a --- /dev/null +++ b/internal/models/llm/protocol/all/all.go @@ -0,0 +1,12 @@ +// Package all registers every built-in protocol driver. +// +// Importing it for its side effects is how a consumer opts into the standard +// protocols without naming each one, and it keeps the drivers themselves free +// of any dependency on each other. +package all + +import ( + _ "github.com/Tencent/WeKnora/internal/models/llm/protocol/anthropic" + _ "github.com/Tencent/WeKnora/internal/models/llm/protocol/openaichat" + _ "github.com/Tencent/WeKnora/internal/models/llm/protocol/responses" +) diff --git a/internal/models/llm/protocol/all/stream_test.go b/internal/models/llm/protocol/all/stream_test.go new file mode 100644 index 00000000000..b7b36d44718 --- /dev/null +++ b/internal/models/llm/protocol/all/stream_test.go @@ -0,0 +1,232 @@ +package all + +import ( + "context" + "strings" + "testing" + + "github.com/Tencent/WeKnora/internal/models/llm/protocol" + "github.com/Tencent/WeKnora/internal/models/llm/spi" + "github.com/Tencent/WeKnora/internal/types" +) + +// The fixtures below follow each vendor's documented stream shape. They exist +// because the three protocols learn about reasoning in three different ways — +// a delta field, a typed event, and a content block — and the whole point of +// the seam is that a consumer downstream cannot tell which one produced a +// given chunk. + +// collect drains a driver's decoded stream into a slice. +func collect(t *testing.T, id spi.ProtocolID, body string) []types.StreamResponse { + t.Helper() + driver, err := protocol.MustGet(id) + if err != nil { + t.Fatalf("driver %s: %v", id, err) + } + + out := make(chan types.StreamResponse, 64) + done := make(chan struct{}) + go func() { + defer close(done) + driver.DecodeStream(context.Background(), strings.NewReader(body), out) + close(out) + }() + + var chunks []types.StreamResponse + for chunk := range out { + chunks = append(chunks, chunk) + } + <-done + return chunks +} + +// transcript renders the decoded stream as a compact "type:content" list, so a +// failure shows the whole sequence rather than one mismatched field. +func transcript(chunks []types.StreamResponse) []string { + out := make([]string, 0, len(chunks)) + for _, c := range chunks { + entry := string(c.ResponseType) + ":" + c.Content + if c.Done { + entry += "|done" + } + out = append(out, entry) + } + return out +} + +func assertTranscript(t *testing.T, got []types.StreamResponse, want []string) { + t.Helper() + gotLines := transcript(got) + if len(gotLines) != len(want) { + t.Fatalf("stream had %d chunks, want %d\ngot: %q\nwant: %q", + len(gotLines), len(want), gotLines, want) + } + for i := range want { + if gotLines[i] != want[i] { + t.Errorf("chunk %d = %q, want %q\ngot: %q\nwant: %q", + i, gotLines[i], want[i], gotLines, want) + } + } +} + +func TestOpenAIChatStreamSeparatesReasoning(t *testing.T) { + body := strings.Join([]string{ + `data: {"choices":[{"delta":{"reasoning_content":"let me"}}]}`, + ``, + `data: {"choices":[{"delta":{"reasoning_content":" think"}}]}`, + ``, + `data: {"choices":[{"delta":{"content":"Paris"}}]}`, + ``, + `data: {"choices":[{"delta":{"content":"."},"finish_reason":"stop"}]}`, + ``, + `data: {"choices":[],"usage":{"prompt_tokens":10,"completion_tokens":5,"total_tokens":15}}`, + ``, + `data: [DONE]`, + ``, + }, "\n") + + chunks := collect(t, spi.ProtocolOpenAIChat, body) + assertTranscript(t, chunks, []string{ + "thinking:let me", + "thinking: think", + "thinking:|done", + "answer:Paris", + "answer:.", + "answer:|done", + }) + + final := chunks[len(chunks)-1] + if final.FinishReason != "stop" { + t.Errorf("finish reason = %q, want stop", final.FinishReason) + } + if final.Usage == nil || final.Usage.TotalTokens != 15 { + t.Errorf("usage = %+v, want total 15", final.Usage) + } +} + +// A model with no reasoning field inlines its reasoning in tags, and +// the tags can straddle chunk boundaries. The decoded stream must be +// indistinguishable from a vendor that uses a dedicated field. +func TestOpenAIChatStreamSplitsInlineThinkTags(t *testing.T) { + body := strings.Join([]string{ + `data: {"choices":[{"delta":{"content":"weigh options"}}]}`, + ``, + `data: {"choices":[{"delta":{"content":"Paris"}}]}`, + ``, + `data: [DONE]`, + ``, + }, "\n") + + assertTranscript(t, collect(t, spi.ProtocolOpenAIChat, body), []string{ + "thinking:weigh options", + "thinking:|done", + "answer:Paris", + "answer:|done", + }) +} + +func TestAnthropicStreamSeparatesThinkingBlocks(t *testing.T) { + body := strings.Join([]string{ + `event: message_start`, + `data: {"type":"message_start","message":{"usage":{"input_tokens":12,"output_tokens":0}}}`, + ``, + `event: content_block_start`, + `data: {"type":"content_block_start","index":0,"content_block":{"type":"thinking"}}`, + ``, + `event: content_block_delta`, + `data: {"type":"content_block_delta","index":0,"delta":{"type":"thinking_delta","thinking":"27 * 453"}}`, + ``, + `event: content_block_delta`, + `data: {"type":"content_block_delta","index":0,"delta":{"type":"signature_delta","signature":"EqQBCgIYAhIM"}}`, + ``, + `event: content_block_stop`, + `data: {"type":"content_block_stop","index":0}`, + ``, + `event: content_block_start`, + `data: {"type":"content_block_start","index":1,"content_block":{"type":"text"}}`, + ``, + `event: content_block_delta`, + `data: {"type":"content_block_delta","index":1,"delta":{"type":"text_delta","text":"12231"}}`, + ``, + `event: message_delta`, + `data: {"type":"message_delta","delta":{"stop_reason":"end_turn"},"usage":{"input_tokens":12,"output_tokens":40}}`, + ``, + `event: message_stop`, + `data: {"type":"message_stop"}`, + ``, + }, "\n") + + chunks := collect(t, spi.ProtocolAnthropicMessages, body) + assertTranscript(t, chunks, []string{ + "thinking:27 * 453", + "thinking:|done", + "answer:12231", + "answer:|done", + }) + + final := chunks[len(chunks)-1] + if final.FinishReason != "end_turn" { + t.Errorf("finish reason = %q, want end_turn", final.FinishReason) + } + // The signature delta must not surface as user-visible text; it exists to + // be replayed, not read. + for _, chunk := range chunks { + if strings.Contains(chunk.Content, "EqQBCgIYAhIM") { + t.Errorf("signature leaked into the stream: %q", chunk.Content) + } + } + if final.Usage == nil || final.Usage.CompletionTokens != 40 { + t.Errorf("usage = %+v, want 40 completion tokens", final.Usage) + } +} + +func TestResponsesStreamSeparatesReasoningSummary(t *testing.T) { + body := strings.Join([]string{ + `event: response.reasoning_summary_text.delta`, + `data: {"type":"response.reasoning_summary_text.delta","delta":"Answering a simple question"}`, + ``, + `event: response.output_text.delta`, + `data: {"type":"response.output_text.delta","delta":"The capital of France is Paris."}`, + ``, + `event: response.completed`, + `data: {"type":"response.completed","response":{"status":"completed","usage":{"input_tokens":14,"output_tokens":9,"total_tokens":23}}}`, + ``, + }, "\n") + + chunks := collect(t, spi.ProtocolOpenAIResponses, body) + assertTranscript(t, chunks, []string{ + "thinking:Answering a simple question", + "thinking:|done", + "answer:The capital of France is Paris.", + "answer:|done", + }) + + final := chunks[len(chunks)-1] + if final.FinishReason != "stop" { + t.Errorf("finish reason = %q, want stop", final.FinishReason) + } + if final.Usage == nil || final.Usage.TotalTokens != 23 { + t.Errorf("usage = %+v, want total 23", final.Usage) + } +} + +// A response that exhausts its output ceiling before writing visible text must +// say so, otherwise an empty answer looks like a normal completion. +func TestResponsesStreamReportsIncomplete(t *testing.T) { + body := strings.Join([]string{ + `data: {"type":"response.incomplete","response":{"status":"incomplete",` + + `"incomplete_details":{"reason":"max_output_tokens"},` + + `"usage":{"input_tokens":14,"output_tokens":300,"total_tokens":314}}}`, + ``, + }, "\n") + + chunks := collect(t, spi.ProtocolOpenAIResponses, body) + final := chunks[len(chunks)-1] + if final.FinishReason != "max_output_tokens" { + t.Errorf("finish reason = %q, want max_output_tokens", final.FinishReason) + } +} diff --git a/internal/models/llm/protocol/anthropic/driver.go b/internal/models/llm/protocol/anthropic/driver.go new file mode 100644 index 00000000000..3d72822a278 --- /dev/null +++ b/internal/models/llm/protocol/anthropic/driver.go @@ -0,0 +1,546 @@ +// Package anthropic implements the Anthropic Messages protocol. +// +// It differs from the OpenAI-shaped protocols in ways that are structural +// rather than cosmetic: the system prompt is a separate top-level field, every +// message body is a list of typed content blocks, tool results travel as user +// blocks instead of a tool role, and reasoning is a first-class block carrying +// a signature that must return unmodified on the next turn. +package anthropic + +import ( + "context" + "encoding/json" + "fmt" + "io" + "net/url" + "strings" + + "github.com/Tencent/WeKnora/internal/models/llm/protocol" + "github.com/Tencent/WeKnora/internal/models/llm/protocol/internal/emit" + "github.com/Tencent/WeKnora/internal/models/llm/spi" + "github.com/Tencent/WeKnora/internal/models/llm/sse" + "github.com/Tencent/WeKnora/internal/types" +) + +// defaultMaxTokens is the fallback ceiling. The Messages API requires +// max_tokens on every request, unlike the OpenAI-shaped protocols where it is +// optional, so the driver must supply one when the vendor descriptor has not. +const defaultMaxTokens = 1024 + +// Driver implements protocol.Driver for the Messages API. +type Driver struct{} + +func init() { protocol.Register(Driver{}) } + +// ID reports the protocol. +func (Driver) ID() spi.ProtocolID { return spi.ProtocolAnthropicMessages } + +// EndpointPath reports the standard path. +func (Driver) EndpointPath() string { return "/v1/messages" } + +// BuildDraft renders the call into a Messages body. +func (d Driver) BuildDraft(call protocol.Call) (*spi.Draft, error) { + draft := spi.NewDraft(d.ID(), call.Model, call.Stream) + draft.Set("model", call.Model) + if call.Stream { + draft.Set("stream", true) + } + + system, messages := convertMessages(call.Messages, call.Replay) + if system != "" { + draft.Set("system", system) + } + draft.Set("messages", messages) + + maxTokens := defaultMaxTokens + opts := call.Options + if opts != nil { + if v := opts.EffectiveMaxTokens(); v > 0 { + maxTokens = v + } + if opts.Temperature > 0 { + draft.Set("temperature", opts.Temperature) + } + if opts.TopP > 0 { + draft.Set("top_p", opts.TopP) + } + if len(opts.Tools) > 0 { + draft.Set("tools", convertTools(opts.Tools)) + } + if choice := convertToolChoice(opts.ToolChoice); choice != nil { + draft.Set("tool_choice", choice) + } + } + draft.Set("max_tokens", maxTokens) + return draft, nil +} + +// convertMessages splits the system prompt out and renders the rest as +// content-block messages. +// +// Consecutive same-role messages are merged because the API rejects a +// conversation that alternates incorrectly, and tool results in particular +// arrive as several separate messages that must become one user turn. +func convertMessages(messages []spi.Message, replay spi.ReasoningReplay) (string, []any) { + var systemParts []string + type turn struct { + role string + blocks []any + } + var turns []turn + + appendBlocks := func(role string, blocks []any) { + if len(blocks) == 0 { + return + } + if n := len(turns); n > 0 && turns[n-1].role == role { + turns[n-1].blocks = append(turns[n-1].blocks, blocks...) + return + } + turns = append(turns, turn{role: role, blocks: blocks}) + } + + for _, msg := range messages { + switch msg.Role { + case "system": + if text := messageText(msg); text != "" { + systemParts = append(systemParts, text) + } + + case "tool": + // A tool result is a user-role block referencing the call it + // answers, not a role of its own. + appendBlocks("user", []any{map[string]any{ + "type": "tool_result", + "tool_use_id": msg.ToolCallID, + "content": msg.Content, + }}) + + case "assistant": + var blocks []any + // Thinking blocks must come first and must return exactly as they + // arrived, signature included, or the API rejects the turn. + if protocol.ShouldReplayReasoning(replay, msg) { + block := map[string]any{"type": "thinking", "thinking": msg.ReasoningContent} + if msg.ReasoningSignature != "" { + block["signature"] = msg.ReasoningSignature + } + blocks = append(blocks, block) + } + blocks = append(blocks, contentBlocks(msg)...) + for _, call := range msg.ToolCalls { + var input any + if err := json.Unmarshal([]byte(call.Function.Arguments), &input); err != nil { + input = map[string]any{} + } + blocks = append(blocks, map[string]any{ + "type": "tool_use", + "id": call.ID, + "name": call.Function.Name, + "input": input, + }) + } + appendBlocks("assistant", blocks) + + default: + appendBlocks("user", contentBlocks(msg)) + } + } + + out := make([]any, 0, len(turns)) + for _, t := range turns { + out = append(out, map[string]any{"role": t.role, "content": t.blocks}) + } + return strings.Join(systemParts, "\n\n"), out +} + +// contentBlocks renders a message body as text and image blocks. +func contentBlocks(msg spi.Message) []any { + var blocks []any + for _, part := range msg.MultiContent { + switch part.Type { + case "text": + if part.Text != "" { + blocks = append(blocks, map[string]any{"type": "text", "text": part.Text}) + } + case "image_url": + if part.ImageURL != nil { + if block := imageBlock(part.ImageURL.URL); block != nil { + blocks = append(blocks, block) + } + } + } + } + for _, img := range msg.Images { + if block := imageBlock(img); block != nil { + blocks = append(blocks, block) + } + } + if msg.Content != "" { + blocks = append(blocks, map[string]any{"type": "text", "text": msg.Content}) + } + return blocks +} + +// imageBlock renders an image reference. Anthropic takes base64 payloads and +// URLs through different source shapes, so a data URI has to be taken apart +// rather than passed along. +func imageBlock(reference string) map[string]any { + if reference == "" { + return nil + } + if strings.HasPrefix(reference, "data:") { + mediaType, data, ok := parseDataURI(reference) + if !ok { + return nil + } + return map[string]any{ + "type": "image", + "source": map[string]any{ + "type": "base64", + "media_type": mediaType, + "data": data, + }, + } + } + if u, err := url.Parse(reference); err != nil || u.Scheme == "" { + return nil + } + return map[string]any{ + "type": "image", + "source": map[string]any{"type": "url", "url": reference}, + } +} + +func parseDataURI(uri string) (mediaType, data string, ok bool) { + rest := strings.TrimPrefix(uri, "data:") + comma := strings.IndexByte(rest, ',') + if comma < 0 { + return "", "", false + } + meta, payload := rest[:comma], rest[comma+1:] + if !strings.HasSuffix(meta, ";base64") { + return "", "", false + } + mediaType = strings.TrimSuffix(meta, ";base64") + if mediaType == "" { + mediaType = "image/png" + } + return mediaType, payload, true +} + +func messageText(msg spi.Message) string { + if msg.Content != "" { + return msg.Content + } + var parts []string + for _, part := range msg.MultiContent { + if part.Type == "text" && part.Text != "" { + parts = append(parts, part.Text) + } + } + return strings.Join(parts, "\n") +} + +func convertTools(tools []spi.Tool) []any { + out := make([]any, 0, len(tools)) + for _, tool := range tools { + var schema any + if len(tool.Function.Parameters) > 0 { + _ = json.Unmarshal(tool.Function.Parameters, &schema) + } + if schema == nil { + schema = map[string]any{"type": "object", "properties": map[string]any{}} + } + out = append(out, map[string]any{ + "name": tool.Function.Name, + "description": tool.Function.Description, + "input_schema": schema, + }) + } + return out +} + +// convertToolChoice maps the neutral choice onto the object form this API uses. +func convertToolChoice(choice string) map[string]any { + switch choice { + case "": + return nil + case "auto": + return map[string]any{"type": "auto"} + case "required": + return map[string]any{"type": "any"} + case "none": + return map[string]any{"type": "none"} + default: + // Anything else names a specific tool. + return map[string]any{"type": "tool", "name": choice} + } +} + +// message mirrors a complete Messages response. +type message struct { + Content []struct { + Type string `json:"type"` + Text string `json:"text"` + Thinking string `json:"thinking"` + Signature string `json:"signature"` + ID string `json:"id"` + Name string `json:"name"` + Input json.RawMessage `json:"input"` + } `json:"content"` + StopReason string `json:"stop_reason"` + Usage *usage `json:"usage"` + Error *struct { + Type string `json:"type"` + Message string `json:"message"` + } `json:"error"` +} + +// ParseResponse decodes a complete Messages response. +func (d Driver) ParseResponse(body []byte) (*types.ChatResponse, error) { + var resp message + if err := json.Unmarshal(body, &resp); err != nil { + return nil, fmt.Errorf("decode message: %w", err) + } + if resp.Error != nil && resp.Error.Message != "" { + return nil, fmt.Errorf("provider error: %s", resp.Error.Message) + } + + out := &types.ChatResponse{FinishReason: resp.StopReason, Usage: resp.Usage.normalize()} + var text, thinking []string + for _, block := range resp.Content { + switch block.Type { + case "text": + text = append(text, block.Text) + case "thinking": + thinking = append(thinking, block.Thinking) + case "redacted_thinking": + // The payload is encrypted and carries no readable text; the turn + // still reasoned, so the fact is worth surfacing. + thinking = append(thinking, "") + case "tool_use": + out.ToolCalls = append(out.ToolCalls, types.LLMToolCall{ + ID: block.ID, + Type: "function", + Function: types.FunctionCall{ + Name: block.Name, + Arguments: string(block.Input), + }, + }) + } + } + out.Content = strings.Join(text, "") + out.ReasoningContent = strings.Join(thinking, "") + return out, nil +} + +// streamEvent mirrors the events this driver consumes. +type streamEvent struct { + Type string `json:"type"` + Index int `json:"index"` + Message *struct { + Usage *usage `json:"usage"` + } `json:"message"` + ContentBlock *struct { + Type string `json:"type"` + ID string `json:"id"` + Name string `json:"name"` + } `json:"content_block"` + Delta *struct { + Type string `json:"type"` + Text string `json:"text"` + Thinking string `json:"thinking"` + Signature string `json:"signature"` + PartialJSON string `json:"partial_json"` + StopReason string `json:"stop_reason"` + } `json:"delta"` + Usage *usage `json:"usage"` + Error *struct { + Type string `json:"type"` + Message string `json:"message"` + } `json:"error"` +} + +// DecodeStream decodes a Messages SSE stream. +func (d Driver) DecodeStream(ctx context.Context, body io.Reader, out chan<- types.StreamResponse) { + reader := sse.NewReader(body) + thinking := &emit.Thinking{} + + var ( + finishReason string + acc *types.TokenUsage + blockType string + ) + + for { + select { + case <-ctx.Done(): + emit.Error(out, ctx.Err().Error()) + return + default: + } + + event, err := reader.Next() + if err == io.EOF { + break + } + if err != nil { + emit.Error(out, fmt.Sprintf("read stream: %v", err)) + return + } + if event.Done || len(event.Data) == 0 { + if event.Done { + break + } + continue + } + + var ev streamEvent + if err := json.Unmarshal(event.Data, &ev); err != nil { + emit.Error(out, fmt.Sprintf("decode stream event: %v", err)) + return + } + if ev.Error != nil && ev.Error.Message != "" { + emit.Error(out, ev.Error.Message) + return + } + + switch ev.Type { + case "message_start": + if ev.Message != nil { + acc = mergeUsage(acc, ev.Message.Usage) + } + + case "content_block_start": + if ev.ContentBlock != nil { + blockType = ev.ContentBlock.Type + } + // Text follows reasoning, so the panel closes here rather than + // waiting for the first token. + if blockType == "text" { + thinking.Finish(out) + } + + case "content_block_delta": + if ev.Delta == nil { + continue + } + switch ev.Delta.Type { + case "thinking_delta": + thinking.Emit(out, ev.Delta.Thinking) + case "text_delta": + thinking.Finish(out) + if ev.Delta.Text != "" { + out <- types.StreamResponse{ + ResponseType: types.ResponseTypeAnswer, + Content: ev.Delta.Text, + } + } + case "signature_delta", "input_json_delta": + // The signature closes a thinking block and the partial JSON + // accumulates a tool call; neither is user-visible text. + } + + case "content_block_stop": + blockType = "" + + case "message_delta": + if ev.Delta != nil && ev.Delta.StopReason != "" { + finishReason = ev.Delta.StopReason + } + acc = mergeUsage(acc, ev.Usage) + + case "message_stop": + // The stream ends after this event. + } + } + + thinking.Finish(out) + out <- types.StreamResponse{ + ResponseType: types.ResponseTypeAnswer, + Done: true, + Usage: acc, + FinishReason: finishReason, + } +} + +// usage mirrors Anthropic's token accounting, where cache counters sit beside +// the input count rather than inside it. +type usage struct { + InputTokens int `json:"input_tokens"` + OutputTokens int `json:"output_tokens"` + CacheCreationInputTokens *int `json:"cache_creation_input_tokens"` + CacheReadInputTokens *int `json:"cache_read_input_tokens"` +} + +// normalize folds Anthropic's counters into the shared usage model. Cache +// reads and writes are additional input tokens here, not a subset of +// input_tokens, so the prompt total is their sum. +func (u *usage) normalize() types.TokenUsage { + if u == nil { + return types.TokenUsage{} + } + read := valueOr(u.CacheReadInputTokens) + write := valueOr(u.CacheCreationInputTokens) + prompt := u.InputTokens + read + write + + out := types.TokenUsage{ + PromptTokens: prompt, + CompletionTokens: u.OutputTokens, + TotalTokens: prompt + u.OutputTokens, + } + reported := u.CacheReadInputTokens != nil || u.CacheCreationInputTokens != nil + miss := prompt - read + if miss < 0 { + miss = 0 + } + out.SetPromptCacheUsage(read, write, miss, reported) + return out +} + +// mergeUsage combines the partial counts a stream reports at its start and end. +func mergeUsage(current *types.TokenUsage, next *usage) *types.TokenUsage { + if next == nil { + return current + } + merged := next.normalize() + if current == nil { + return &merged + } + // message_start reports input counts and message_delta reports the final + // output count, so each field takes the larger of the two observations. + if merged.PromptTokens > current.PromptTokens { + current.PromptTokens = merged.PromptTokens + current.CacheReadTokens = merged.CacheReadTokens + current.CacheWriteTokens = merged.CacheWriteTokens + current.CacheMissTokens = merged.CacheMissTokens + current.CacheReported = current.CacheReported || merged.CacheReported + } + if merged.CompletionTokens > current.CompletionTokens { + current.CompletionTokens = merged.CompletionTokens + } + current.TotalTokens = current.PromptTokens + current.CompletionTokens + return current +} + +func valueOr(v *int) int { + if v == nil { + return 0 + } + return *v +} + +// Endpoint reports the full URL for a base URL, tolerating the spellings users +// paste: a bare host, a versioned base, or the endpoint itself. +func (d Driver) Endpoint(baseURL string) string { + base := strings.TrimRight(baseURL, "/") + switch { + case strings.HasSuffix(base, "/messages"): + return base + case strings.HasSuffix(base, "/v1"), strings.HasSuffix(base, "/v1beta"): + return base + "/messages" + default: + return base + d.EndpointPath() + } +} diff --git a/internal/models/llm/protocol/internal/emit/emit.go b/internal/models/llm/protocol/internal/emit/emit.go new file mode 100644 index 00000000000..cee71159a22 --- /dev/null +++ b/internal/models/llm/protocol/internal/emit/emit.go @@ -0,0 +1,168 @@ +// Package emit holds the stream-emission bookkeeping every protocol driver +// shares: the reasoning-to-answer hand-off, the inline convention, and +// error termination. +// +// Centralizing it is what keeps three protocols behaving identically from the +// UI's point of view, even though each learns about reasoning differently — +// Chat Completions through a delta field, Responses through typed events, and +// Anthropic through content blocks. +package emit + +import ( + "strings" + + "github.com/Tencent/WeKnora/internal/types" +) + +const ( + thinkOpen = "" + thinkClose = "" +) + +// Thinking owns the reasoning-to-answer hand-off: reasoning chunks flow as +// they arrive, and exactly one done marker is emitted before the first answer +// token, or when the stream ends without one. +// +// Consumers rely on that single marker to close the thinking panel, so +// emitting it twice or never both misrender. +type Thinking struct { + active bool +} + +// Emit forwards a reasoning chunk and records that a done marker is owed. +func (t *Thinking) Emit(out chan<- types.StreamResponse, content string) { + if content == "" { + return + } + t.active = true + out <- types.StreamResponse{ + ResponseType: types.ResponseTypeThinking, + Content: content, + } +} + +// Finish emits the done marker if one is owed. It is safe to call repeatedly. +func (t *Thinking) Finish(out chan<- types.StreamResponse) { + if !t.active { + return + } + t.active = false + out <- types.StreamResponse{ + ResponseType: types.ResponseTypeThinking, + Done: true, + } +} + +// Active reports whether a done marker is still owed. +func (t *Thinking) Active() bool { return t.active } + +// Error terminates a stream with an error chunk. +func Error(out chan<- types.StreamResponse, message string) { + out <- types.StreamResponse{ + ResponseType: types.ResponseTypeError, + Content: message, + Done: true, + } +} + +// InlineSplitter separates reasoning that a model inlines in tags from +// the answer that follows, across chunk boundaries. +// +// Some open-weight models have no reasoning field and emit the tags in the +// content instead. Handling that in the splitter rather than at the end lets +// the reasoning stream live, and the state machine is needed because a tag can +// straddle two chunks: "" arrive separately often enough that a +// per-chunk check silently leaks tags into the answer. +type InlineSplitter struct { + inThinking bool + started bool + pending string +} + +// Feed consumes one content chunk and reports the answer text and the inline +// reasoning text it contained. +func (s *InlineSplitter) Feed(chunk string) (answer, thinking string) { + s.pending += chunk + + var answerParts, thinkingParts []string + for s.pending != "" { + if s.inThinking { + idx := strings.Index(s.pending, thinkClose) + if idx < 0 { + // Hold back a possible partial closing tag; release the rest. + safe, held := splitPartial(s.pending, thinkClose) + thinkingParts = append(thinkingParts, safe) + s.pending = held + break + } + thinkingParts = append(thinkingParts, s.pending[:idx]) + s.pending = s.pending[idx+len(thinkClose):] + s.inThinking = false + continue + } + + // The convention only applies when the tag opens the message; a + // mid-answer "" is ordinary text a model may legitimately write. + if !s.started { + trimmed := strings.TrimLeft(s.pending, " \t\r\n") + if strings.HasPrefix(trimmed, thinkOpen) { + s.started = true + s.inThinking = true + s.pending = trimmed[len(thinkOpen):] + continue + } + if isPrefixOf(trimmed, thinkOpen) { + // Could still become an opening tag once more arrives. + break + } + s.started = true + } + + answerParts = append(answerParts, s.pending) + s.pending = "" + } + + return strings.Join(answerParts, ""), strings.Join(thinkingParts, "") +} + +// splitPartial returns the portion of s that cannot be the start of tag, plus +// the trailing portion that might be. +func splitPartial(s, tag string) (safe, held string) { + maxHold := len(tag) - 1 + if maxHold > len(s) { + maxHold = len(s) + } + for hold := maxHold; hold > 0; hold-- { + if isPrefixOf(s[len(s)-hold:], tag) { + return s[:len(s)-hold], s[len(s)-hold:] + } + } + return s, "" +} + +// isPrefixOf reports whether s could be the beginning of tag. +func isPrefixOf(s, tag string) bool { + if s == "" || len(s) >= len(tag) { + return false + } + return strings.HasPrefix(tag, s) +} + +// SplitInlineThinking separates inlined reasoning from the answer in a +// complete, non-streamed message. +func SplitInlineThinking(content string) (answer, thinking string) { + trimmed := strings.TrimSpace(content) + if !strings.HasPrefix(trimmed, thinkOpen) { + return content, "" + } + // The last closing tag wins, so a model that nests or repeats the tag does + // not leak its reasoning into the answer. + end := strings.LastIndex(trimmed, thinkClose) + if end < 0 { + // Reasoning was truncated before it closed; there is no answer yet. + return "", strings.TrimSpace(trimmed[len(thinkOpen):]) + } + thinking = strings.TrimSpace(trimmed[len(thinkOpen):end]) + answer = strings.TrimSpace(trimmed[end+len(thinkClose):]) + return answer, thinking +} diff --git a/internal/models/llm/protocol/openaichat/driver.go b/internal/models/llm/protocol/openaichat/driver.go new file mode 100644 index 00000000000..c33be773c94 --- /dev/null +++ b/internal/models/llm/protocol/openaichat/driver.go @@ -0,0 +1,453 @@ +// Package openaichat implements the OpenAI Chat Completions protocol, the +// baseline most vendors serve. +package openaichat + +import ( + "context" + "encoding/json" + "fmt" + "io" + "strings" + + "github.com/Tencent/WeKnora/internal/models/llm/protocol" + "github.com/Tencent/WeKnora/internal/models/llm/protocol/internal/emit" + "github.com/Tencent/WeKnora/internal/models/llm/spi" + "github.com/Tencent/WeKnora/internal/models/llm/sse" + "github.com/Tencent/WeKnora/internal/types" +) + +// Driver implements protocol.Driver for Chat Completions. +type Driver struct{} + +func init() { protocol.Register(Driver{}) } + +// ID reports the protocol. +func (Driver) ID() spi.ProtocolID { return spi.ProtocolOpenAIChat } + +// EndpointPath reports the standard path. +func (Driver) EndpointPath() string { return "/chat/completions" } + +// BuildDraft renders the call into a Chat Completions body. +// +// Only fields the caller actually set are written. An OpenAI-shaped API treats +// an explicit zero as a real value, so writing every field would override +// model defaults the caller never touched — and would also defeat the vendor +// layer, whose forbidden-parameter contract is about removing fields that +// would otherwise be present. +func (d Driver) BuildDraft(call protocol.Call) (*spi.Draft, error) { + draft := spi.NewDraft(d.ID(), call.Model, call.Stream) + draft.Set("model", call.Model) + draft.Set("messages", convertMessages(call.Messages, call.Replay)) + if call.Stream { + draft.Set("stream", true) + // Without this, usage is absent from streamed responses and every + // token count downstream silently reads zero. + draft.Set("stream_options", map[string]any{"include_usage": true}) + } + + opts := call.Options + if opts == nil { + return draft, nil + } + if opts.Temperature > 0 { + draft.Set("temperature", opts.Temperature) + } + if opts.TopP > 0 { + draft.Set("top_p", opts.TopP) + } + if opts.FrequencyPenalty != 0 { + draft.Set("frequency_penalty", opts.FrequencyPenalty) + } + if opts.PresencePenalty != 0 { + draft.Set("presence_penalty", opts.PresencePenalty) + } + if opts.Seed != 0 { + draft.Set("seed", opts.Seed) + } + if maxTokens := opts.EffectiveMaxTokens(); maxTokens > 0 { + draft.Set("max_tokens", maxTokens) + } + if len(opts.Tools) > 0 { + draft.Set("tools", convertTools(opts.Tools)) + } + if opts.ToolChoice != "" { + draft.Set("tool_choice", opts.ToolChoice) + } + if opts.ParallelToolCalls != nil { + draft.Set("parallel_tool_calls", *opts.ParallelToolCalls) + } + if len(opts.Format) > 0 { + var format any + if err := json.Unmarshal(opts.Format, &format); err != nil { + return nil, fmt.Errorf("decode response format: %w", err) + } + draft.Set("response_format", format) + } + return draft, nil +} + +// convertMessages renders the conversation in Chat Completions shape. +func convertMessages(messages []spi.Message, replay spi.ReasoningReplay) []any { + out := make([]any, 0, len(messages)) + for _, msg := range messages { + wire := map[string]any{"role": msg.Role} + + switch { + case len(msg.MultiContent) > 0: + wire["content"] = convertParts(msg.MultiContent) + case len(msg.Images) > 0 && msg.Role == "user": + parts := make([]any, 0, len(msg.Images)+1) + for _, img := range msg.Images { + parts = append(parts, map[string]any{ + "type": "image_url", + "image_url": map[string]any{"url": img}, + }) + } + parts = append(parts, map[string]any{"type": "text", "text": msg.Content}) + wire["content"] = parts + default: + wire["content"] = msg.Content + } + + if len(msg.ToolCalls) > 0 { + wire["tool_calls"] = convertToolCalls(msg.ToolCalls) + } + if msg.Role == "tool" { + wire["tool_call_id"] = msg.ToolCallID + if msg.Name != "" { + wire["name"] = msg.Name + } + } else if msg.Name != "" { + wire["name"] = msg.Name + } + if protocol.ShouldReplayReasoning(replay, msg) { + wire["reasoning_content"] = msg.ReasoningContent + } + out = append(out, wire) + } + return out +} + +func convertParts(parts []spi.MessageContentPart) []any { + out := make([]any, 0, len(parts)) + for _, part := range parts { + switch part.Type { + case "text": + out = append(out, map[string]any{"type": "text", "text": part.Text}) + case "image_url": + if part.ImageURL == nil { + continue + } + image := map[string]any{"url": part.ImageURL.URL} + if part.ImageURL.Detail != "" { + image["detail"] = part.ImageURL.Detail + } + out = append(out, map[string]any{"type": "image_url", "image_url": image}) + } + } + return out +} + +func convertToolCalls(calls []spi.ToolCall) []any { + out := make([]any, 0, len(calls)) + for _, call := range calls { + callType := call.Type + if callType == "" { + callType = "function" + } + wire := map[string]any{ + "id": call.ID, + "type": callType, + "function": map[string]any{ + "name": call.Function.Name, + "arguments": call.Function.Arguments, + }, + } + // Vendor state travels back exactly as it arrived. Gemini's thought + // signatures ride here, and a replayed call without them is rejected. + for key, raw := range call.ProviderMetadata { + var value any + if err := json.Unmarshal(raw, &value); err == nil { + wire[key] = value + } + } + out = append(out, wire) + } + return out +} + +func convertTools(tools []spi.Tool) []any { + out := make([]any, 0, len(tools)) + for _, tool := range tools { + var params any + if len(tool.Function.Parameters) > 0 { + _ = json.Unmarshal(tool.Function.Parameters, ¶ms) + } + out = append(out, map[string]any{ + "type": "function", + "function": map[string]any{ + "name": tool.Function.Name, + "description": tool.Function.Description, + "parameters": params, + }, + }) + } + return out +} + +// completion mirrors the response fields this driver consumes. Vendors add +// their own, which json ignores. +type completion struct { + Choices []struct { + Message struct { + Content string `json:"content"` + ReasoningContent string `json:"reasoning_content"` + Reasoning string `json:"reasoning"` + ToolCalls []json.RawMessage `json:"tool_calls"` + } `json:"message"` + FinishReason string `json:"finish_reason"` + } `json:"choices"` + Usage *usage `json:"usage"` + Error *apiError `json:"error"` +} + +type apiError struct { + Message string `json:"message"` + Type string `json:"type"` + Code any `json:"code"` +} + +// ParseResponse decodes a complete Chat Completions response. +func (d Driver) ParseResponse(body []byte) (*types.ChatResponse, error) { + var resp completion + if err := json.Unmarshal(body, &resp); err != nil { + return nil, fmt.Errorf("decode chat completion: %w", err) + } + if resp.Error != nil && resp.Error.Message != "" { + return nil, fmt.Errorf("provider error: %s", resp.Error.Message) + } + if len(resp.Choices) == 0 { + return nil, fmt.Errorf("chat completion contained no choices") + } + + choice := resp.Choices[0] + reasoning := choice.Message.ReasoningContent + if reasoning == "" { + // Some OpenAI-compatible servers spell it `reasoning`. + reasoning = choice.Message.Reasoning + } + content, inlineReasoning := emit.SplitInlineThinking(choice.Message.Content) + if reasoning == "" { + reasoning = inlineReasoning + } + + out := &types.ChatResponse{ + Content: content, + ReasoningContent: reasoning, + FinishReason: choice.FinishReason, + ToolCalls: decodeToolCalls(choice.Message.ToolCalls), + Usage: resp.Usage.normalize(), + } + return out, nil +} + +// decodeToolCalls decodes tool calls while preserving any vendor fields, which +// the standard shape cannot describe but the next turn must carry back. +func decodeToolCalls(raw []json.RawMessage) []types.LLMToolCall { + if len(raw) == 0 { + return nil + } + out := make([]types.LLMToolCall, 0, len(raw)) + for _, item := range raw { + var call struct { + ID string `json:"id"` + Type string `json:"type"` + Function struct { + Name string `json:"name"` + Arguments string `json:"arguments"` + } `json:"function"` + } + if err := json.Unmarshal(item, &call); err != nil { + continue + } + out = append(out, types.LLMToolCall{ + ID: call.ID, + Type: call.Type, + Function: types.FunctionCall{ + Name: call.Function.Name, + Arguments: call.Function.Arguments, + }, + }) + } + return out +} + +// DecodeStream decodes a Chat Completions SSE stream. +func (d Driver) DecodeStream(ctx context.Context, body io.Reader, out chan<- types.StreamResponse) { + reader := sse.NewReader(body) + thinking := &emit.Thinking{} + splitter := &emit.InlineSplitter{} + + var ( + finishReason string + acc *types.TokenUsage + ) + + for { + select { + case <-ctx.Done(): + emit.Error(out, ctx.Err().Error()) + return + default: + } + + event, err := reader.Next() + if err == io.EOF { + break + } + if err != nil { + emit.Error(out, fmt.Sprintf("read stream: %v", err)) + return + } + if event.Done { + break + } + if len(event.Data) == 0 { + continue + } + + var chunk streamChunk + if err := json.Unmarshal(event.Data, &chunk); err != nil { + emit.Error(out, fmt.Sprintf("decode stream chunk: %v", err)) + return + } + if chunk.Error != nil && chunk.Error.Message != "" { + emit.Error(out, chunk.Error.Message) + return + } + if chunk.Usage != nil { + usage := chunk.Usage.normalize() + acc = &usage + } + if len(chunk.Choices) == 0 { + continue + } + + choice := chunk.Choices[0] + if choice.FinishReason != "" { + finishReason = choice.FinishReason + } + + reasoning := choice.Delta.ReasoningContent + if reasoning == "" { + reasoning = choice.Delta.Reasoning + } + if reasoning != "" { + thinking.Emit(out, reasoning) + } + + if choice.Delta.Content != "" { + // A model that inlines its reasoning in tags is routed to + // the same thinking channel, so the UI shows one behavior + // regardless of which convention the vendor picked. + answer, inline := splitter.Feed(choice.Delta.Content) + if inline != "" { + thinking.Emit(out, inline) + } + if answer != "" { + thinking.Finish(out) + out <- types.StreamResponse{ + ResponseType: types.ResponseTypeAnswer, + Content: answer, + } + } + } + if len(choice.Delta.ToolCalls) > 0 { + thinking.Finish(out) + } + } + + thinking.Finish(out) + out <- types.StreamResponse{ + ResponseType: types.ResponseTypeAnswer, + Done: true, + Usage: acc, + FinishReason: finishReason, + } +} + +type streamChunk struct { + Choices []struct { + Delta struct { + Content string `json:"content"` + ReasoningContent string `json:"reasoning_content"` + Reasoning string `json:"reasoning"` + ToolCalls []json.RawMessage `json:"tool_calls"` + } `json:"delta"` + FinishReason string `json:"finish_reason"` + } `json:"choices"` + Usage *usage `json:"usage"` + Error *apiError `json:"error"` +} + +// usage mirrors the token accounting, including the two cache spellings in +// circulation: OpenAI nests cached tokens under prompt_tokens_details, while +// DeepSeek reports hit and miss counters at the top level. +type usage struct { + PromptTokens int `json:"prompt_tokens"` + CompletionTokens int `json:"completion_tokens"` + TotalTokens int `json:"total_tokens"` + PromptCacheHitTokens *int `json:"prompt_cache_hit_tokens"` + PromptCacheMissTokens *int `json:"prompt_cache_miss_tokens"` + PromptTokensDetails *struct { + CachedTokens int `json:"cached_tokens"` + } `json:"prompt_tokens_details"` + CompletionTokensDetails *struct { + ReasoningTokens int `json:"reasoning_tokens"` + } `json:"completion_tokens_details"` +} + +func (u *usage) normalize() types.TokenUsage { + if u == nil { + return types.TokenUsage{} + } + out := types.TokenUsage{ + PromptTokens: u.PromptTokens, + CompletionTokens: u.CompletionTokens, + TotalTokens: u.TotalTokens, + } + if out.TotalTokens == 0 { + out.TotalTokens = out.PromptTokens + out.CompletionTokens + } + + read, reported := 0, false + if u.PromptCacheHitTokens != nil { + read, reported = *u.PromptCacheHitTokens, true + } else if u.PromptTokensDetails != nil { + read, reported = u.PromptTokensDetails.CachedTokens, true + } + miss := out.PromptTokens - read + if u.PromptCacheMissTokens != nil { + miss = *u.PromptCacheMissTokens + } + if miss < 0 { + miss = 0 + } + out.SetPromptCacheUsage(read, 0, miss, reported) + return out +} + +// trimEndpoint keeps a caller-provided base URL from producing a doubled path +// when it already names the endpoint. +func trimEndpoint(baseURL, path string) string { + base := strings.TrimRight(baseURL, "/") + if strings.HasSuffix(base, path) { + return base + } + return base + path +} + +// Endpoint reports the full URL for a base URL. +func (d Driver) Endpoint(baseURL string) string { + return trimEndpoint(baseURL, d.EndpointPath()) +} diff --git a/internal/models/llm/protocol/protocol.go b/internal/models/llm/protocol/protocol.go new file mode 100644 index 00000000000..5c877a4c3fa --- /dev/null +++ b/internal/models/llm/protocol/protocol.go @@ -0,0 +1,123 @@ +// Package protocol implements the wire protocols a model plugin can speak. +// +// A protocol owns four things and nothing else: the canonical request body, +// the endpoint path, how a complete response decodes, and how a stream +// decodes. Everything vendor-specific — which extra fields ride along, which +// standard ones are forbidden, how reasoning is spelled — belongs to the +// vendor descriptor that selects the protocol. +// +// That split is what makes a new OpenAI-compatible vendor a declaration +// instead of an implementation. +package protocol + +import ( + "context" + "fmt" + "io" + "sync" + + "github.com/Tencent/WeKnora/internal/models/llm/spi" + "github.com/Tencent/WeKnora/internal/types" +) + +// Call is one chat request in neutral terms, before any vendor encoding. +type Call struct { + // Model is the model identifier to send. + Model string + // Stream reports whether the response is streamed. + Stream bool + // Messages is the conversation. + Messages []spi.Message + // Options are the caller's generation settings. + Options *spi.Options + // Replay is the descriptor's reasoning-replay rule, which decides whether + // prior-turn reasoning is rendered back onto the wire. + Replay spi.ReasoningReplay +} + +// Driver implements one wire protocol. +type Driver interface { + // ID reports the protocol this driver implements. + ID() spi.ProtocolID + // BuildDraft renders the call into a canonical request body. Vendor + // encoders run over the result, so a driver writes what its protocol + // defines and leaves the extensions alone. + BuildDraft(call Call) (*spi.Draft, error) + // EndpointPath is the path appended to a base URL, including the leading + // slash. A descriptor may override it. + EndpointPath() string + // ParseResponse decodes a complete response body. + ParseResponse(body []byte) (*types.ChatResponse, error) + // DecodeStream decodes a streaming body, sending neutral chunks to out and + // closing nothing: the caller owns the channel's lifetime because it also + // owns the surrounding request. + DecodeStream(ctx context.Context, body io.Reader, out chan<- types.StreamResponse) +} + +// registry holds the available drivers. +var ( + registryMu sync.RWMutex + registry = map[spi.ProtocolID]Driver{} +) + +// Register adds a driver, returning the function that removes it. +func Register(d Driver) func() { + registryMu.Lock() + defer registryMu.Unlock() + registry[d.ID()] = d + id := d.ID() + return func() { + registryMu.Lock() + defer registryMu.Unlock() + delete(registry, id) + } +} + +// Get reports the driver for a protocol. +func Get(id spi.ProtocolID) (Driver, bool) { + registryMu.RLock() + defer registryMu.RUnlock() + d, ok := registry[id] + return d, ok +} + +// MustGet reports the driver for a protocol, erroring when none is registered. +func MustGet(id spi.ProtocolID) (Driver, error) { + d, ok := Get(id) + if !ok { + return nil, fmt.Errorf("no driver registered for protocol %q", id) + } + return d, nil +} + +// IDs reports the registered protocols. +func IDs() []spi.ProtocolID { + registryMu.RLock() + defer registryMu.RUnlock() + out := make([]spi.ProtocolID, 0, len(registry)) + for id := range registry { + out = append(out, id) + } + return out +} + +// ShouldReplayReasoning reports whether a turn's prior reasoning must be +// rendered back onto the wire under a replay rule. +// +// The rule is a vendor fact rather than a preference: DeepSeek answers 400 +// when a tool-calling turn's reasoning is dropped, while replaying it to a +// vendor that does not want it is merely ignored. When in doubt the safe +// direction is to send it, so only an explicit "never" suppresses it. +func ShouldReplayReasoning(replay spi.ReasoningReplay, msg spi.Message) bool { + if msg.ReasoningContent == "" { + return false + } + switch replay { + case spi.ReplayAlways: + return true + case spi.ReplayWithTools: + return len(msg.ToolCalls) > 0 + default: + return false + } +} diff --git a/internal/models/llm/protocol/responses/driver.go b/internal/models/llm/protocol/responses/driver.go new file mode 100644 index 00000000000..124cefbc725 --- /dev/null +++ b/internal/models/llm/protocol/responses/driver.go @@ -0,0 +1,438 @@ +// Package responses implements the OpenAI Responses protocol. +// +// Responses is not a renamed Chat Completions. Input is a list of typed items +// rather than messages, reasoning is an output item of its own instead of a +// field beside the content, tool calls are top-level items rather than a +// property of an assistant message, and the stream is a sequence of named +// events rather than deltas on a choice. +package responses + +import ( + "context" + "encoding/json" + "fmt" + "io" + "strings" + + "github.com/Tencent/WeKnora/internal/models/llm/protocol" + "github.com/Tencent/WeKnora/internal/models/llm/protocol/internal/emit" + "github.com/Tencent/WeKnora/internal/models/llm/spi" + "github.com/Tencent/WeKnora/internal/models/llm/sse" + "github.com/Tencent/WeKnora/internal/types" +) + +// Driver implements protocol.Driver for the Responses API. +type Driver struct{} + +func init() { protocol.Register(Driver{}) } + +// ID reports the protocol. +func (Driver) ID() spi.ProtocolID { return spi.ProtocolOpenAIResponses } + +// EndpointPath reports the standard path. +func (Driver) EndpointPath() string { return "/responses" } + +// BuildDraft renders the call into a Responses body. +func (d Driver) BuildDraft(call protocol.Call) (*spi.Draft, error) { + draft := spi.NewDraft(d.ID(), call.Model, call.Stream) + draft.Set("model", call.Model) + if call.Stream { + draft.Set("stream", true) + } + + instructions, input := convertInput(call.Messages) + if instructions != "" { + draft.Set("instructions", instructions) + } + draft.Set("input", input) + + // Server-side conversation state is opt-in, and storing prompts on the + // vendor's side is a decision for the deployment rather than a default. + draft.Set("store", false) + + opts := call.Options + if opts == nil { + return draft, nil + } + if opts.Temperature > 0 { + draft.Set("temperature", opts.Temperature) + } + if opts.TopP > 0 { + draft.Set("top_p", opts.TopP) + } + if maxTokens := opts.EffectiveMaxTokens(); maxTokens > 0 { + draft.Set("max_output_tokens", maxTokens) + } + if len(opts.Tools) > 0 { + draft.Set("tools", convertTools(opts.Tools)) + } + if opts.ToolChoice != "" { + draft.Set("tool_choice", opts.ToolChoice) + } + if opts.ParallelToolCalls != nil { + draft.Set("parallel_tool_calls", *opts.ParallelToolCalls) + } + return draft, nil +} + +// convertInput splits the system prompt into `instructions` and renders the +// rest as input items. +func convertInput(messages []spi.Message) (string, []any) { + var instructions []string + items := make([]any, 0, len(messages)) + + for _, msg := range messages { + switch msg.Role { + case "system": + if text := plainText(msg); text != "" { + instructions = append(instructions, text) + } + + case "tool": + // A tool result is its own item referencing the call it answers. + items = append(items, map[string]any{ + "type": "function_call_output", + "call_id": msg.ToolCallID, + "output": msg.Content, + }) + + case "assistant": + if content := contentParts(msg, "output_text"); len(content) > 0 { + items = append(items, map[string]any{ + "type": "message", + "role": "assistant", + "content": content, + }) + } + for _, call := range msg.ToolCalls { + items = append(items, map[string]any{ + "type": "function_call", + "call_id": call.ID, + "name": call.Function.Name, + "arguments": call.Function.Arguments, + }) + } + + default: + if content := contentParts(msg, "input_text"); len(content) > 0 { + items = append(items, map[string]any{ + "type": "message", + "role": "user", + "content": content, + }) + } + } + } + return strings.Join(instructions, "\n\n"), items +} + +// contentParts renders a message body as Responses content parts. The text +// part is named differently on input and output items, which is why the caller +// passes the type in. +func contentParts(msg spi.Message, textType string) []any { + var parts []any + for _, part := range msg.MultiContent { + switch part.Type { + case "text": + if part.Text != "" { + parts = append(parts, map[string]any{"type": textType, "text": part.Text}) + } + case "image_url": + if part.ImageURL != nil && part.ImageURL.URL != "" { + parts = append(parts, map[string]any{ + "type": "input_image", + "image_url": part.ImageURL.URL, + }) + } + } + } + for _, img := range msg.Images { + if img != "" { + parts = append(parts, map[string]any{"type": "input_image", "image_url": img}) + } + } + if msg.Content != "" { + parts = append(parts, map[string]any{"type": textType, "text": msg.Content}) + } + return parts +} + +func plainText(msg spi.Message) string { + if msg.Content != "" { + return msg.Content + } + var parts []string + for _, part := range msg.MultiContent { + if part.Type == "text" && part.Text != "" { + parts = append(parts, part.Text) + } + } + return strings.Join(parts, "\n") +} + +// convertTools renders tools in the flat shape Responses uses, where the +// function fields sit on the tool itself rather than in a nested object. +func convertTools(tools []spi.Tool) []any { + out := make([]any, 0, len(tools)) + for _, tool := range tools { + var params any + if len(tool.Function.Parameters) > 0 { + _ = json.Unmarshal(tool.Function.Parameters, ¶ms) + } + out = append(out, map[string]any{ + "type": "function", + "name": tool.Function.Name, + "description": tool.Function.Description, + "parameters": params, + }) + } + return out +} + +// response mirrors a complete Responses body. +type response struct { + Status string `json:"status"` + Output []struct { + Type string `json:"type"` + Content []struct { + Type string `json:"type"` + Text string `json:"text"` + } `json:"content"` + Summary []struct { + Type string `json:"type"` + Text string `json:"text"` + } `json:"summary"` + // Reasoning text is exposed directly by some deployments and only as a + // summary by others. + Text string `json:"text"` + CallID string `json:"call_id"` + Name string `json:"name"` + Arguments string `json:"arguments"` + } `json:"output"` + IncompleteDetails *struct { + Reason string `json:"reason"` + } `json:"incomplete_details"` + Usage *usage `json:"usage"` + Error *struct { + Message string `json:"message"` + } `json:"error"` +} + +// ParseResponse decodes a complete Responses body. +func (d Driver) ParseResponse(body []byte) (*types.ChatResponse, error) { + var resp response + if err := json.Unmarshal(body, &resp); err != nil { + return nil, fmt.Errorf("decode response: %w", err) + } + if resp.Error != nil && resp.Error.Message != "" { + return nil, fmt.Errorf("provider error: %s", resp.Error.Message) + } + + out := &types.ChatResponse{Usage: resp.Usage.normalize()} + var text, reasoning []string + + for _, item := range resp.Output { + switch item.Type { + case "message": + for _, part := range item.Content { + if part.Type == "output_text" { + text = append(text, part.Text) + } + } + case "reasoning": + if item.Text != "" { + reasoning = append(reasoning, item.Text) + } + for _, part := range item.Summary { + if part.Text != "" { + reasoning = append(reasoning, part.Text) + } + } + case "function_call": + out.ToolCalls = append(out.ToolCalls, types.LLMToolCall{ + ID: item.CallID, + Type: "function", + Function: types.FunctionCall{Name: item.Name, Arguments: item.Arguments}, + }) + } + } + + out.Content = strings.Join(text, "") + out.ReasoningContent = strings.Join(reasoning, "") + // A response can stop on the output ceiling before producing visible text, + // and reporting that as a normal completion hides why the answer is empty. + switch { + case resp.IncompleteDetails != nil && resp.IncompleteDetails.Reason != "": + out.FinishReason = resp.IncompleteDetails.Reason + case len(out.ToolCalls) > 0: + out.FinishReason = "tool_calls" + case resp.Status == "completed": + out.FinishReason = "stop" + default: + out.FinishReason = resp.Status + } + return out, nil +} + +// streamEvent mirrors the typed events this driver consumes. Responses names +// roughly forty event types; the ones absent here carry no information this +// seam's neutral stream can express. +type streamEvent struct { + Type string `json:"type"` + Delta string `json:"delta"` + Text string `json:"text"` + Response *response `json:"response"` + Item *struct { + Type string `json:"type"` + CallID string `json:"call_id"` + Name string `json:"name"` + Arguments string `json:"arguments"` + } `json:"item"` + Error *struct { + Message string `json:"message"` + } `json:"error"` +} + +// DecodeStream decodes a Responses SSE stream. +func (d Driver) DecodeStream(ctx context.Context, body io.Reader, out chan<- types.StreamResponse) { + reader := sse.NewReader(body) + thinking := &emit.Thinking{} + + var ( + acc *types.TokenUsage + finishReason string + ) + + for { + select { + case <-ctx.Done(): + emit.Error(out, ctx.Err().Error()) + return + default: + } + + event, err := reader.Next() + if err == io.EOF { + break + } + if err != nil { + emit.Error(out, fmt.Sprintf("read stream: %v", err)) + return + } + if event.Done { + break + } + if len(event.Data) == 0 { + continue + } + + var ev streamEvent + if err := json.Unmarshal(event.Data, &ev); err != nil { + emit.Error(out, fmt.Sprintf("decode stream event: %v", err)) + return + } + // The event name and the payload type agree, but only the payload is + // guaranteed present on every deployment. + kind := ev.Type + if kind == "" { + kind = event.Name + } + + switch kind { + case "response.reasoning_summary_text.delta", "response.reasoning_text.delta": + thinking.Emit(out, ev.Delta) + + case "response.output_text.delta": + thinking.Finish(out) + if ev.Delta != "" { + out <- types.StreamResponse{ + ResponseType: types.ResponseTypeAnswer, + Content: ev.Delta, + } + } + + case "response.output_item.done": + if ev.Item != nil && ev.Item.Type == "function_call" { + thinking.Finish(out) + finishReason = "tool_calls" + } + + case "response.completed", "response.incomplete", "response.failed": + if ev.Response != nil { + if u := ev.Response.Usage.normalize(); u.TotalTokens > 0 { + acc = &u + } + if ev.Response.IncompleteDetails != nil && ev.Response.IncompleteDetails.Reason != "" { + finishReason = ev.Response.IncompleteDetails.Reason + } else if finishReason == "" && ev.Response.Status == "completed" { + finishReason = "stop" + } + if ev.Response.Error != nil && ev.Response.Error.Message != "" { + emit.Error(out, ev.Response.Error.Message) + return + } + } + + case "error": + if ev.Error != nil { + emit.Error(out, ev.Error.Message) + return + } + } + } + + thinking.Finish(out) + out <- types.StreamResponse{ + ResponseType: types.ResponseTypeAnswer, + Done: true, + Usage: acc, + FinishReason: finishReason, + } +} + +// usage mirrors the Responses token accounting, which names its fields for +// input and output rather than prompt and completion. +type usage struct { + InputTokens int `json:"input_tokens"` + OutputTokens int `json:"output_tokens"` + TotalTokens int `json:"total_tokens"` + InputTokensDetails *struct { + CachedTokens int `json:"cached_tokens"` + } `json:"input_tokens_details"` + OutputTokensDetails *struct { + ReasoningTokens int `json:"reasoning_tokens"` + } `json:"output_tokens_details"` +} + +func (u *usage) normalize() types.TokenUsage { + if u == nil { + return types.TokenUsage{} + } + out := types.TokenUsage{ + PromptTokens: u.InputTokens, + CompletionTokens: u.OutputTokens, + TotalTokens: u.TotalTokens, + } + if out.TotalTokens == 0 { + out.TotalTokens = out.PromptTokens + out.CompletionTokens + } + read, reported := 0, false + if u.InputTokensDetails != nil { + read, reported = u.InputTokensDetails.CachedTokens, true + } + miss := out.PromptTokens - read + if miss < 0 { + miss = 0 + } + out.SetPromptCacheUsage(read, 0, miss, reported) + return out +} + +// Endpoint reports the full URL for a base URL. +func (d Driver) Endpoint(baseURL string) string { + base := strings.TrimRight(baseURL, "/") + if strings.HasSuffix(base, d.EndpointPath()) { + return base + } + return base + d.EndpointPath() +} diff --git a/internal/models/llm/spi/message.go b/internal/models/llm/spi/message.go new file mode 100644 index 00000000000..f582a18db34 --- /dev/null +++ b/internal/models/llm/spi/message.go @@ -0,0 +1,190 @@ +package spi + +import ( + "encoding/json" + + "github.com/Tencent/WeKnora/internal/types" +) + +// This file holds the neutral request vocabulary: the messages, tools, and +// options a caller expresses without naming a vendor. Protocol drivers render +// it onto their own wire, so it lives with the seam definition rather than +// inside any one protocol or the chat package that predates them. +// +// The chat package aliases these types, so existing callers are unaffected and +// there is exactly one definition of a message in the codebase. + +// MessageContentPart is one span of a multi-part message, which is how text +// and images travel together. +type MessageContentPart struct { + // Type is "text" or "image_url". + Type string `json:"type"` + // Text carries the span when Type is "text". + Text string `json:"text,omitempty"` + // ImageURL carries the span when Type is "image_url". + ImageURL *ImageURL `json:"image_url,omitempty"` +} + +// ImageURL references an image by URL or inline data URI. +type ImageURL struct { + // URL is an http(s) URL or a base64 data URI. + URL string `json:"url"` + // Detail requests a resolution tier: "auto", "low", or "high". + Detail string `json:"detail,omitempty"` +} + +// FunctionCall is a tool invocation's name and JSON-encoded arguments. +type FunctionCall struct { + Name string `json:"name"` + // Arguments is a JSON object encoded as a string, as every protocol + // transports it. + Arguments string `json:"arguments"` +} + +// ToolCall is one tool invocation an assistant turn produced. +type ToolCall struct { + ID string `json:"id"` + Type string `json:"type"` + + Function FunctionCall `json:"function"` + // ProviderMetadata carries vendor state that must survive a round trip, + // such as Gemini's thought signatures or Anthropic's block signature. + // Dropping it changes the meaning of the replayed turn, so it travels with + // the call rather than being reconstructed. + ProviderMetadata types.ToolCallMetadata `json:"provider_metadata,omitempty"` +} + +// Message is one turn of the conversation. +type Message struct { + // Role is "system", "user", "assistant", or "tool". + Role string `json:"role"` + // Content is the plain-text body. + Content string `json:"content"` + // MultiContent carries a mixed text and image body. + MultiContent []MessageContentPart `json:"multi_content,omitempty"` + // Name is the tool name on a tool-role message. + Name string `json:"name,omitempty"` + // ToolCallID links a tool-role message to the call it answers. + ToolCallID string `json:"tool_call_id,omitempty"` + // ToolCalls are the calls an assistant turn made. + ToolCalls []ToolCall `json:"tool_calls,omitempty"` + // Images are image references attached to the current user turn. + Images []string `json:"images,omitempty"` + // ReasoningContent is the reasoning this assistant turn produced. + // + // Whether it must be replayed is a vendor fact, declared as a descriptor's + // ReasoningReplay: DeepSeek answers 400 when a tool-calling turn's + // reasoning is dropped, Anthropic requires thinking blocks back verbatim, + // and vendors that need neither ignore the field. + ReasoningContent string `json:"reasoning_content,omitempty"` + // ReasoningSignature is the vendor's integrity token over the reasoning, + // replayed alongside it where the vendor issues one. + ReasoningSignature string `json:"reasoning_signature,omitempty"` +} + +// FunctionDef declares a callable tool. +type FunctionDef struct { + Name string `json:"name"` + Description string `json:"description"` + // Parameters is a JSON Schema object. + Parameters json.RawMessage `json:"parameters"` +} + +// Tool is a tool the model may call. +type Tool struct { + // Type is "function". + Type string `json:"type"` + Function FunctionDef `json:"function"` +} + +// Options are the generation settings a caller supplies. +// +// It keeps the flat, vendor-neutral shape callers already use. The plugin layer +// translates it into parameter values, which is where the vendor differences +// live — a caller sets Thinking, and whether that becomes enable_thinking, +// thinking.type, or a reasoning effort rung is the descriptor's business. +type Options struct { + Temperature float64 `json:"temperature"` + TopP float64 `json:"top_p"` + Seed int `json:"seed"` + MaxTokens int `json:"max_tokens"` + MaxCompletionTokens int `json:"max_completion_tokens"` + FrequencyPenalty float64 `json:"frequency_penalty"` + PresencePenalty float64 `json:"presence_penalty"` + // Thinking is the neutral reasoning toggle; nil defers to the model. + Thinking *bool `json:"thinking"` + // ThinkingEffort and ThinkingBudget are the depth controls, empty or zero + // when the caller has no opinion. A value a vendor does not accept is + // reported in the plan rather than sent. + ThinkingEffort string `json:"thinking_effort,omitempty"` + ThinkingBudget int `json:"thinking_budget,omitempty"` + + Tools []Tool `json:"tools,omitempty"` + ToolChoice string `json:"tool_choice,omitempty"` + ParallelToolCalls *bool `json:"parallel_tool_calls,omitempty"` + Format json.RawMessage `json:"format,omitempty"` +} + +// ParamValues translates the caller's options into neutral parameter values. +// +// Only fields the caller actually set become values: a zero temperature is +// indistinguishable from an unset one in this flat struct, and sending it +// would override a model default the caller never meant to touch. That +// asymmetry is the price of the legacy shape, and confining it to this one +// function keeps it out of every plugin. +func (o *Options) ParamValues() map[ParamID]Value { + values := map[ParamID]Value{} + if o == nil { + return values + } + if o.Temperature > 0 { + values[ParamTemperature] = FloatValue(o.Temperature) + } + if o.TopP > 0 { + values[ParamTopP] = FloatValue(o.TopP) + } + if o.FrequencyPenalty != 0 { + values[ParamFrequencyPenalty] = FloatValue(o.FrequencyPenalty) + } + if o.PresencePenalty != 0 { + values[ParamPresencePenalty] = FloatValue(o.PresencePenalty) + } + if o.Seed != 0 { + values[ParamSeed] = IntValue(o.Seed) + } + if maxTokens := o.EffectiveMaxTokens(); maxTokens > 0 { + values[ParamMaxTokens] = IntValue(maxTokens) + } + if o.Thinking != nil { + mode := ThinkingOff + if *o.Thinking { + mode = ThinkingOn + } + values[ParamThinkingMode] = EnumValue(mode) + } + if o.ThinkingEffort != "" { + values[ParamThinkingEffort] = EnumValue(o.ThinkingEffort) + } + if o.ThinkingBudget > 0 { + values[ParamThinkingBudget] = IntValue(o.ThinkingBudget) + } + if o.ToolChoice != "" { + values[ParamToolChoice] = EnumValue(o.ToolChoice) + } + if o.ParallelToolCalls != nil { + values[ParamParallelToolCalls] = BoolValue(*o.ParallelToolCalls) + } + return values +} + +// EffectiveMaxTokens reports the output ceiling, accepting either spelling the +// caller may have used. +func (o *Options) EffectiveMaxTokens() int { + if o == nil { + return 0 + } + if o.MaxTokens > 0 { + return o.MaxTokens + } + return o.MaxCompletionTokens +} diff --git a/internal/models/llm/sse/reader.go b/internal/models/llm/sse/reader.go new file mode 100644 index 00000000000..5687e04bd50 --- /dev/null +++ b/internal/models/llm/sse/reader.go @@ -0,0 +1,109 @@ +// Package sse reads Server-Sent Events streams from model APIs. +// +// It keeps the event name alongside the payload. The OpenAI Chat Completions +// protocol needs only the data lines, but the Responses and Anthropic +// protocols are event-typed, and a reader that discards `event:` forces every +// consumer to re-derive the type from the payload — which works only as long +// as every vendor also repeats it there. +package sse + +import ( + "bufio" + "io" + "strings" +) + +// maxLineBytes bounds one event line. Reasoning-capable models emit very long +// single-line payloads, and the default scanner limit truncates them into +// invalid JSON. +const maxLineBytes = 1024 * 1024 + +// Event is one parsed Server-Sent Event. +type Event struct { + // Name is the `event:` field, empty when the stream does not send one. + Name string + // Data is the concatenated `data:` payload. + Data []byte + // Done reports the `[DONE]` sentinel that OpenAI-shaped streams end with. + Done bool +} + +// Reader parses an SSE stream. +type Reader struct { + scanner *bufio.Scanner +} + +// NewReader returns a reader over an SSE body. +func NewReader(r io.Reader) *Reader { + scanner := bufio.NewScanner(r) + scanner.Buffer(make([]byte, maxLineBytes), maxLineBytes) + return &Reader{scanner: scanner} +} + +// Next returns the next event, or io.EOF when the stream ends. +// +// Per the SSE specification an event is terminated by a blank line and may +// carry several data lines, which are joined with newlines. Vendors mostly +// send one line per event, but honoring the specification costs nothing and +// avoids truncating the ones that do not. +func (r *Reader) Next() (*Event, error) { + event := &Event{} + var data []string + + for r.scanner.Scan() { + line := r.scanner.Text() + + if line == "" { + // Blank line dispatches the event, unless nothing accumulated. + if len(data) == 0 && event.Name == "" { + continue + } + return finish(event, data), nil + } + if strings.HasPrefix(line, ":") { + // A comment line, commonly used as a keep-alive. + continue + } + + field, value := splitField(line) + switch field { + case "event": + event.Name = value + case "data": + if value == "[DONE]" { + event.Done = true + return event, nil + } + data = append(data, value) + default: + // `id:` and `retry:` carry no meaning for these APIs. + } + } + + if err := r.scanner.Err(); err != nil { + return nil, err + } + // A stream that ends without a trailing blank line still owes its last + // event, which is common when a connection closes right after the payload. + if len(data) > 0 || event.Name != "" { + return finish(event, data), nil + } + return nil, io.EOF +} + +func finish(event *Event, data []string) *Event { + event.Data = []byte(strings.Join(data, "\n")) + return event +} + +// splitField splits an SSE line into its field name and value, tolerating a +// missing space after the colon. +func splitField(line string) (field, value string) { + idx := strings.IndexByte(line, ':') + if idx < 0 { + return line, "" + } + field = line[:idx] + value = line[idx+1:] + return field, strings.TrimPrefix(value, " ") +} From 2025fc70b684264977b26c377a791b9c637e99d1 Mon Sep 17 00:00:00 2001 From: lyingbug <[email> Date: Fri, 14 Aug 2026 00:04:03 +0000 Subject: [PATCH 3/9] refactor(models): drive chat requests through the model-plugin seam Replace the provider-adapter branches and the four-value thinking_control enum with the declarative plugins. The chat layer now resolves a descriptor and applies its plan; which field carries a toggle, which sampling knobs a model forbids, and which value is pinned are all declared per vendor. - chat.Message and friends become aliases of the seam vocabulary, so there is one definition of a message and protocol drivers can render it. - provider.go keeps only what a declaration cannot express: signing, a non-derived endpoint, a message rewrite, and tool-call metadata. - The OpenAI-compatible transport applies the plan to the request's own JSON and compares the bytes to decide whether the SDK path still fits, instead of maintaining a second list of providers that need raw HTTP. - Anthropic routes through the Messages driver, replacing a text-only implementation that had no tools, no thinking blocks, and no images. The Responses protocol becomes selectable per model. - A stored thinking_control still forces the wire format it names, so existing configurations keep their behavior. Two real bugs surfaced and are fixed: an unset thinking mode was treated as off, which discarded a depth setting the request would have honored, and max_completion_tokens lost to max_tokens when both were set. Legacy AnthropicChat is deleted; its coverage is replaced by end-to-end tests against the new driver that also cover tools, thinking replay, and the base-URL spellings users paste. Co-authored-by: lyingbug --- internal/models/chat/anthropic.go | 537 ----------------------- internal/models/chat/anthropic_test.go | 269 +----------- internal/models/chat/chat.go | 97 ++-- internal/models/chat/plugin_chat.go | 314 +++++++++++++ internal/models/chat/plugin_chat_test.go | 291 ++++++++++++ internal/models/chat/provider.go | 149 +------ internal/models/chat/provider_test.go | 80 ++-- internal/models/chat/remote_api.go | 131 ++++-- internal/models/chat/remote_api_test.go | 48 +- internal/models/chat/thinking.go | 254 +++++------ internal/models/chat/thinking_test.go | 214 ++++----- internal/models/llm/spi/message.go | 9 +- 12 files changed, 1050 insertions(+), 1343 deletions(-) delete mode 100644 internal/models/chat/anthropic.go create mode 100644 internal/models/chat/plugin_chat.go create mode 100644 internal/models/chat/plugin_chat_test.go diff --git a/internal/models/chat/anthropic.go b/internal/models/chat/anthropic.go deleted file mode 100644 index 5da902f147c..00000000000 --- a/internal/models/chat/anthropic.go +++ /dev/null @@ -1,537 +0,0 @@ -package chat - -import ( - "bytes" - "context" - "encoding/json" - "fmt" - "io" - "net/http" - "net/url" - "strings" - - "github.com/Tencent/WeKnora/internal/models/provider" - "github.com/Tencent/WeKnora/internal/types" - secutils "github.com/Tencent/WeKnora/internal/utils" -) - -const anthropicVersion = "2023-06-01" - -type AnthropicChat struct { - modelName string - modelID string - baseURL string - apiKey string - customHeaders map[string]string -} - -type anthropicMessage struct { - Role string `json:"role"` - Content string `json:"content"` -} - -type anthropicRequest struct { - Model string `json:"model"` - MaxTokens int `json:"max_tokens"` - Stream bool `json:"stream,omitempty"` - System string `json:"system,omitempty"` - Messages []anthropicMessage `json:"messages"` - Temperature *float64 `json:"temperature,omitempty"` - TopP *float64 `json:"top_p,omitempty"` -} - -type anthropicResponse struct { - ID string `json:"id"` - Type string `json:"type"` - Role string `json:"role"` - Content []struct { - Type string `json:"type"` - Text string `json:"text"` - } `json:"content"` - StopReason string `json:"stop_reason"` - Usage struct { - InputTokens int `json:"input_tokens"` - OutputTokens int `json:"output_tokens"` - CacheCreationInputTokens *int `json:"cache_creation_input_tokens"` - CacheReadInputTokens *int `json:"cache_read_input_tokens"` - } `json:"usage"` - Error *struct { - Type string `json:"type"` - Message string `json:"message"` - } `json:"error,omitempty"` -} - -type anthropicStreamEvent struct { - Type string `json:"type"` - Message *struct { - Usage struct { - InputTokens int `json:"input_tokens"` - OutputTokens int `json:"output_tokens"` - CacheCreationInputTokens *int `json:"cache_creation_input_tokens"` - CacheReadInputTokens *int `json:"cache_read_input_tokens"` - } `json:"usage"` - } `json:"message,omitempty"` - Delta *struct { - Type string `json:"type"` - Text string `json:"text"` - StopReason string `json:"stop_reason"` - } `json:"delta,omitempty"` - Usage *struct { - InputTokens int `json:"input_tokens"` - OutputTokens int `json:"output_tokens"` - CacheCreationInputTokens *int `json:"cache_creation_input_tokens"` - CacheReadInputTokens *int `json:"cache_read_input_tokens"` - } `json:"usage,omitempty"` - Error *struct { - Type string `json:"type"` - Message string `json:"message"` - } `json:"error,omitempty"` -} - -func NewAnthropicChat(config *ChatConfig) (*AnthropicChat, error) { - if config.BaseURL != "" { - if err := secutils.ValidateURLForSSRF(config.BaseURL); err != nil { - return nil, fmt.Errorf("baseURL SSRF check failed: %w", err) - } - } - if strings.TrimSpace(config.APIKey) == "" { - return nil, fmt.Errorf("Anthropic provider: API key is required") - } - - baseURL := strings.TrimRight(config.BaseURL, "/") - if baseURL == "" { - baseURL = provider.AnthropicBaseURL - } - - return &AnthropicChat{ - modelName: config.ModelName, - modelID: config.ModelID, - baseURL: baseURL, - apiKey: config.APIKey, - customHeaders: config.CustomHeaders, - }, nil -} - -func (c *AnthropicChat) Chat(ctx context.Context, messages []Message, opts *ChatOptions) (*types.ChatResponse, error) { - reqBody := c.buildRequest(messages, opts) - jsonData, err := json.Marshal(reqBody) - if err != nil { - return nil, fmt.Errorf("marshal request: %w", err) - } - - ctx, cancel := withLLMTimeout(ctx, defaultChatTimeout) - defer cancel() - - endpoint := c.endpoint() - if err := secutils.ValidateURLForSSRF(endpoint); err != nil { - return nil, fmt.Errorf("endpoint SSRF check failed: %w", err) - } - - httpReq, err := http.NewRequestWithContext(ctx, http.MethodPost, endpoint, bytes.NewBuffer(jsonData)) - if err != nil { - return nil, fmt.Errorf("create request: %w", err) - } - httpReq.Header.Set("Content-Type", "application/json") - httpReq.Header.Set("x-api-key", c.apiKey) - httpReq.Header.Set("anthropic-version", anthropicVersion) - secutils.ApplyCustomHeaders(httpReq, c.customHeaders) - - resp, err := rawHTTPClient.Do(httpReq) - if err != nil { - return nil, fmt.Errorf("send request: %w", err) - } - defer resp.Body.Close() - - body, err := io.ReadAll(resp.Body) - if err != nil { - return nil, fmt.Errorf("read response: %w", err) - } - - if strings.Contains(strings.ToLower(resp.Header.Get("Content-Type")), "text/event-stream") { - chatResp, err := parseAnthropicSSE(bytes.NewReader(body)) - if err != nil { - return nil, err - } - if resp.StatusCode < http.StatusOK || resp.StatusCode >= http.StatusMultipleChoices { - return nil, fmt.Errorf("API request failed with status %d: %s", resp.StatusCode, chatResp.Content) - } - logUsage(ctx, c.modelName, &chatResp.Usage) - return chatResp, nil - } - - var chatResp anthropicResponse - if err := json.Unmarshal(body, &chatResp); err != nil { - return nil, fmt.Errorf("decode response: %w", err) - } - if resp.StatusCode < http.StatusOK || resp.StatusCode >= http.StatusMultipleChoices { - if chatResp.Error != nil && chatResp.Error.Message != "" { - return nil, fmt.Errorf("API request failed with status %d: %s", resp.StatusCode, chatResp.Error.Message) - } - return nil, fmt.Errorf("API request failed with status %d: %s", resp.StatusCode, string(body)) - } - - result := c.parseResponse(&chatResp) - logUsage(ctx, c.modelName, &result.Usage) - return result, nil -} - -func (c *AnthropicChat) ChatStream(ctx context.Context, messages []Message, opts *ChatOptions) (<-chan types.StreamResponse, error) { - reqBody := c.buildRequest(messages, opts) - reqBody.Stream = true - jsonData, err := json.Marshal(reqBody) - if err != nil { - return nil, fmt.Errorf("marshal request: %w", err) - } - - endpoint := c.endpoint() - if err := secutils.ValidateURLForSSRF(endpoint); err != nil { - return nil, fmt.Errorf("endpoint SSRF check failed: %w", err) - } - - httpReq, err := http.NewRequestWithContext(ctx, http.MethodPost, endpoint, bytes.NewBuffer(jsonData)) - if err != nil { - return nil, fmt.Errorf("create request: %w", err) - } - httpReq.Header.Set("Content-Type", "application/json") - httpReq.Header.Set("Accept", "text/event-stream") - httpReq.Header.Set("x-api-key", c.apiKey) - httpReq.Header.Set("anthropic-version", anthropicVersion) - secutils.ApplyCustomHeaders(httpReq, c.customHeaders) - - resp, err := rawHTTPClient.Do(httpReq) - if err != nil { - return nil, fmt.Errorf("send request: %w", err) - } - if resp.StatusCode < http.StatusOK || resp.StatusCode >= http.StatusMultipleChoices { - body, _ := io.ReadAll(resp.Body) - resp.Body.Close() - return nil, fmt.Errorf("API request failed with status %d: %s", resp.StatusCode, string(body)) - } - - streamChan := make(chan types.StreamResponse) - go processAnthropicStream(ctx, c.modelName, resp, streamChan) - return streamChan, nil -} - -func (c *AnthropicChat) GetModelName() string { - return c.modelName -} - -func (c *AnthropicChat) GetModelID() string { - return c.modelID -} - -func (c *AnthropicChat) endpoint() string { - baseURL := strings.TrimRight(c.baseURL, "/") - if isAnthropicMessagesEndpoint(baseURL) { - return baseURL - } - if isAnthropicVersionedBaseURL(baseURL) { - return baseURL + "/messages" - } - return baseURL + "/v1/messages" -} - -func isAnthropicMessagesEndpoint(baseURL string) bool { - u, err := url.Parse(baseURL) - if err != nil { - return false - } - path := strings.TrimRight(u.Path, "/") - return strings.HasSuffix(path, "/messages") -} - -func isAnthropicVersionedBaseURL(baseURL string) bool { - u, err := url.Parse(baseURL) - if err != nil { - return false - } - path := strings.TrimRight(u.Path, "/") - return strings.HasSuffix(path, "/v1") || strings.HasSuffix(path, "/v1beta") -} - -func (c *AnthropicChat) buildRequest(messages []Message, opts *ChatOptions) anthropicRequest { - req := anthropicRequest{ - Model: c.modelName, - MaxTokens: 1024, - Messages: make([]anthropicMessage, 0, len(messages)), - } - if opts != nil { - if opts.MaxTokens > 0 { - req.MaxTokens = opts.MaxTokens - } else if opts.MaxCompletionTokens > 0 { - req.MaxTokens = opts.MaxCompletionTokens - } - if opts.Temperature > 0 { - temperature := opts.Temperature - req.Temperature = &temperature - } - if opts.TopP > 0 { - topP := opts.TopP - req.TopP = &topP - } - } - - var systemParts []string - for _, msg := range messages { - content := strings.TrimSpace(msg.Content) - if content == "" { - content = textFromMultiContent(msg.MultiContent) - } - if content == "" { - continue - } - switch msg.Role { - case "system": - systemParts = append(systemParts, content) - case "assistant": - req.Messages = append(req.Messages, anthropicMessage{Role: "assistant", Content: content}) - case "user": - req.Messages = append(req.Messages, anthropicMessage{Role: "user", Content: content}) - default: - req.Messages = append(req.Messages, anthropicMessage{Role: "user", Content: content}) - } - } - req.System = strings.Join(systemParts, "\n\n") - return req -} - -func textFromMultiContent(parts []MessageContentPart) string { - if len(parts) == 0 { - return "" - } - textParts := make([]string, 0, len(parts)) - for _, part := range parts { - if part.Type == "text" && strings.TrimSpace(part.Text) != "" { - textParts = append(textParts, strings.TrimSpace(part.Text)) - } - } - return strings.Join(textParts, "\n") -} - -func (c *AnthropicChat) parseResponse(resp *anthropicResponse) *types.ChatResponse { - parts := make([]string, 0, len(resp.Content)) - for _, part := range resp.Content { - if part.Type == "text" && part.Text != "" { - parts = append(parts, part.Text) - } - } - inputTokens := resp.Usage.InputTokens - outputTokens := resp.Usage.OutputTokens - cacheRead := valueOrZero(resp.Usage.CacheReadInputTokens) - cacheWrite := valueOrZero(resp.Usage.CacheCreationInputTokens) - promptTokens := inputTokens + cacheRead + cacheWrite - usage := types.TokenUsage{ - PromptTokens: promptTokens, - CompletionTokens: outputTokens, - TotalTokens: promptTokens + outputTokens, - } - usage.SetPromptCacheUsage(cacheRead, cacheWrite, max(0, promptTokens-cacheRead), - resp.Usage.CacheReadInputTokens != nil || resp.Usage.CacheCreationInputTokens != nil) - return &types.ChatResponse{ - Content: strings.Join(parts, ""), - FinishReason: resp.StopReason, - Usage: usage, - } -} - -func parseAnthropicSSE(reader io.Reader) (*types.ChatResponse, error) { - sseReader := NewSSEReader(reader) - var contentParts []string - var finishReason string - var inputTokens int - var outputTokens int - var cacheReadTokens int - var cacheWriteTokens int - var cacheReported bool - - for { - event, err := sseReader.ReadEvent() - if err == io.EOF { - break - } - if err != nil { - return nil, fmt.Errorf("read SSE response: %w", err) - } - if event.Done { - break - } - if len(event.Data) == 0 { - continue - } - - var streamEvent anthropicStreamEvent - if err := json.Unmarshal(event.Data, &streamEvent); err != nil { - return nil, fmt.Errorf("decode SSE response: %w", err) - } - if streamEvent.Error != nil && streamEvent.Error.Message != "" { - return nil, fmt.Errorf("API stream error: %s", streamEvent.Error.Message) - } - if streamEvent.Message != nil { - inputTokens = max(inputTokens, streamEvent.Message.Usage.InputTokens) - outputTokens = max(outputTokens, streamEvent.Message.Usage.OutputTokens) - cacheReadTokens, cacheWriteTokens, cacheReported = mergeAnthropicCacheCounters( - cacheReadTokens, cacheWriteTokens, cacheReported, - streamEvent.Message.Usage.CacheReadInputTokens, - streamEvent.Message.Usage.CacheCreationInputTokens, - ) - } - if streamEvent.Delta != nil { - if streamEvent.Delta.Type == "text_delta" && streamEvent.Delta.Text != "" { - contentParts = append(contentParts, streamEvent.Delta.Text) - } - if streamEvent.Delta.StopReason != "" { - finishReason = streamEvent.Delta.StopReason - } - } - if streamEvent.Usage != nil { - inputTokens = max(inputTokens, streamEvent.Usage.InputTokens) - outputTokens = max(outputTokens, streamEvent.Usage.OutputTokens) - cacheReadTokens, cacheWriteTokens, cacheReported = mergeAnthropicCacheCounters( - cacheReadTokens, cacheWriteTokens, cacheReported, - streamEvent.Usage.CacheReadInputTokens, - streamEvent.Usage.CacheCreationInputTokens, - ) - } - } - promptTokens := inputTokens + cacheReadTokens + cacheWriteTokens - usage := types.TokenUsage{ - PromptTokens: promptTokens, - CompletionTokens: outputTokens, - TotalTokens: promptTokens + outputTokens, - } - usage.SetPromptCacheUsage(cacheReadTokens, cacheWriteTokens, - max(0, promptTokens-cacheReadTokens), cacheReported) - - return &types.ChatResponse{ - Content: strings.Join(contentParts, ""), - FinishReason: finishReason, - Usage: usage, - }, nil -} - -func processAnthropicStream(ctx context.Context, model string, resp *http.Response, streamChan chan types.StreamResponse) { - defer close(streamChan) - defer resp.Body.Close() - - sseReader := NewSSEReader(resp.Body) - var usage *types.TokenUsage - var finishReason string - - for { - event, err := sseReader.ReadEvent() - if err != nil { - if err == io.EOF { - logUsage(ctx, model, usage) - streamChan <- types.StreamResponse{ - ResponseType: types.ResponseTypeAnswer, - Content: "", - Done: true, - Usage: usage, - FinishReason: finishReason, - } - } else { - streamChan <- types.StreamResponse{ - ResponseType: types.ResponseTypeError, - Content: err.Error(), - Done: true, - } - } - return - } - if event.Done { - logUsage(ctx, model, usage) - streamChan <- types.StreamResponse{ - ResponseType: types.ResponseTypeAnswer, - Content: "", - Done: true, - Usage: usage, - FinishReason: finishReason, - } - return - } - if len(event.Data) == 0 { - continue - } - - var streamEvent anthropicStreamEvent - if err := json.Unmarshal(event.Data, &streamEvent); err != nil { - streamChan <- types.StreamResponse{ - ResponseType: types.ResponseTypeError, - Content: fmt.Sprintf("decode SSE response: %v", err), - Done: true, - } - return - } - if streamEvent.Error != nil && streamEvent.Error.Message != "" { - streamChan <- types.StreamResponse{ - ResponseType: types.ResponseTypeError, - Content: streamEvent.Error.Message, - Done: true, - } - return - } - if streamEvent.Message != nil { - usage = mergeAnthropicUsage(usage, streamEvent.Message.Usage.InputTokens, - streamEvent.Message.Usage.OutputTokens, - streamEvent.Message.Usage.CacheReadInputTokens, - streamEvent.Message.Usage.CacheCreationInputTokens) - } - if streamEvent.Delta != nil { - if streamEvent.Delta.StopReason != "" { - finishReason = streamEvent.Delta.StopReason - } - if streamEvent.Delta.Type == "text_delta" && streamEvent.Delta.Text != "" { - streamChan <- types.StreamResponse{ - ResponseType: types.ResponseTypeAnswer, - Content: streamEvent.Delta.Text, - Done: false, - } - } - } - if streamEvent.Usage != nil { - usage = mergeAnthropicUsage(usage, streamEvent.Usage.InputTokens, - streamEvent.Usage.OutputTokens, - streamEvent.Usage.CacheReadInputTokens, - streamEvent.Usage.CacheCreationInputTokens) - } - } -} - -func mergeAnthropicUsage( - current *types.TokenUsage, - inputTokens, outputTokens int, - cacheRead, cacheWrite *int, -) *types.TokenUsage { - if current == nil { - current = &types.TokenUsage{} - } - read, write, reported := mergeAnthropicCacheCounters( - current.CacheReadTokens, current.CacheWriteTokens, current.CacheReported, - cacheRead, cacheWrite, - ) - uncachedInput := max(0, current.PromptTokens-current.CacheReadTokens-current.CacheWriteTokens) - uncachedInput = max(uncachedInput, inputTokens) - current.PromptTokens = uncachedInput + read + write - current.CompletionTokens = max(current.CompletionTokens, outputTokens) - current.TotalTokens = current.PromptTokens + current.CompletionTokens - current.SetPromptCacheUsage(read, write, max(0, current.PromptTokens-read), reported) - return current -} - -func mergeAnthropicCacheCounters( - currentRead, currentWrite int, - currentReported bool, - cacheRead, cacheWrite *int, -) (read, write int, reported bool) { - read = currentRead - write = currentWrite - reported = currentReported || cacheRead != nil || cacheWrite != nil - if cacheRead != nil { - read = max(read, *cacheRead) - } - if cacheWrite != nil { - write = max(write, *cacheWrite) - } - return read, write, reported -} diff --git a/internal/models/chat/anthropic_test.go b/internal/models/chat/anthropic_test.go index 1651caa15ca..79d2691c7ab 100644 --- a/internal/models/chat/anthropic_test.go +++ b/internal/models/chat/anthropic_test.go @@ -1,277 +1,30 @@ package chat import ( - "context" - "encoding/json" - "net/http" - "net/http/httptest" "testing" + "github.com/Tencent/WeKnora/internal/models/llm/spi" "github.com/Tencent/WeKnora/internal/models/provider" "github.com/Tencent/WeKnora/internal/types" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) -func TestAnthropicChat(t *testing.T) { - t.Setenv("SSRF_WHITELIST", "127.0.0.1") - - var capturedHeaders http.Header - var capturedRequest anthropicRequest - - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - assert.Equal(t, "/v1/messages", r.URL.Path) - capturedHeaders = r.Header.Clone() - require.NoError(t, json.NewDecoder(r.Body).Decode(&capturedRequest)) - - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{ - "id":"msg_123", - "type":"message", - "role":"assistant", - "content":[{"type":"text","text":"hello"}], - "stop_reason":"end_turn", - "usage":{"input_tokens":3,"output_tokens":2} - }`)) - })) - defer server.Close() - - chat, err := NewAnthropicChat(&ChatConfig{ +// Anthropic now routes through the plugin client, which speaks the full +// Messages protocol — content blocks, tool use, and thinking blocks — instead +// of the text-only implementation this replaced. +func TestNewRemoteChat_AnthropicUsesTheMessagesPlugin(t *testing.T) { + client, err := NewRemoteChat(&ChatConfig{ Source: types.ModelSourceRemote, - BaseURL: server.URL, ModelName: "claude-sonnet-4-5", APIKey: "test-key", Provider: string(provider.ProviderAnthropic), - CustomHeaders: map[string]string{ - "anthropic-beta": "test-beta", - }, - }) - require.NoError(t, err) - - resp, err := chat.Chat(context.Background(), []Message{ - {Role: "system", Content: "You are helpful."}, - {Role: "user", Content: "Hi"}, - }, &ChatOptions{MaxTokens: 7, Temperature: 0.2}) - require.NoError(t, err) - - assert.Equal(t, "test-key", capturedHeaders.Get("x-api-key")) - assert.Equal(t, anthropicVersion, capturedHeaders.Get("anthropic-version")) - assert.Equal(t, "test-beta", capturedHeaders.Get("anthropic-beta")) - assert.Equal(t, "claude-sonnet-4-5", capturedRequest.Model) - assert.Equal(t, 7, capturedRequest.MaxTokens) - assert.Equal(t, "You are helpful.", capturedRequest.System) - require.Len(t, capturedRequest.Messages, 1) - assert.Equal(t, "user", capturedRequest.Messages[0].Role) - assert.Equal(t, "Hi", capturedRequest.Messages[0].Content) - assert.Equal(t, "hello", resp.Content) - assert.Equal(t, "end_turn", resp.FinishReason) - assert.Equal(t, 3, resp.Usage.PromptTokens) - assert.Equal(t, 2, resp.Usage.CompletionTokens) - assert.Equal(t, 5, resp.Usage.TotalTokens) -} - -func TestAnthropicChat_CacheUsage(t *testing.T) { - t.Setenv("SSRF_WHITELIST", "127.0.0.1") - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{ - "id":"msg_cached","type":"message","role":"assistant", - "content":[{"type":"text","text":"hello"}],"stop_reason":"end_turn", - "usage":{"input_tokens":24,"output_tokens":2,"cache_creation_input_tokens":100,"cache_read_input_tokens":900} - }`)) - })) - defer server.Close() - - chat, err := NewAnthropicChat(&ChatConfig{ - Source: types.ModelSourceRemote, BaseURL: server.URL, ModelName: "claude-sonnet-4-5", - APIKey: "test-key", Provider: string(provider.ProviderAnthropic), }) require.NoError(t, err) - resp, err := chat.Chat(context.Background(), []Message{{Role: "user", Content: "Hi"}}, nil) - require.NoError(t, err) - assert.Equal(t, 1024, resp.Usage.PromptTokens) - assert.Equal(t, 900, resp.Usage.CacheReadTokens) - assert.Equal(t, 100, resp.Usage.CacheWriteTokens) - assert.Equal(t, 124, resp.Usage.CacheMissTokens) - assert.True(t, resp.Usage.CacheReported) - assert.Equal(t, types.PromptCacheStatusHit, resp.Usage.CacheStatus) -} - -func TestAnthropicChat_FullEndpoint(t *testing.T) { - t.Setenv("SSRF_WHITELIST", "127.0.0.1") - - var capturedPath string - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - capturedPath = r.URL.Path - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{ - "id":"msg_123", - "type":"message", - "role":"assistant", - "content":[{"type":"text","text":"hello"}], - "stop_reason":"end_turn", - "usage":{"input_tokens":3,"output_tokens":2} - }`)) - })) - defer server.Close() - - chat, err := NewAnthropicChat(&ChatConfig{ - Source: types.ModelSourceRemote, - BaseURL: server.URL + "/api/proxy/forward", - ModelName: "gpt-5.5", - APIKey: "test-key", - Provider: string(provider.ProviderAnthropic), - }) - require.NoError(t, err) - - resp, err := chat.Chat(context.Background(), []Message{{Role: "user", Content: "Hi"}}, nil) - require.NoError(t, err) - - assert.Equal(t, "/api/proxy/forward/v1/messages", capturedPath) - assert.Equal(t, "hello", resp.Content) -} - -func TestAnthropicChat_MessagesEndpoint(t *testing.T) { - t.Setenv("SSRF_WHITELIST", "127.0.0.1") - - var capturedPath string - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - capturedPath = r.URL.Path - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{ - "id":"msg_123", - "type":"message", - "role":"assistant", - "content":[{"type":"text","text":"hello"}], - "stop_reason":"end_turn", - "usage":{"input_tokens":3,"output_tokens":2} - }`)) - })) - defer server.Close() - - chat, err := NewAnthropicChat(&ChatConfig{ - Source: types.ModelSourceRemote, - BaseURL: server.URL + "/api/v1/messages", - ModelName: "gpt-5.5", - APIKey: "test-key", - Provider: string(provider.ProviderAnthropic), - }) - require.NoError(t, err) - - resp, err := chat.Chat(context.Background(), []Message{{Role: "user", Content: "Hi"}}, nil) - require.NoError(t, err) - - assert.Equal(t, "/api/v1/messages", capturedPath) - assert.Equal(t, "hello", resp.Content) -} - -func TestAnthropicChat_SSEResponse(t *testing.T) { - t.Setenv("SSRF_WHITELIST", "127.0.0.1") - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - assert.Equal(t, "/v1/messages", r.URL.Path) - w.Header().Set("Content-Type", "text/event-stream") - _, _ = w.Write([]byte(`event: message_start -data: {"type":"message_start","message":{"usage":{"input_tokens":14,"output_tokens":0,"cache_creation_input_tokens":0,"cache_read_input_tokens":100}}} - -event: content_block_delta -data: {"type":"content_block_delta","delta":{"type":"text_delta","text":"pong"}} - -event: message_delta -data: {"type":"message_delta","delta":{"stop_reason":"end_turn"},"usage":{"output_tokens":5}} - -event: message_stop -data: {"type":"message_stop"} - -`)) - })) - defer server.Close() - - chat, err := NewAnthropicChat(&ChatConfig{ - Source: types.ModelSourceRemote, - BaseURL: server.URL, - ModelName: "gpt-5.5", - APIKey: "test-key", - Provider: string(provider.ProviderAnthropic), - }) - require.NoError(t, err) - - resp, err := chat.Chat(context.Background(), []Message{{Role: "user", Content: "ping"}}, nil) - require.NoError(t, err) - - assert.Equal(t, "pong", resp.Content) - assert.Equal(t, "end_turn", resp.FinishReason) - assert.Equal(t, 114, resp.Usage.PromptTokens) - assert.Equal(t, 5, resp.Usage.CompletionTokens) - assert.Equal(t, 119, resp.Usage.TotalTokens) - assert.Equal(t, 100, resp.Usage.CacheReadTokens) - assert.Equal(t, 14, resp.Usage.CacheMissTokens) -} - -func TestAnthropicChat_ChatStream(t *testing.T) { - t.Setenv("SSRF_WHITELIST", "127.0.0.1") - - var capturedRequest anthropicRequest - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - assert.Equal(t, "/v1/messages", r.URL.Path) - assert.Equal(t, "text/event-stream", r.Header.Get("Accept")) - require.NoError(t, json.NewDecoder(r.Body).Decode(&capturedRequest)) - w.Header().Set("Content-Type", "text/event-stream") - _, _ = w.Write([]byte(`event: message_start -data: {"type":"message_start","message":{"usage":{"input_tokens":14,"output_tokens":0,"cache_creation_input_tokens":0,"cache_read_input_tokens":100}}} - -event: content_block_delta -data: {"type":"content_block_delta","delta":{"type":"text_delta","text":"pong"}} - -event: message_delta -data: {"type":"message_delta","delta":{"stop_reason":"end_turn"},"usage":{"output_tokens":5}} - -event: message_stop -data: {"type":"message_stop"} - -`)) - })) - defer server.Close() - - chat, err := NewAnthropicChat(&ChatConfig{ - Source: types.ModelSourceRemote, - BaseURL: server.URL, - ModelName: "gpt-5.5", - APIKey: "test-key", - Provider: string(provider.ProviderAnthropic), - }) - require.NoError(t, err) - - ch, err := chat.ChatStream(context.Background(), []Message{{Role: "user", Content: "ping"}}, nil) - require.NoError(t, err) - - var chunks []types.StreamResponse - for chunk := range ch { - chunks = append(chunks, chunk) - } - - require.Len(t, chunks, 2) - assert.True(t, capturedRequest.Stream) - assert.Equal(t, "pong", chunks[0].Content) - assert.False(t, chunks[0].Done) - assert.True(t, chunks[1].Done) - assert.Equal(t, "end_turn", chunks[1].FinishReason) - require.NotNil(t, chunks[1].Usage) - assert.Equal(t, 114, chunks[1].Usage.PromptTokens) - assert.Equal(t, 5, chunks[1].Usage.CompletionTokens) - assert.Equal(t, 119, chunks[1].Usage.TotalTokens) - assert.Equal(t, 100, chunks[1].Usage.CacheReadTokens) - assert.Equal(t, 14, chunks[1].Usage.CacheMissTokens) -} - -func TestNewRemoteChat_AnthropicProvider(t *testing.T) { - chat, err := NewRemoteChat(&ChatConfig{ - Source: types.ModelSourceRemote, - ModelName: "claude-sonnet-4-5", - APIKey: "test-key", - Provider: string(provider.ProviderAnthropic), - }) - require.NoError(t, err) - _, ok := chat.(*AnthropicChat) - assert.True(t, ok) + plugin, ok := client.(*PluginChat) + require.True(t, ok, "expected the plugin client, got %T", client) + assert.Equal(t, spi.ProtocolAnthropicMessages, plugin.Descriptor().Protocol) + assert.Equal(t, spi.ReplayAlways, plugin.Descriptor().EffectiveReplay(), + "thinking blocks must be replayed verbatim") } diff --git a/internal/models/chat/chat.go b/internal/models/chat/chat.go index aa2dc5efaf4..591eaf13cef 100644 --- a/internal/models/chat/chat.go +++ b/internal/models/chat/chat.go @@ -2,86 +2,42 @@ package chat import ( "context" - "encoding/json" "fmt" "strings" - "github.com/Tencent/WeKnora/internal/models/provider" + "github.com/Tencent/WeKnora/internal/models/llm/spi" "github.com/Tencent/WeKnora/internal/models/utils/ollama" "github.com/Tencent/WeKnora/internal/types" ) +// The request vocabulary lives in the model-plugin seam (internal/models/llm/spi) +// so protocol drivers can render it without importing this package. These +// aliases keep every existing caller working while leaving exactly one +// definition of a message in the codebase. + // Tool represents a function/tool definition -type Tool struct { - Type string `json:"type"` // "function" - Function FunctionDef `json:"function"` -} +type Tool = spi.Tool // FunctionDef represents a function definition -type FunctionDef struct { - Name string `json:"name"` - Description string `json:"description"` - Parameters json.RawMessage `json:"parameters"` -} +type FunctionDef = spi.FunctionDef // ChatOptions 聊天选项 -type ChatOptions struct { - Temperature float64 `json:"temperature"` // 温度参数 - TopP float64 `json:"top_p"` // Top P 参数 - Seed int `json:"seed"` // 随机种子 - MaxTokens int `json:"max_tokens"` // 最大 token 数 - MaxCompletionTokens int `json:"max_completion_tokens"` // 最大完成 token 数 - FrequencyPenalty float64 `json:"frequency_penalty"` // 频率惩罚 - PresencePenalty float64 `json:"presence_penalty"` // 存在惩罚 - Thinking *bool `json:"thinking"` // 是否启用思考 - Tools []Tool `json:"tools,omitempty"` // 可用工具列表 - ToolChoice string `json:"tool_choice,omitempty"` // "auto", "required", "none", or specific tool - ParallelToolCalls *bool `json:"parallel_tool_calls,omitempty"` // 是否允许并行工具调用(默认 nil 表示由模型决定) - Format json.RawMessage `json:"format,omitempty"` // 响应格式定义 -} +type ChatOptions = spi.Options // MessageContentPart represents a part of multi-content message -type MessageContentPart struct { - Type string `json:"type"` // "text" or "image_url" - Text string `json:"text,omitempty"` // For type="text" - ImageURL *ImageURL `json:"image_url,omitempty"` // For type="image_url" -} +type MessageContentPart = spi.MessageContentPart // ImageURL represents the image URL structure -type ImageURL struct { - URL string `json:"url"` // URL or base64 data URI - Detail string `json:"detail,omitempty"` // "auto", "low", "high" -} +type ImageURL = spi.ImageURL // Message 表示聊天消息 -type Message struct { - Role string `json:"role"` // 角色:system, user, assistant, tool - Content string `json:"content"` // 消息内容 - MultiContent []MessageContentPart `json:"multi_content,omitempty"` // 多内容消息(文本+图片) - Name string `json:"name,omitempty"` // Function/tool name (for tool role) - ToolCallID string `json:"tool_call_id,omitempty"` // Tool call ID (for tool role) - ToolCalls []ToolCall `json:"tool_calls,omitempty"` // Tool calls (for assistant role) - Images []string `json:"images,omitempty"` // Image URLs for multimodal (only for current user message) - // ReasoningContent 是 assistant 推理类模型(DeepSeek thinking、小米 MiMo、vLLM reasoning 等) - // 上一轮输出的思考内容。部分供应商(MiMo、DeepSeek V3.2/V4 thinking 模式)要求多轮对话中 - // 把 assistant 的 reasoning_content 原样回传,否则会以 400 拒绝请求;其他不要求的供应商 - // 会忽略未知字段,无副作用。 - ReasoningContent string `json:"reasoning_content,omitempty"` -} +type Message = spi.Message // ToolCall represents a tool call in a message -type ToolCall struct { - ID string `json:"id"` - Type string `json:"type"` // "function" - Function FunctionCall `json:"function"` - ProviderMetadata types.ToolCallMetadata `json:"provider_metadata,omitempty"` -} +type ToolCall = spi.ToolCall // FunctionCall represents a function call -type FunctionCall struct { - Name string `json:"name"` - Arguments string `json:"arguments"` // JSON string -} +type FunctionCall = spi.FunctionCall // Chat 定义了聊天接口 type Chat interface { @@ -157,16 +113,23 @@ func NewChat(config *ChatConfig, ollamaService *ollama.OllamaService) (Chat, err return wrapChatConcurrency(c, config.MaxConcurrency, err) } -// NewRemoteChat 根据 provider 创建远程聊天实例。 -// Anthropic 走独立的 Messages 协议实现;其余 OpenAI 兼容供应商统一由 -// RemoteAPIChat 处理,provider 特定行为在构造时通过 providerAdapter 解析。 +// NewRemoteChat 根据解析出的模型插件创建远程聊天实例。 +// +// Protocol selection is the plugin's, not a provider name's: a descriptor +// declares the wire it speaks, so Anthropic Messages and OpenAI Responses go +// through the plugin client that implements them, while the OpenAI-compatible +// majority keeps the established RemoteAPIChat transport. Either way the +// vendor's parameter dispositions come from the same descriptor. func NewRemoteChat(config *ChatConfig) (Chat, error) { - providerName := provider.ProviderName(config.Provider) - if providerName == "" { - providerName = provider.DetectProvider(config.BaseURL) - } - if providerName == provider.ProviderAnthropic { - return NewAnthropicChat(config) + desc, ok := resolveDescriptor(config) + if ok && desc.Protocol != spi.ProtocolOpenAIChat { + client, resolved, err := NewPluginChat(config) + if err != nil { + return nil, err + } + if resolved { + return client, nil + } } return NewRemoteAPIChat(config) } diff --git a/internal/models/chat/plugin_chat.go b/internal/models/chat/plugin_chat.go new file mode 100644 index 00000000000..51232f63d35 --- /dev/null +++ b/internal/models/chat/plugin_chat.go @@ -0,0 +1,314 @@ +package chat + +import ( + "bytes" + "context" + "encoding/json" + "fmt" + "io" + "net/http" + "strings" + + "github.com/Tencent/WeKnora/internal/logger" + "github.com/Tencent/WeKnora/internal/models/llm/protocol" + _ "github.com/Tencent/WeKnora/internal/models/llm/protocol/all" // register the standard protocol drivers + "github.com/Tencent/WeKnora/internal/models/llm/spi" + _ "github.com/Tencent/WeKnora/internal/models/llm/vendors" // register the built-in vendor plugins + "github.com/Tencent/WeKnora/internal/models/provider" + "github.com/Tencent/WeKnora/internal/types" + secutils "github.com/Tencent/WeKnora/internal/utils" +) + +// PluginChat drives a model through the plugin seam: a vendor descriptor +// supplies the parameter dispositions, a protocol driver supplies the wire. +// +// It owns only what neither of those does — the HTTP round trip, credentials, +// and timeouts — which is why it stays short while covering three protocols +// and every registered vendor. +type PluginChat struct { + descriptor spi.Descriptor + driver protocol.Driver + endpoint string + + modelName string + modelID string + creds spi.Credentials + + customHeaders map[string]string +} + +// endpointer is implemented by drivers that derive a full URL from a base URL, +// which is how a protocol absorbs the base-URL spellings users actually paste. +type endpointer interface { + Endpoint(baseURL string) string +} + +// NewPluginChat builds a plugin-driven chat client, reporting whether the +// configuration resolves to a registered plugin. A caller that gets false +// should fall back to the legacy path rather than fail: an unregistered vendor +// is a gap in the catalog, not a broken model. +func NewPluginChat(config *ChatConfig) (*PluginChat, bool, error) { + vendor := resolveVendorName(config) + desc, ok := spi.Resolve(spi.Query{ + Vendor: string(vendor), + Kind: spi.KindChat, + Model: config.ModelName, + Protocol: configuredProtocol(config), + }) + if !ok { + return nil, false, nil + } + + driver, err := protocol.MustGet(desc.Protocol) + if err != nil { + return nil, false, err + } + + baseURL := strings.TrimRight(config.BaseURL, "/") + if baseURL == "" { + baseURL = desc.DefaultBaseURL + } + if baseURL == "" { + return nil, false, fmt.Errorf("%s: base URL is required", desc.Vendor) + } + if err := secutils.ValidateURLForSSRF(baseURL); err != nil { + return nil, false, fmt.Errorf("baseURL SSRF check failed: %w", err) + } + + endpoint := buildEndpoint(driver, desc, baseURL) + if err := secutils.ValidateURLForSSRF(endpoint); err != nil { + return nil, false, fmt.Errorf("endpoint SSRF check failed: %w", err) + } + + if desc.Auth.EffectiveKind() == spi.AuthSigned { + if config.AppID == "" || config.AppSecret == "" { + return nil, false, fmt.Errorf("%s: application credentials are required", desc.Vendor) + } + } + + return &PluginChat{ + descriptor: desc, + driver: driver, + endpoint: endpoint, + modelName: remoteModelName(config), + modelID: config.ModelID, + creds: spi.Credentials{ + APIKey: config.APIKey, + AppID: config.AppID, + AppSecret: config.AppSecret, + }, + customHeaders: config.CustomHeaders, + }, true, nil +} + +// buildEndpoint resolves the request URL, letting a descriptor override the +// protocol's standard path for vendors serving a compatible protocol elsewhere. +func buildEndpoint(driver protocol.Driver, desc spi.Descriptor, baseURL string) string { + if desc.EndpointPath != "" { + return baseURL + desc.EndpointPath + } + if e, ok := driver.(endpointer); ok { + return e.Endpoint(baseURL) + } + return baseURL + driver.EndpointPath() +} + +// resolveVendorName reports the provider identity, falling back to detection +// from the base URL exactly as the rest of the model layer does. +func resolveVendorName(config *ChatConfig) provider.ProviderName { + name := provider.ProviderName(strings.TrimSpace(config.Provider)) + if name == "" { + name = provider.DetectProvider(config.BaseURL) + } + return name +} + +// configuredProtocol reports the protocol pinned in the model configuration. +// It is empty for the great majority of models, where the vendor offers one. +func configuredProtocol(config *ChatConfig) spi.ProtocolID { + if config.ExtraConfig == nil { + return "" + } + return spi.ProtocolID(strings.TrimSpace(config.ExtraConfig[ExtraConfigProtocol])) +} + +// remoteModelName reports the identifier to send, honoring the override some +// deployments need when the stored name is a local label. +func remoteModelName(config *ChatConfig) string { + if config.ExtraConfig != nil { + if override := strings.TrimSpace(config.ExtraConfig["remote_model_name"]); override != "" { + return override + } + } + return config.ModelName +} + +// GetModelName reports the model identifier. +func (c *PluginChat) GetModelName() string { return c.modelName } + +// GetModelID reports the stored model id. +func (c *PluginChat) GetModelID() string { return c.modelID } + +// Descriptor reports the resolved plugin, for diagnostics. +func (c *PluginChat) Descriptor() spi.Descriptor { return c.descriptor } + +// Plan resolves what a call would send without sending it. The debug endpoint +// uses it to report the real request shape rather than re-deriving it. +func (c *PluginChat) Plan(opts *ChatOptions, stream bool) (*spi.Plan, error) { + return c.descriptor.Plan(spi.Request{ + Model: c.modelName, + Stream: stream, + Values: opts.ParamValues(), + }) +} + +// build renders one call into its outbound bytes, running the protocol driver +// first and the vendor plan over the result. +func (c *PluginChat) build(messages []Message, opts *ChatOptions, stream bool) ([]byte, *spi.Plan, error) { + draft, err := c.driver.BuildDraft(protocol.Call{ + Model: c.modelName, + Stream: stream, + Messages: messages, + Options: opts, + Replay: c.descriptor.EffectiveReplay(), + }) + if err != nil { + return nil, nil, fmt.Errorf("build request: %w", err) + } + + plan, err := c.Plan(opts, stream) + if err != nil { + return nil, nil, err + } + if err := plan.Apply(draft); err != nil { + return nil, nil, err + } + + body, err := json.Marshal(draft.Body) + if err != nil { + return nil, nil, fmt.Errorf("marshal request: %w", err) + } + return body, plan, nil +} + +// newRequest builds the authenticated HTTP request for a body. +func (c *PluginChat) newRequest(ctx context.Context, body []byte, stream bool) (*http.Request, error) { + req, err := http.NewRequestWithContext(ctx, http.MethodPost, c.endpoint, bytes.NewReader(body)) + if err != nil { + return nil, fmt.Errorf("create request: %w", err) + } + req.Header.Set("Content-Type", "application/json") + if stream { + req.Header.Set("Accept", "text/event-stream") + } + + auth := c.descriptor.Auth + for key, value := range auth.Static { + req.Header.Set(key, value) + } + switch auth.EffectiveKind() { + case spi.AuthHeader: + req.Header.Set(auth.Header, c.creds.APIKey) + case spi.AuthSigned: + headers, err := auth.Signer.Sign(body, c.creds) + if err != nil { + return nil, fmt.Errorf("sign request: %w", err) + } + for key, value := range headers { + req.Header.Set(key, value) + } + default: + req.Header.Set("Authorization", "Bearer "+c.creds.APIKey) + } + + // User headers come last but never displace the ones above; the helper + // skips reserved names. + secutils.ApplyCustomHeaders(req, c.customHeaders) + return req, nil +} + +// logPlan records the adjustments a vendor's rules forced, so a request that +// came back without reasoning can be explained from the logs alone. +func (c *PluginChat) logPlan(ctx context.Context, plan *spi.Plan, body []byte) { + logger.Infof(ctx, "[LLM Request] vendor=%s protocol=%s model=%s endpoint=%s request:\n%s", + plan.Vendor, plan.Protocol, c.modelName, c.endpoint, + secutils.CompactImageDataURLForLog(string(body))) + for _, note := range plan.Notes { + logger.Infof(ctx, "[LLM Plan] %s %s: %s", note.Param, note.Reason, note.Detail) + } +} + +// Chat performs a non-streaming call. +func (c *PluginChat) Chat(ctx context.Context, messages []Message, opts *ChatOptions) (*types.ChatResponse, error) { + ctx, cancel := withLLMTimeout(ctx, defaultChatTimeout) + defer cancel() + + body, plan, err := c.build(messages, opts, false) + if err != nil { + return nil, err + } + c.logPlan(ctx, plan, body) + + req, err := c.newRequest(ctx, body, false) + if err != nil { + return nil, err + } + resp, err := rawHTTPClient.Do(req) + if err != nil { + return nil, fmt.Errorf("send request: %w", err) + } + defer resp.Body.Close() + + payload, err := io.ReadAll(resp.Body) + if err != nil { + return nil, fmt.Errorf("read response: %w", err) + } + if resp.StatusCode < http.StatusOK || resp.StatusCode >= http.StatusMultipleChoices { + return nil, fmt.Errorf("API request failed with status %d: %s", resp.StatusCode, string(payload)) + } + + result, err := c.driver.ParseResponse(payload) + if err != nil { + return nil, err + } + logUsage(ctx, c.modelName, &result.Usage) + return result, nil +} + +// ChatStream performs a streaming call. +func (c *PluginChat) ChatStream(ctx context.Context, messages []Message, opts *ChatOptions) (<-chan types.StreamResponse, error) { + streamCtx, cancel := withLLMTimeout(ctx, defaultStreamTimeout) + + body, plan, err := c.build(messages, opts, true) + if err != nil { + cancel() + return nil, err + } + c.logPlan(streamCtx, plan, body) + + req, err := c.newRequest(streamCtx, body, true) + if err != nil { + cancel() + return nil, err + } + resp, err := rawHTTPClient.Do(req) + if err != nil { + cancel() + return nil, fmt.Errorf("send request: %w", err) + } + if resp.StatusCode < http.StatusOK || resp.StatusCode >= http.StatusMultipleChoices { + payload, _ := io.ReadAll(resp.Body) + resp.Body.Close() + cancel() + return nil, fmt.Errorf("API request failed with status %d: %s", resp.StatusCode, string(payload)) + } + + out := make(chan types.StreamResponse) + go func() { + defer cancel() + defer close(out) + defer resp.Body.Close() + c.driver.DecodeStream(streamCtx, resp.Body, out) + }() + return out, nil +} diff --git a/internal/models/chat/plugin_chat_test.go b/internal/models/chat/plugin_chat_test.go new file mode 100644 index 00000000000..510da4044c1 --- /dev/null +++ b/internal/models/chat/plugin_chat_test.go @@ -0,0 +1,291 @@ +package chat + +import ( + "context" + "encoding/json" + "io" + "net/http" + "net/http/httptest" + "strings" + "testing" + + "github.com/Tencent/WeKnora/internal/models/llm/spi" + "github.com/Tencent/WeKnora/internal/types" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +// These tests drive the plugin client end to end against a stub server, so +// they cover what a unit test of either half would miss: that the descriptor's +// auth and parameter dispositions, the protocol driver's body, and the HTTP +// layer agree on one request. + +// capture records the request a client sends and replies with a canned body. +type capture struct { + body map[string]any + headers http.Header + path string +} + +func stubServer(t *testing.T, reply string, contentType string) (*httptest.Server, *capture) { + t.Helper() + // The client validates every outbound URL against the SSRF guard, which + // blocks loopback by default; the stub server lives there. + t.Setenv("SSRF_WHITELIST", "127.0.0.1") + got := &capture{} + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + raw, err := io.ReadAll(r.Body) + require.NoError(t, err) + if len(raw) > 0 { + require.NoError(t, json.Unmarshal(raw, &got.body)) + } + got.headers = r.Header.Clone() + got.path = r.URL.Path + w.Header().Set("Content-Type", contentType) + _, _ = w.Write([]byte(reply)) + })) + t.Cleanup(server.Close) + return server, got +} + +func newPluginChat(t *testing.T, config *ChatConfig) *PluginChat { + t.Helper() + client, ok, err := NewPluginChat(config) + require.NoError(t, err) + require.True(t, ok, "expected a registered plugin for provider %q", config.Provider) + return client +} + +const anthropicReply = `{ + "content": [ + {"type": "thinking", "thinking": "27 * 453"}, + {"type": "text", "text": "12231"} + ], + "stop_reason": "end_turn", + "usage": {"input_tokens": 10, "output_tokens": 4, + "cache_read_input_tokens": 6, "cache_creation_input_tokens": 2} +}` + +func TestPluginChat_AnthropicRequestShape(t *testing.T) { + server, got := stubServer(t, anthropicReply, "application/json") + + client := newPluginChat(t, &ChatConfig{ + Source: types.ModelSourceRemote, + BaseURL: server.URL, + ModelName: "claude-sonnet-4-5", + APIKey: "test-key", + Provider: "anthropic", + }) + + resp, err := client.Chat(context.Background(), []Message{ + {Role: "system", Content: "be brief"}, + {Role: "user", Content: "27 * 453?"}, + }, &ChatOptions{Thinking: ptrBool(true), ThinkingBudget: 2000}) + require.NoError(t, err) + + // The system prompt is a top-level field here, not a message. + assert.Equal(t, "be brief", got.body["system"]) + assert.Equal(t, "/v1/messages", got.path) + assert.Equal(t, "test-key", got.headers.Get("x-api-key")) + assert.Equal(t, "2023-06-01", got.headers.Get("anthropic-version")) + assert.Empty(t, got.headers.Get("Authorization"), "the key must not also travel as a bearer token") + + thinking, ok := got.body["thinking"].(map[string]any) + require.True(t, ok, "body: %v", got.body) + assert.Equal(t, "enabled", thinking["type"]) + assert.EqualValues(t, 2000, thinking["budget_tokens"]) + // The budget must fit under the ceiling, which the plugin raises for it. + assert.Greater(t, got.body["max_tokens"], float64(2000)) + + messages, ok := got.body["messages"].([]any) + require.True(t, ok) + require.Len(t, messages, 1) + first := messages[0].(map[string]any) + assert.Equal(t, "user", first["role"]) + blocks := first["content"].([]any) + assert.Equal(t, "text", blocks[0].(map[string]any)["type"]) + + assert.Equal(t, "12231", resp.Content) + assert.Equal(t, "27 * 453", resp.ReasoningContent) + assert.Equal(t, "end_turn", resp.FinishReason) + // Cache counters are additional input tokens in this protocol, not a + // subset of them. + assert.Equal(t, 18, resp.Usage.PromptTokens) + assert.Equal(t, 6, resp.Usage.CacheReadTokens) + assert.True(t, resp.Usage.CacheReported) +} + +// A tool result is a user-role block referencing its call, and a replayed +// thinking block must lead the assistant turn with its signature intact. +func TestPluginChat_AnthropicReplaysThinkingAndToolResults(t *testing.T) { + server, got := stubServer(t, anthropicReply, "application/json") + + client := newPluginChat(t, &ChatConfig{ + Source: types.ModelSourceRemote, + BaseURL: server.URL, + ModelName: "claude-sonnet-4-5", + APIKey: "k", + Provider: "anthropic", + }) + + _, err := client.Chat(context.Background(), []Message{ + {Role: "user", Content: "weather?"}, + { + Role: "assistant", + ReasoningContent: "need the tool", + ReasoningSignature: "sig-abc", + ToolCalls: []ToolCall{{ + ID: "call_1", + Type: "function", + Function: FunctionCall{Name: "get_weather", Arguments: `{"city":"SF"}`}, + }}, + }, + {Role: "tool", ToolCallID: "call_1", Content: "sunny"}, + }, &ChatOptions{}) + require.NoError(t, err) + + messages := got.body["messages"].([]any) + require.Len(t, messages, 3) + + assistant := messages[1].(map[string]any) + assert.Equal(t, "assistant", assistant["role"]) + blocks := assistant["content"].([]any) + lead := blocks[0].(map[string]any) + assert.Equal(t, "thinking", lead["type"], "thinking must lead the turn") + assert.Equal(t, "sig-abc", lead["signature"], "the signature must return unmodified") + assert.Equal(t, "tool_use", blocks[1].(map[string]any)["type"]) + + result := messages[2].(map[string]any) + assert.Equal(t, "user", result["role"]) + block := result["content"].([]any)[0].(map[string]any) + assert.Equal(t, "tool_result", block["type"]) + assert.Equal(t, "call_1", block["tool_use_id"]) +} + +// Users paste the base URL in several shapes; all of them must reach the same +// endpoint rather than a doubled or truncated path. +func TestPluginChat_AnthropicEndpointSpellings(t *testing.T) { + for _, suffix := range []string{"", "/", "/v1", "/v1/messages"} { + t.Run("base"+suffix, func(t *testing.T) { + server, got := stubServer(t, anthropicReply, "application/json") + client := newPluginChat(t, &ChatConfig{ + Source: types.ModelSourceRemote, + BaseURL: server.URL + suffix, + ModelName: "claude-sonnet-4-5", + APIKey: "k", + Provider: "anthropic", + }) + _, err := client.Chat(context.Background(), []Message{{Role: "user", Content: "hi"}}, &ChatOptions{}) + require.NoError(t, err) + assert.Equal(t, "/v1/messages", got.path) + }) + } +} + +func TestPluginChat_AnthropicStream(t *testing.T) { + stream := strings.Join([]string{ + `event: message_start`, + `data: {"type":"message_start","message":{"usage":{"input_tokens":10,"output_tokens":0}}}`, + ``, + `event: content_block_delta`, + `data: {"type":"content_block_delta","delta":{"type":"thinking_delta","thinking":"hmm"}}`, + ``, + `event: content_block_delta`, + `data: {"type":"content_block_delta","delta":{"type":"text_delta","text":"hello"}}`, + ``, + `event: message_delta`, + `data: {"type":"message_delta","delta":{"stop_reason":"end_turn"},"usage":{"input_tokens":10,"output_tokens":3}}`, + ``, + }, "\n") + server, got := stubServer(t, stream, "text/event-stream") + + client := newPluginChat(t, &ChatConfig{ + Source: types.ModelSourceRemote, + BaseURL: server.URL, + ModelName: "claude-sonnet-4-5", + APIKey: "k", + Provider: "anthropic", + }) + + ch, err := client.ChatStream(context.Background(), []Message{{Role: "user", Content: "hi"}}, &ChatOptions{}) + require.NoError(t, err) + + var thinking, answer strings.Builder + var final types.StreamResponse + for chunk := range ch { + switch chunk.ResponseType { + case types.ResponseTypeThinking: + thinking.WriteString(chunk.Content) + case types.ResponseTypeAnswer: + answer.WriteString(chunk.Content) + if chunk.Done { + final = chunk + } + case types.ResponseTypeError: + t.Fatalf("stream error: %s", chunk.Content) + } + } + + assert.Equal(t, true, got.body["stream"]) + assert.Equal(t, "hmm", thinking.String()) + assert.Equal(t, "hello", answer.String()) + assert.Equal(t, "end_turn", final.FinishReason) + require.NotNil(t, final.Usage) + assert.Equal(t, 3, final.Usage.CompletionTokens) +} + +// The Responses protocol is reachable by pinning it on a vendor that offers +// more than one, which is the case the protocol selector exists for. +func TestPluginChat_ResponsesProtocol(t *testing.T) { + reply := `{ + "status": "completed", + "output": [ + {"type": "reasoning", "summary": [{"type": "summary_text", "text": "brief"}]}, + {"type": "message", "content": [{"type": "output_text", "text": "Paris"}]} + ], + "usage": {"input_tokens": 9, "output_tokens": 2, "total_tokens": 11} + }` + server, got := stubServer(t, reply, "application/json") + + client := newPluginChat(t, &ChatConfig{ + Source: types.ModelSourceRemote, + BaseURL: server.URL, + ModelName: "gpt-5", + APIKey: "k", + Provider: "openai", + ExtraConfig: map[string]string{ExtraConfigProtocol: string(spi.ProtocolOpenAIResponses)}, + }) + require.Equal(t, spi.ProtocolOpenAIResponses, client.Descriptor().Protocol) + + resp, err := client.Chat(context.Background(), []Message{ + {Role: "system", Content: "be brief"}, + {Role: "user", Content: "capital of France?"}, + }, &ChatOptions{Thinking: ptrBool(true), ThinkingEffort: "low", MaxTokens: 300}) + require.NoError(t, err) + + assert.Equal(t, "/responses", got.path) + assert.Equal(t, "be brief", got.body["instructions"], "the system prompt becomes instructions") + assert.EqualValues(t, 300, got.body["max_output_tokens"]) + assert.NotContains(t, got.body, "max_tokens") + reasoning := got.body["reasoning"].(map[string]any) + assert.Equal(t, "low", reasoning["effort"]) + + assert.Equal(t, "Paris", resp.Content) + assert.Equal(t, "brief", resp.ReasoningContent) + assert.Equal(t, "stop", resp.FinishReason) + assert.Equal(t, 11, resp.Usage.TotalTokens) +} + +// An unregistered provider must report "not resolved" rather than erroring, so +// the caller can fall back to the generic transport. +func TestNewPluginChat_UnregisteredProvider(t *testing.T) { + client, ok, err := NewPluginChat(&ChatConfig{ + Source: types.ModelSourceRemote, + BaseURL: "https://example.com", + ModelName: "x", + Provider: "not-a-registered-vendor", + }) + require.NoError(t, err) + assert.False(t, ok) + assert.Nil(t, client) +} diff --git a/internal/models/chat/provider.go b/internal/models/chat/provider.go index c596a354739..aa9b0a5c2d1 100644 --- a/internal/models/chat/provider.go +++ b/internal/models/chat/provider.go @@ -12,6 +12,17 @@ import ( "github.com/sashabaranov/go-openai" ) +// This file holds the transport-level provider behavior that a declarative +// descriptor cannot express: request signing, an endpoint that is not derived +// from the base URL, a message rewrite, and vendor state that must survive a +// tool-call round trip. +// +// Everything else that used to live here — which field carries the thinking +// toggle, which sampling parameters a reasoning model rejects, which model +// pins its temperature — is now declared in internal/models/llm/vendors and +// applied through the resolved plan. Those were the parts that had to be kept +// in sync with a frontend copy, and they no longer exist in two places. + // authCreds carries the credentials a providerAdapter needs to authenticate a // raw HTTP request. APIKey covers the common Bearer / api-key cases; AppID and // AppSecret are only used by signing providers (WeKnoraCloud). @@ -21,22 +32,15 @@ type authCreds struct { AppSecret string } -// providerAdapter captures everything provider-specific about an -// OpenAI-compatible chat backend. Every method has a sensible default on -// baseProvider, so a new provider is added by embedding baseProvider and -// overriding only the one or two methods that actually differ. +// providerAdapter captures the transport-level behavior of an +// OpenAI-compatible chat backend. Every method has a default on baseProvider, +// so an adapter overrides only what genuinely differs. type providerAdapter interface { // Name is the provider this adapter handles. Name() provider.ProviderName // Matches reports whether this adapter applies to the given model name. - // Used for sub-provider routing (e.g. Qwen thinking models within Aliyun, - // reasoning models within OpenAI). Default: true. + // Default: true. Matches(model string) bool - // Thinking is how this provider encodes ChatOptions.Thinking. Default: none. - Thinking() ThinkingStrategy - // ShapeRequest applies in-place parameter quirks to the standard request - // (stripping unsupported fields, pinning temperature, …). Default: noop. - ShapeRequest(req *openai.ChatCompletionRequest, opts *ChatOptions, isStream bool) // TransformMessages rewrites the converted messages (e.g. downgrading // multi-content to plain text). Default: identity. TransformMessages(msgs []openai.ChatCompletionMessage) []openai.ChatCompletionMessage @@ -46,7 +50,8 @@ type providerAdapter interface { // Auth sets authentication headers on a raw HTTP request. Default: Bearer. Auth(req *http.Request, creds authCreds, body []byte) // ForceRawHTTP forces the raw HTTP path even when the body is standard - // (needed by providers that must sign the exact request bytes). Default: false. + // (needed by providers that must sign the exact request bytes, or whose + // response carries counters the SDK type cannot represent). Default: false. ForceRawHTTP() bool // ExtractToolCallMetadata captures provider-specific state from a raw // OpenAI-compatible tool_call object. Default: nil. @@ -58,13 +63,11 @@ type providerAdapter interface { // baseProvider supplies the default behavior for every providerAdapter method. // It is also the fallback returned by resolveProvider for unknown providers: -// Bearer auth, standard endpoint, no thinking, no request shaping. +// Bearer auth, standard endpoint, no request shaping. type baseProvider struct{} -func (baseProvider) Name() provider.ProviderName { return "" } -func (baseProvider) Matches(string) bool { return true } -func (baseProvider) Thinking() ThinkingStrategy { return noThinking{} } -func (baseProvider) ShapeRequest(*openai.ChatCompletionRequest, *ChatOptions, bool) {} +func (baseProvider) Name() provider.ProviderName { return "" } +func (baseProvider) Matches(string) bool { return true } func (baseProvider) TransformMessages(msgs []openai.ChatCompletionMessage) []openai.ChatCompletionMessage { return msgs } @@ -119,29 +122,7 @@ func (weKnoraCloudProvider) TransformMessages(messages []openai.ChatCompletionMe return result } -// --- Aliyun Qwen thinking models: enable_thinking (always sent, forced off non-stream) --- - -type qwenThinkingProvider struct{ baseProvider } - -func (qwenThinkingProvider) Name() provider.ProviderName { return provider.ProviderAliyun } -func (qwenThinkingProvider) Matches(model string) bool { return provider.IsQwenThinkingModel(model) } -func (qwenThinkingProvider) Thinking() ThinkingStrategy { - return enableThinking{alwaysSend: true, disableOnNonStream: true} -} - -// --- LKEAP: thinking via { "thinking": { "type": ... } }, only for DeepSeek V3.x --- -// R1 series enables chain-of-thought by default and is left untouched (falls -// back to baseProvider). See https://cloud.tencent.com/document/product/1772/115963 - -type lkeapProvider struct{ baseProvider } - -func (lkeapProvider) Name() provider.ProviderName { return provider.ProviderLKEAP } -func (lkeapProvider) Matches(model string) bool { - return strings.Contains(strings.ToLower(model), "deepseek-v3") -} -func (lkeapProvider) Thinking() ThinkingStrategy { return thinkingTypeField{} } - -// --- DeepSeek: does not support tool_choice --- +// --- DeepSeek: native cache counters the SDK response type cannot represent --- type deepseekProvider struct{ baseProvider } @@ -150,23 +131,6 @@ func (deepseekProvider) Name() provider.ProviderName { return provider.ProviderD // Native DeepSeek cache counters are not represented by go-openai v1.41.2; // use the raw path so prompt_cache_hit_tokens/miss_tokens remain observable. func (deepseekProvider) ForceRawHTTP() bool { return true } -func (deepseekProvider) ShapeRequest(req *openai.ChatCompletionRequest, opts *ChatOptions, _ bool) { - if opts != nil && opts.ToolChoice != "" { - req.ToolChoice = nil - } -} - -// --- Generic (vLLM) / NVIDIA: thinking via chat_template_kwargs --- - -type genericProvider struct{ baseProvider } - -func (genericProvider) Name() provider.ProviderName { return provider.ProviderGeneric } -func (genericProvider) Thinking() ThinkingStrategy { return chatTemplateKwargs{} } - -type nvidiaProvider struct{ baseProvider } - -func (nvidiaProvider) Name() provider.ProviderName { return provider.ProviderNvidia } -func (nvidiaProvider) Thinking() ThinkingStrategy { return chatTemplateKwargs{} } // --- Gemini OpenAI compatibility: tool thought signatures live in extra_content --- @@ -187,6 +151,7 @@ func (geminiProvider) ExtractToolCallMetadata(raw json.RawMessage) types.ToolCal } return types.ToolCallMetadata{"google": google} } + func (geminiProvider) InjectToolCallMetadata(toolCall map[string]any, metadata types.ToolCallMetadata) { if len(metadata) == 0 { return @@ -202,14 +167,7 @@ func (geminiProvider) InjectToolCallMetadata(toolCall map[string]any, metadata t toolCall["extra_content"] = map[string]any{"google": googleValue} } -// --- Volcengine (火山引擎 Ark): thinking via { "thinking": { "type": ... } } --- - -type volcengineProvider struct{ baseProvider } - -func (volcengineProvider) Name() provider.ProviderName { return provider.ProviderVolcengine } -func (volcengineProvider) Thinking() ThinkingStrategy { return thinkingTypeField{} } - -// --- Azure OpenAI: api-key auth (reasoning variant also strips sampling params) --- +// --- Azure OpenAI: api-key auth --- type azureProvider struct{ baseProvider } @@ -218,76 +176,17 @@ func (azureProvider) Auth(req *http.Request, creds authCreds, _ []byte) { req.Header.Set("api-key", creds.APIKey) } -type azureReasoningProvider struct{ azureProvider } - -func (azureReasoningProvider) Matches(model string) bool { - return provider.IsOpenAIReasoningOrGPT5Model(model) -} -func (azureReasoningProvider) ShapeRequest(req *openai.ChatCompletionRequest, _ *ChatOptions, _ bool) { - shapeOpenAIReasoning(req) -} - -// --- OpenAI reasoning / GPT-5: no sampling params, must use max_completion_tokens --- - -type openAIReasoningProvider struct{ baseProvider } - -func (openAIReasoningProvider) Name() provider.ProviderName { return provider.ProviderOpenAI } -func (openAIReasoningProvider) Matches(model string) bool { - return provider.IsOpenAIReasoningOrGPT5Model(model) -} -func (openAIReasoningProvider) ShapeRequest(req *openai.ChatCompletionRequest, _ *ChatOptions, _ bool) { - shapeOpenAIReasoning(req) -} - -// --- Moonshot: v1 models accept only temperature=1 --- - -type moonshotProvider struct{ baseProvider } - -func (moonshotProvider) Name() provider.ProviderName { return provider.ProviderMoonshot } -func (moonshotProvider) Matches(model string) bool { - return provider.IsMoonshotFixedTempModel(model) -} -func (moonshotProvider) ShapeRequest(req *openai.ChatCompletionRequest, _ *ChatOptions, _ bool) { - // Pin temperature to 1 and drop the other sampling params, matching the - // pre-refactor behavior where these fields were never set for this model. - req.Temperature = 1 - req.TopP = 0 - req.FrequencyPenalty = 0 - req.PresencePenalty = 0 -} - -// shapeOpenAIReasoning strips sampling params (unsupported by o-series / GPT-5) -// and migrates max_tokens to max_completion_tokens. See issue #1283. -func shapeOpenAIReasoning(req *openai.ChatCompletionRequest) { - req.Temperature = 0 - req.TopP = 0 - req.FrequencyPenalty = 0 - req.PresencePenalty = 0 - if req.MaxCompletionTokens == 0 && req.MaxTokens > 0 { - req.MaxCompletionTokens = req.MaxTokens - } - req.MaxTokens = 0 -} - // providerRegistry is ordered: more specific adapters (those with a real // Matches predicate) must precede the generic catch-all for the same provider. var providerRegistry = []providerAdapter{ weKnoraCloudProvider{}, - qwenThinkingProvider{}, - lkeapProvider{}, deepseekProvider{}, - genericProvider{}, geminiProvider{}, - volcengineProvider{}, - nvidiaProvider{}, - azureReasoningProvider{}, azureProvider{}, - openAIReasoningProvider{}, - moonshotProvider{}, } // resolveProvider returns the adapter handling the given provider+model, or -// baseProvider{} (Bearer auth, standard endpoint, no thinking) when none matches. +// baseProvider{} (Bearer auth, standard endpoint, no shaping) when none matches. func resolveProvider(name provider.ProviderName, model string) providerAdapter { for _, p := range providerRegistry { if p.Name() == name && p.Matches(model) { diff --git a/internal/models/chat/provider_test.go b/internal/models/chat/provider_test.go index bc1049fbc8d..73ae1e0cc42 100644 --- a/internal/models/chat/provider_test.go +++ b/internal/models/chat/provider_test.go @@ -11,9 +11,10 @@ import ( "github.com/stretchr/testify/require" ) -// TestResolveProvider pins the provider+model routing table, including the -// sub-model matchers (reasoning models, Qwen thinking, LKEAP DeepSeek V3) and -// the baseProvider fallback for everything else. +// TestResolveProvider pins the transport-level adapter table. It is short +// because parameter behavior moved to the model plugins: what remains here is +// only what a declaration cannot express — signing, a non-derived endpoint, a +// message rewrite, and tool-call metadata. func TestResolveProvider(t *testing.T) { cases := []struct { name string @@ -21,21 +22,12 @@ func TestResolveProvider(t *testing.T) { model string want providerAdapter }{ - {"deepseek", provider.ProviderDeepSeek, "deepseek-chat", deepseekProvider{}}, - {"lkeap v3", provider.ProviderLKEAP, "deepseek-v3.1", lkeapProvider{}}, - {"lkeap r1 falls back", provider.ProviderLKEAP, "deepseek-r1", baseProvider{}}, - {"qwen thinking", provider.ProviderAliyun, "qwen3-32b", qwenThinkingProvider{}}, - {"generic", provider.ProviderGeneric, "anything", genericProvider{}}, - {"gemini", provider.ProviderGemini, "gemini-3-flash-preview", geminiProvider{}}, - {"nvidia", provider.ProviderNvidia, "anything", nvidiaProvider{}}, - {"volcengine", provider.ProviderVolcengine, "doubao", volcengineProvider{}}, - {"openai non-reasoning falls back", provider.ProviderOpenAI, "gpt-4o", baseProvider{}}, - {"openai reasoning", provider.ProviderOpenAI, "gpt-5", openAIReasoningProvider{}}, - {"azure non-reasoning", provider.ProviderAzureOpenAI, "gpt-4", azureProvider{}}, - {"azure reasoning", provider.ProviderAzureOpenAI, "gpt-5-mini", azureReasoningProvider{}}, - {"moonshot fixed temp", provider.ProviderMoonshot, "moonshot-v1-8k", moonshotProvider{}}, - {"moonshot other falls back", provider.ProviderMoonshot, "kimi-latest", baseProvider{}}, - {"weknora cloud", provider.ProviderWeKnoraCloud, "anything", weKnoraCloudProvider{}}, + {"deepseek keeps the raw path for cache counters", provider.ProviderDeepSeek, "deepseek-chat", deepseekProvider{}}, + {"gemini carries thought signatures", provider.ProviderGemini, "gemini-3-flash-preview", geminiProvider{}}, + {"azure authenticates with api-key", provider.ProviderAzureOpenAI, "gpt-4", azureProvider{}}, + {"weknora cloud signs its requests", provider.ProviderWeKnoraCloud, "anything", weKnoraCloudProvider{}}, + {"openai needs no transport adapter", provider.ProviderOpenAI, "gpt-5", baseProvider{}}, + {"aliyun needs no transport adapter", provider.ProviderAliyun, "qwen3-32b", baseProvider{}}, {"unknown falls back", provider.ProviderName("nope"), "x", baseProvider{}}, } for _, tc := range cases { @@ -59,13 +51,22 @@ func newOutboundChat(t *testing.T, providerName, model string, extra map[string] return c } -// TestBuildOutbound_Thinking is the characterization suite for the merged -// thinking-control path: it asserts that buildOutbound produces the same wire -// formats the pre-refactor provider customizers did. +// TestBuildOutbound_Thinking asserts the wire format each vendor's plugin +// produces through the OpenAI-compatible transport. The expectations match the +// ones this path produced before the plugin seam existed, so a regression in +// the new resolution shows up as a changed body rather than as silence. func TestBuildOutbound_Thinking(t *testing.T) { msgs := []Message{{Role: "user", Content: "hi"}} - t.Run("generic explicit thinking_type overrides legacy kwargs", func(t *testing.T) { + t.Run("generic deployment uses chat_template_kwargs", func(t *testing.T) { + c := newOutboundChat(t, string(provider.ProviderGeneric), "qwen", nil) + body, _, useRaw, err := c.buildOutbound(msgs, &ChatOptions{Thinking: ptrBool(false)}, true) + require.NoError(t, err) + require.True(t, useRaw) + assert.Contains(t, mustJSON(t, body), "chat_template_kwargs") + }) + + t.Run("stored thinking_type override wins over the plugin default", func(t *testing.T) { c := newOutboundChat(t, string(provider.ProviderGeneric), "deepseek-v4-flash", map[string]string{ExtraConfigThinkingControl: "thinking_type"}) body, _, useRaw, err := c.buildOutbound(msgs, &ChatOptions{Thinking: ptrBool(false)}, true) @@ -77,15 +78,7 @@ func TestBuildOutbound_Thinking(t *testing.T) { assert.NotContains(t, js, "chat_template_kwargs") }) - t.Run("generic legacy chat_template_kwargs", func(t *testing.T) { - c := newOutboundChat(t, string(provider.ProviderGeneric), "qwen", nil) - body, _, useRaw, err := c.buildOutbound(msgs, &ChatOptions{Thinking: ptrBool(false)}, true) - require.NoError(t, err) - require.True(t, useRaw) - assert.Contains(t, mustJSON(t, body), "chat_template_kwargs") - }) - - t.Run("none keeps the standard SDK request", func(t *testing.T) { + t.Run("stored none override sends no thinking field", func(t *testing.T) { c := newOutboundChat(t, string(provider.ProviderGeneric), "x", map[string]string{ExtraConfigThinkingControl: "none"}) body, _, useRaw, err := c.buildOutbound(msgs, &ChatOptions{Thinking: ptrBool(false)}, true) @@ -128,7 +121,7 @@ func TestBuildOutbound_Thinking(t *testing.T) { assert.Contains(t, mustJSON(t, body), `"thinking"`) }) - t.Run("lkeap r1 left untouched", func(t *testing.T) { + t.Run("lkeap r1 reasons unconditionally and takes no toggle", func(t *testing.T) { c := newOutboundChat(t, string(provider.ProviderLKEAP), "deepseek-r1", nil) body, _, useRaw, err := c.buildOutbound(msgs, &ChatOptions{Thinking: ptrBool(false)}, true) require.NoError(t, err) @@ -138,9 +131,9 @@ func TestBuildOutbound_Thinking(t *testing.T) { }) } -// TestBuildOutbound_ShapeRequest covers the param-shaping providers that used -// to live inline in BuildChatCompletionRequest. -func TestBuildOutbound_ShapeRequest(t *testing.T) { +// TestBuildOutbound_ParameterDispositions covers the vendors whose plugins +// forbid or pin a standard parameter. +func TestBuildOutbound_ParameterDispositions(t *testing.T) { msgs := []Message{{Role: "user", Content: "hi"}} t.Run("deepseek strips tool_choice", func(t *testing.T) { @@ -157,9 +150,20 @@ func TestBuildOutbound_ShapeRequest(t *testing.T) { c := newOutboundChat(t, string(provider.ProviderMoonshot), "moonshot-v1-8k", nil) body, _, _, err := c.buildOutbound(msgs, &ChatOptions{Temperature: 0.7, TopP: 0.9}, false) require.NoError(t, err) - req := body.(*openai.ChatCompletionRequest) - assert.EqualValues(t, 1, req.Temperature) - assert.EqualValues(t, 0, req.TopP) + js := mustJSON(t, body) + assert.Contains(t, js, `"temperature":1`) + }) + + t.Run("openai reasoning model drops sampling and renames the ceiling", func(t *testing.T) { + c := newOutboundChat(t, string(provider.ProviderOpenAI), "o3-mini", nil) + body, _, useRaw, err := c.buildOutbound(msgs, &ChatOptions{Temperature: 0.7, MaxTokens: 4096}, false) + require.NoError(t, err) + require.True(t, useRaw) + request, ok := body.(map[string]any) + require.True(t, ok) + assert.NotContains(t, request, "temperature") + assert.NotContains(t, request, "max_tokens") + assert.EqualValues(t, 4096, request["max_completion_tokens"]) }) } diff --git a/internal/models/chat/remote_api.go b/internal/models/chat/remote_api.go index 02b93d0f02b..1ad76336533 100644 --- a/internal/models/chat/remote_api.go +++ b/internal/models/chat/remote_api.go @@ -10,6 +10,7 @@ import ( "strings" "github.com/Tencent/WeKnora/internal/logger" + "github.com/Tencent/WeKnora/internal/models/llm/spi" "github.com/Tencent/WeKnora/internal/models/provider" "github.com/Tencent/WeKnora/internal/types" secutils "github.com/Tencent/WeKnora/internal/utils" @@ -32,10 +33,16 @@ type RemoteAPIChat struct { // customHeaders 为用户在模型配置中指定的自定义 HTTP 请求头(类似 OpenAI Python SDK 的 extra_headers)。 customHeaders map[string]string - // adapter 承载所有 provider 特定行为(thinking / 参数特判 / endpoint / 鉴权 / 消息变换)。 + // adapter carries the transport-level provider behavior (endpoint, auth, + // message transform, tool-call metadata). adapter providerAdapter - // thinkingOverride 来自 extra_config.thinking_control,非 nil 时覆盖 adapter.Thinking()。 - thinkingOverride ThinkingStrategy + // descriptor is the resolved model plugin. It supplies every parameter + // disposition — the thinking field, forbidden sampling knobs, pinned + // values, renamed ceilings — so this type no longer knows any of them. + // It is absent only for a provider with no registered plugin, in which + // case the request is sent exactly as the caller built it. + descriptor *spi.Descriptor + hasDescriptor bool } // NewRemoteAPIChat 创建远程 API 聊天实例 @@ -98,19 +105,28 @@ func NewRemoteAPIChat(chatConfig *ChatConfig) (*RemoteAPIChat, error) { } } - return &RemoteAPIChat{ - modelName: modelName, - client: openai.NewClientWithConfig(config), - modelID: chatConfig.ModelID, - baseURL: strings.TrimRight(config.BaseURL, "/"), - apiKey: apiKey, - provider: providerName, - appID: chatConfig.AppID, - appSecret: chatConfig.AppSecret, - customHeaders: chatConfig.CustomHeaders, - adapter: resolveProvider(providerName, modelName), - thinkingOverride: parseThinkingOverride(chatConfig.ExtraConfig), - }, nil + remote := &RemoteAPIChat{ + modelName: modelName, + client: openai.NewClientWithConfig(config), + modelID: chatConfig.ModelID, + baseURL: strings.TrimRight(config.BaseURL, "/"), + apiKey: apiKey, + provider: providerName, + appID: chatConfig.AppID, + appSecret: chatConfig.AppSecret, + customHeaders: chatConfig.CustomHeaders, + adapter: resolveProvider(providerName, modelName), + } + // Resolve against the model name actually sent, which an override may have + // changed: a descriptor's model matcher must see the same name the vendor + // will. + resolveConfig := *chatConfig + resolveConfig.Provider = string(providerName) + resolveConfig.ModelName = modelName + if desc, ok := resolveDescriptor(&resolveConfig); ok { + remote.descriptor, remote.hasDescriptor = &desc, true + } + return remote, nil } // authCreds bundles the credentials passed to the adapter's Auth method. @@ -119,42 +135,97 @@ func (c *RemoteAPIChat) authCreds() authCreds { } // shapedRequest builds the standard request and applies the adapter's message -// transform and parameter shaping (but not thinking, which may wrap the body). +// transform. func (c *RemoteAPIChat) shapedRequest(messages []Message, opts *ChatOptions, isStream bool) openai.ChatCompletionRequest { req := c.BuildChatCompletionRequest(messages, opts, isStream) req.Messages = c.adapter.TransformMessages(req.Messages) - c.adapter.ShapeRequest(&req, opts, isStream) return req } // buildOutbound assembles the final outbound request: the body to send, the -// endpoint override (empty for the standard endpoint), and whether the raw HTTP -// path is required. This is the single place that composes adapter + thinking, -// replacing the former buildRequestCustomizer plumbing. +// endpoint override (empty for the standard endpoint), and whether the raw +// HTTP path is required. +// +// The vendor plan decides the body. Applying it to the SDK request's own JSON +// and comparing the result is what selects the transport: an unchanged body +// means the plugin had nothing vendor-specific to add, so the SDK path stays +// available, while any difference must be sent verbatim because the SDK type +// cannot carry the fields the plan wrote. Deriving that from the actual bytes +// beats maintaining a second list of which providers "need raw HTTP". func (c *RemoteAPIChat) buildOutbound( messages []Message, opts *ChatOptions, isStream bool, ) (body any, endpoint string, useRawHTTP bool, err error) { req := c.shapedRequest(messages, opts, isStream) - thinking := c.thinkingOverride - if thinking == nil { - thinking = c.adapter.Thinking() - } - customBody, useRaw := thinking.Apply(&req, opts, isStream) - body = &req - if customBody != nil { - body = customBody + planChangedBody := false + if c.hasDescriptor { + planned, changed, planErr := c.applyPlan(&req, opts, isStream) + if planErr != nil { + return nil, "", false, planErr + } + if changed { + body, planChangedBody = planned, true + } } + body, err = c.shapeProviderRequest(body, req, messages) if err != nil { return nil, "", false, err } endpoint = c.adapter.Endpoint(c.baseURL, c.modelID, isStream) - useRawHTTP = useRaw || c.adapter.ForceRawHTTP() || endpoint != "" + useRawHTTP = planChangedBody || c.adapter.ForceRawHTTP() || endpoint != "" return body, endpoint, useRawHTTP, nil } +// applyPlan resolves the vendor plan and applies it to the SDK request's JSON +// form, reporting the resulting body and whether it differs. +func (c *RemoteAPIChat) applyPlan( + req *openai.ChatCompletionRequest, opts *ChatOptions, isStream bool, +) (map[string]any, bool, error) { + encoded, err := json.Marshal(req) + if err != nil { + return nil, false, fmt.Errorf("marshal request: %w", err) + } + var canonical map[string]any + if err := json.Unmarshal(encoded, &canonical); err != nil { + return nil, false, fmt.Errorf("decode request: %w", err) + } + // Re-marshal from the map before comparing. A struct marshals in field + // order while a map marshals in key order, so comparing against the + // struct's own bytes would report a difference on every request and force + // the raw path for all of them. + before, err := json.Marshal(canonical) + if err != nil { + return nil, false, fmt.Errorf("marshal canonical request: %w", err) + } + + plan, err := c.descriptor.Plan(spi.Request{ + Model: c.modelName, + Stream: isStream, + Values: opts.ParamValues(), + }) + if err != nil { + return nil, false, err + } + + draft := spi.NewDraft(c.descriptor.Protocol, c.modelName, isStream) + draft.Body = canonical + if err := plan.Apply(draft); err != nil { + return nil, false, err + } + + after, err := json.Marshal(draft.Body) + if err != nil { + return nil, false, fmt.Errorf("marshal planned request: %w", err) + } + // Only a change to the bytes forces the raw path. Notes are diagnostics: a + // pinned value the caller already matched, or a parameter the vendor does + // not offer, both leave the body untouched and should not cost the SDK + // path. + return draft.Body, !bytes.Equal(before, after), nil +} + // logRequest 记录请求日志 func (c *RemoteAPIChat) logRequest(ctx context.Context, req any, isStream bool) { if jsonData, err := json.MarshalIndent(req, "", " "); err == nil { diff --git a/internal/models/chat/remote_api_test.go b/internal/models/chat/remote_api_test.go index 8652373970a..0a5c4c2d041 100644 --- a/internal/models/chat/remote_api_test.go +++ b/internal/models/chat/remote_api_test.go @@ -168,42 +168,52 @@ func TestBuildChatCompletionRequest_GPT5MaxCompletionTokens(t *testing.T) { {"AzureOpenAI gpt-4 (unchanged)", "azure_openai", "gpt-4", false}, } + // The assertions target the outbound body rather than the SDK struct: the + // reasoning dispositions are now declared by the OpenAI plugin and applied + // to the request that is actually sent, so the body is where the contract + // lives. + outbound := func(t *testing.T, c *RemoteAPIChat, opts *ChatOptions) map[string]any { + t.Helper() + body, _, _, err := c.buildOutbound(messages, opts, false) + require.NoError(t, err) + encoded, err := json.Marshal(body) + require.NoError(t, err) + var out map[string]any + require.NoError(t, json.Unmarshal(encoded, &out)) + return out + } + for _, tc := range cases { t.Run(tc.name, func(t *testing.T) { c := build(t, tc.provider, tc.model) - opts := &ChatOptions{ + req := outbound(t, c, &ChatOptions{ Temperature: 0.7, TopP: 0.9, MaxTokens: 128, FrequencyPenalty: 0.1, PresencePenalty: 0.2, - } - req := c.shapedRequest(messages, opts, false) + }) if tc.shouldRewriteMaxT { - assert.Equal(t, 0, req.MaxTokens, "MaxTokens must NOT be sent for GPT-5/o-series") - assert.Equal(t, 128, req.MaxCompletionTokens, "MaxCompletionTokens should be populated from MaxTokens") - assert.EqualValues(t, 0, req.Temperature, "temperature must be omitted") - assert.EqualValues(t, 0, req.TopP, "top_p must be omitted") - assert.EqualValues(t, 0, req.FrequencyPenalty, "frequency_penalty must be omitted") - assert.EqualValues(t, 0, req.PresencePenalty, "presence_penalty must be omitted") + assert.NotContains(t, req, "max_tokens", "max_tokens must NOT be sent for GPT-5/o-series") + assert.EqualValues(t, 128, req["max_completion_tokens"], "max_completion_tokens carries the ceiling") + assert.NotContains(t, req, "temperature", "temperature must be omitted") + assert.NotContains(t, req, "top_p", "top_p must be omitted") + assert.NotContains(t, req, "frequency_penalty", "frequency_penalty must be omitted") + assert.NotContains(t, req, "presence_penalty", "presence_penalty must be omitted") } else { - assert.Equal(t, 128, req.MaxTokens) - assert.Equal(t, 0, req.MaxCompletionTokens) - assert.InDelta(t, 0.7, req.Temperature, 1e-6) + assert.EqualValues(t, 128, req["max_tokens"]) + assert.NotContains(t, req, "max_completion_tokens") + assert.InDelta(t, 0.7, req["temperature"], 1e-6) } }) } t.Run("MaxCompletionTokens takes precedence over MaxTokens", func(t *testing.T) { c := build(t, "openai", "gpt-5.2") - opts := &ChatOptions{ - MaxTokens: 128, - MaxCompletionTokens: 2048, - } - req := c.shapedRequest(messages, opts, false) - assert.Equal(t, 0, req.MaxTokens) - assert.Equal(t, 2048, req.MaxCompletionTokens) + req := outbound(t, c, &ChatOptions{MaxTokens: 128, MaxCompletionTokens: 2048}) + assert.NotContains(t, req, "max_tokens") + assert.EqualValues(t, 2048, req["max_completion_tokens"]) }) } diff --git a/internal/models/chat/thinking.go b/internal/models/chat/thinking.go index 23e92372f7f..9939ba5d7f9 100644 --- a/internal/models/chat/thinking.go +++ b/internal/models/chat/thinking.go @@ -3,170 +3,142 @@ package chat import ( "strings" + "github.com/Tencent/WeKnora/internal/models/llm/encoding" + "github.com/Tencent/WeKnora/internal/models/llm/spi" "github.com/Tencent/WeKnora/internal/models/provider" - "github.com/sashabaranov/go-openai" ) -// ExtraConfigThinkingControl is the model parameters.extra_config key for -// selecting how ChatOptions.Thinking is translated to provider HTTP fields. -// The accepted values mirror the strings the frontend writes (see -// ModelEditorDialog.vue): "none", "enable_thinking", "thinking_type", -// "chat_template_kwargs". -const ExtraConfigThinkingControl = "thinking_control" - -// Wire-format request bodies used by providers that express extended-thinking -// through a non-standard top-level field. They embed the standard OpenAI -// request so all other fields are marshalled unchanged. - -// QwenChatCompletionRequest adds Aliyun Qwen's `enable_thinking` boolean. -type QwenChatCompletionRequest struct { - openai.ChatCompletionRequest - EnableThinking *bool `json:"enable_thinking,omitempty"` -} - -// ThinkingConfig is the `{ "type": "enabled"|"disabled" }` block used by -// LKEAP / Volcengine style providers. -type ThinkingConfig struct { - Type string `json:"type"` -} - -// ThinkingChatCompletionRequest adds the `thinking` object for providers that -// use the `{ "thinking": { "type": ... } }` wire format. -type ThinkingChatCompletionRequest struct { - openai.ChatCompletionRequest - Thinking *ThinkingConfig `json:"thinking,omitempty"` -} - -// ThinkingStrategy encodes how ChatOptions.Thinking is mapped onto a provider's -// HTTP request. Apply returns (customBody, useRawHTTP): -// - (nil, false) means "send the standard OpenAI request unchanged" (the -// caller keeps using the SDK path). -// - a non-nil customBody must be sent verbatim over raw HTTP because it -// carries fields the OpenAI SDK would strip. -// -// When opts.Thinking is nil most strategies emit nothing, deferring to the -// model's own default; the exception is enableThinking{alwaysSend: true} -// (Aliyun Qwen), which must always pin the field. -type ThinkingStrategy interface { - Apply(req *openai.ChatCompletionRequest, opts *ChatOptions, isStream bool) (customBody any, useRawHTTP bool) -} - -// noThinking sends no thinking-related fields at all. -type noThinking struct{} +// Model configuration keys read from parameters.extra_config. +const ( + // ExtraConfigThinkingControl was the manual override selecting how the + // thinking toggle reached a provider. The plugin descriptors now derive + // that from the vendor's documentation, so the key is honored only to keep + // existing configurations working: a stored value still forces the wire + // shape it names. + // + // Deprecated: leave it unset and let the resolved plugin decide. + ExtraConfigThinkingControl = "thinking_control" + + // ExtraConfigProtocol pins the wire protocol for vendors that offer more + // than one, such as OpenAI's Chat Completions and Responses APIs. Empty + // selects the vendor's default. + ExtraConfigProtocol = "protocol" +) -func (noThinking) Apply(*openai.ChatCompletionRequest, *ChatOptions, bool) (any, bool) { - return nil, false +// ThinkingControlNone is the reported control for a model whose plugin sends +// no thinking field at all. +const ThinkingControlNone = "none" + +// legacyThinkingEncoders maps the historical thinking_control values onto the +// encoders they selected, so a stored override keeps doing what it did. +var legacyThinkingEncoders = map[string]spi.Encoder{ + "enable_thinking": encoding.EnableThinkingBool{Key: "enable_thinking"}, + "thinking_type": encoding.ThinkingObject{ + Key: "thinking", On: "enabled", Off: "disabled", + }, + "chat_template_kwargs": encoding.ChatTemplateKwargs{ + Key: "chat_template_kwargs", Arg: "enable_thinking", + }, } -// enableThinking encodes thinking via Qwen's `enable_thinking` boolean. +// applyLegacyThinkingOverride honors a stored extra_config.thinking_control by +// rewriting the resolved descriptor's thinking encoder. // -// - alwaysSend: pin the field even when opts.Thinking is nil (Aliyun Qwen -// thinking models require it on every request; default value is false). -// - disableOnNonStream: force enable_thinking=false for non-stream requests -// (Qwen3 rejects thinking in non-stream mode). -type enableThinking struct { - alwaysSend bool - disableOnNonStream bool -} - -func (s enableThinking) Apply(req *openai.ChatCompletionRequest, opts *ChatOptions, isStream bool) (any, bool) { - thinking := false - switch { - case opts != nil && opts.Thinking != nil: - thinking = *opts.Thinking - case !s.alwaysSend: - return nil, false +// The override predates the plugin catalog and exists because the catalog used +// to be a guess; keeping it working means an operator who pinned a wire format +// against a misdetected provider does not silently start sending a different +// one after this refactor. Descriptors are values, so the override produces a +// copy and never mutates the registry. +func applyLegacyThinkingOverride(desc spi.Descriptor, extraConfig map[string]string) spi.Descriptor { + control := strings.ToLower(strings.TrimSpace(extraConfig[ExtraConfigThinkingControl])) + if control == "" { + return desc } - if s.disableOnNonStream && !isStream { - thinking = false + + params := make([]spi.Param, 0, len(desc.Params)+1) + var mode *spi.Param + for _, p := range desc.Params { + if p.ID == spi.ParamThinkingMode { + clone := p + mode = &clone + continue + } + params = append(params, p) } - qwenReq := QwenChatCompletionRequest{ChatCompletionRequest: *req} - qwenReq.EnableThinking = &thinking - return qwenReq, true -} -// thinkingTypeField encodes thinking via the `{ "thinking": { "type": ... } }` -// object (LKEAP / Volcengine). Emits nothing when opts.Thinking is unset. -type thinkingTypeField struct{} + if control == ThinkingControlNone { + // "none" means send no thinking field, so the parameter goes away + // entirely rather than becoming a knob that does nothing. + desc.Params = params + return desc + } -func (thinkingTypeField) Apply(req *openai.ChatCompletionRequest, opts *ChatOptions, _ bool) (any, bool) { - if opts == nil || opts.Thinking == nil { - return nil, false + encoder, ok := legacyThinkingEncoders[control] + if !ok { + return desc } - r := ThinkingChatCompletionRequest{ChatCompletionRequest: *req} - thinkingType := "disabled" - if *opts.Thinking { - thinkingType = "enabled" + if mode == nil { + // The override also enables the toggle on a model whose plugin does + // not declare one, which is the case it was originally added for. + mode = &spi.Param{} + *mode = encoding.ThinkingMode(encoder, spi.ThinkingOn, spi.ThinkingOff) + } else { + mode.Encode = encoder + // A vendor default belongs to the vendor's own field; forcing another + // one should not also force a value the caller never asked for. + mode.Default = nil } - r.Thinking = &ThinkingConfig{Type: thinkingType} - return r, true + desc.Params = append([]spi.Param{*mode}, params...) + return desc } -// chatTemplateKwargs encodes thinking via the standard request's -// `chat_template_kwargs.enable_thinking` (vLLM / NVIDIA / generic local -// deployments). Emits nothing when opts.Thinking is unset. -type chatTemplateKwargs struct{} - -func (chatTemplateKwargs) Apply(req *openai.ChatCompletionRequest, opts *ChatOptions, _ bool) (any, bool) { - if opts == nil || opts.Thinking == nil { - return nil, false +// EffectiveThinkingControl reports the wire field that will carry the thinking +// toggle for a model, or "none" when the model has no toggle. +// +// It resolves through the same registry and descriptor the request path uses, +// so the answer is derived rather than predicted. That matters because this +// value is what the model editor and the debug drawer show, and the previous +// arrangement — a backend table plus a frontend copy of the same heuristics — +// could disagree with what was actually sent. +func EffectiveThinkingControl(config *ChatConfig) string { + desc, ok := resolveDescriptor(config) + if !ok { + return ThinkingControlNone } - req.ChatTemplateKwargs = map[string]interface{}{ - "enable_thinking": *opts.Thinking, + param, ok := desc.Param(spi.ParamThinkingMode) + if !ok || param.Encode == nil || param.EffectiveSupport() == spi.SupportForbidden { + return ThinkingControlNone } - return req, true + return param.Encode.ID() } -// parseThinkingOverride reads extra_config.thinking_control and returns the -// strategy it selects, or nil when unset (the provider adapter's default -// strategy then applies). An unrecognized non-empty value falls back to -// chat_template_kwargs, preserving the legacy default-mode behavior. -func parseThinkingOverride(extraConfig map[string]string) ThinkingStrategy { - if extraConfig == nil { - return nil - } - switch strings.ToLower(strings.TrimSpace(extraConfig[ExtraConfigThinkingControl])) { - case "": - return nil - case "none": - return noThinking{} - case "enable_thinking": - return enableThinking{} - case "thinking_type": - return thinkingTypeField{} - default: - // "chat_template_kwargs" and any unknown non-empty value. - return chatTemplateKwargs{} - } +// SupportsThinking reports whether a model exposes a reasoning toggle. The UI +// asks this before offering the control, instead of re-deriving it from a +// provider name. +func SupportsThinking(config *ChatConfig) bool { + return EffectiveThinkingControl(config) != ThinkingControlNone } -// EffectiveThinkingControl reports the provider field that will carry -// ChatOptions.Thinking. It intentionally shares the same adapter/override -// resolution as the real request path so diagnostics do not guess from the -// frontend selection. -func EffectiveThinkingControl(config *ChatConfig) string { +// resolveDescriptor looks up the plugin backing a chat configuration, with any +// stored legacy override already applied. Every caller — the request path, the +// reporting helpers, the debug endpoint — goes through here, so none of them +// can disagree about what will be sent. +func resolveDescriptor(config *ChatConfig) (spi.Descriptor, bool) { if config == nil { - return "none" + return spi.Descriptor{}, false } - if override := parseThinkingOverride(config.ExtraConfig); override != nil { - return thinkingStrategyName(override) + name := provider.ProviderName(strings.TrimSpace(config.Provider)) + if name == "" { + name = provider.DetectProvider(config.BaseURL) } - providerName := provider.ProviderName(config.Provider) - if providerName == "" { - providerName = provider.DetectProvider(config.BaseURL) - } - return thinkingStrategyName(resolveProvider(providerName, config.ModelName).Thinking()) -} - -func thinkingStrategyName(strategy ThinkingStrategy) string { - switch strategy.(type) { - case enableThinking: - return "enable_thinking" - case thinkingTypeField: - return "thinking_type" - case chatTemplateKwargs: - return "chat_template_kwargs" - default: - return "none" + desc, ok := spi.Resolve(spi.Query{ + Vendor: string(name), + Kind: spi.KindChat, + Model: config.ModelName, + Protocol: configuredProtocol(config), + }) + if !ok { + return spi.Descriptor{}, false } + return applyLegacyThinkingOverride(desc, config.ExtraConfig), true } diff --git a/internal/models/chat/thinking_test.go b/internal/models/chat/thinking_test.go index 13996997861..cc2adb0ecd0 100644 --- a/internal/models/chat/thinking_test.go +++ b/internal/models/chat/thinking_test.go @@ -1,144 +1,108 @@ package chat import ( - "encoding/json" "testing" - "github.com/sashabaranov/go-openai" "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" ) func ptrBool(b bool) *bool { return &b } -// TestThinkingStrategy_NilThinking verifies the strategies that defer to the -// model default emit nothing when ChatOptions.Thinking is unset. -func TestThinkingStrategy_NilThinking(t *testing.T) { - req := openai.ChatCompletionRequest{Model: "test"} - strategies := []ThinkingStrategy{ - noThinking{}, - enableThinking{}, // not alwaysSend - thinkingTypeField{}, - chatTemplateKwargs{}, - } - for _, s := range strategies { - custom, raw := s.Apply(&req, nil, true) - assert.Nil(t, custom, "%T", s) - assert.False(t, raw, "%T", s) +// TestEffectiveThinkingControl asserts that the reported thinking field is the +// one the request path will actually use. +// +// The value is now the wire field itself rather than a category name, because +// that is the question the setting was always standing in for, and because it +// is read out of the resolved plugin instead of predicted by a second table. +// The frontend copy of those predictions is what this replaces. +func TestEffectiveThinkingControl(t *testing.T) { + cases := []struct { + name string + config *ChatConfig + want string + }{ + { + name: "aliyun qwen uses the boolean", + config: &ChatConfig{Provider: "aliyun", ModelName: "qwen3-32b"}, + want: "enable_thinking", + }, + { + name: "deepseek uses the thinking object", + config: &ChatConfig{Provider: "deepseek", ModelName: "deepseek-chat"}, + want: "thinking.type", + }, + { + name: "self-hosted deployments use the chat template", + config: &ChatConfig{Provider: "generic", ModelName: "qwen3"}, + want: "chat_template_kwargs.enable_thinking", + }, + { + name: "openai reasoning models toggle through the effort ladder", + config: &ChatConfig{Provider: "openai", ModelName: "gpt-5"}, + want: "reasoning_effort", + }, + { + name: "anthropic uses the thinking object", + config: &ChatConfig{Provider: "anthropic", ModelName: "claude-sonnet-4-6"}, + want: "thinking.type", + }, + { + name: "a model without a toggle reports none", + config: &ChatConfig{Provider: "openai", ModelName: "gpt-4o"}, + want: ThinkingControlNone, + }, + { + name: "an unregistered provider reports none", + config: &ChatConfig{Provider: "not-a-vendor", ModelName: "x"}, + want: ThinkingControlNone, + }, + { + name: "nil config reports none", + config: nil, + want: ThinkingControlNone, + }, } -} - -// TestEnableThinking_QwenSemantics pins the Aliyun Qwen behavior: thinking is -// always sent, defaults to false, and is forced off on non-stream requests. -func TestEnableThinking_QwenSemantics(t *testing.T) { - s := enableThinking{alwaysSend: true, disableOnNonStream: true} - req := openai.ChatCompletionRequest{Model: "qwen3-32b"} - - t.Run("non-stream forces false even when requested true", func(t *testing.T) { - custom, raw := s.Apply(&req, &ChatOptions{Thinking: ptrBool(true)}, false) - require.True(t, raw) - qwen, ok := custom.(QwenChatCompletionRequest) - require.True(t, ok) - require.NotNil(t, qwen.EnableThinking) - assert.False(t, *qwen.EnableThinking) - }) - - t.Run("stream honors requested true", func(t *testing.T) { - custom, raw := s.Apply(&req, &ChatOptions{Thinking: ptrBool(true)}, true) - require.True(t, raw) - qwen := custom.(QwenChatCompletionRequest) - require.NotNil(t, qwen.EnableThinking) - assert.True(t, *qwen.EnableThinking) - }) - - t.Run("stream defaults to false when unset", func(t *testing.T) { - custom, raw := s.Apply(&req, nil, true) - require.True(t, raw) - qwen := custom.(QwenChatCompletionRequest) - require.NotNil(t, qwen.EnableThinking) - assert.False(t, *qwen.EnableThinking) - }) -} - -// TestEnableThinking_ExtraConfigSemantics pins the extra_config "enable_thinking" -// override: only sent when explicitly requested. -func TestEnableThinking_ExtraConfigSemantics(t *testing.T) { - s := enableThinking{} - req := openai.ChatCompletionRequest{Model: "qwen3"} - - custom, raw := s.Apply(&req, &ChatOptions{Thinking: ptrBool(true)}, true) - require.True(t, raw) - qwen := custom.(QwenChatCompletionRequest) - require.NotNil(t, qwen.EnableThinking) - assert.True(t, *qwen.EnableThinking) - - custom, raw = s.Apply(&req, nil, true) - assert.Nil(t, custom) - assert.False(t, raw) -} - -func TestThinkingTypeField(t *testing.T) { - s := thinkingTypeField{} - req := openai.ChatCompletionRequest{Model: "ds-v3"} - - custom, raw := s.Apply(&req, &ChatOptions{Thinking: ptrBool(false)}, true) - require.True(t, raw) - typed, ok := custom.(ThinkingChatCompletionRequest) - require.True(t, ok) - require.NotNil(t, typed.Thinking) - assert.Equal(t, "disabled", typed.Thinking.Type) - - custom, raw = s.Apply(&req, &ChatOptions{Thinking: ptrBool(true)}, true) - require.True(t, raw) - assert.Equal(t, "enabled", custom.(ThinkingChatCompletionRequest).Thinking.Type) -} - -func TestChatTemplateKwargs(t *testing.T) { - s := chatTemplateKwargs{} - req := openai.ChatCompletionRequest{Model: "vllm"} - custom, raw := s.Apply(&req, &ChatOptions{Thinking: ptrBool(true)}, true) - require.True(t, raw) - out, ok := custom.(*openai.ChatCompletionRequest) - require.True(t, ok) - assert.Equal(t, true, out.ChatTemplateKwargs["enable_thinking"]) - - body, err := json.Marshal(custom) - require.NoError(t, err) - assert.Contains(t, string(body), "chat_template_kwargs") + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + assert.Equal(t, tc.want, EffectiveThinkingControl(tc.config)) + }) + } } -func TestParseThinkingOverride(t *testing.T) { - cases := map[string]ThinkingStrategy{ - "none": noThinking{}, - "enable_thinking": enableThinking{}, - "thinking_type": thinkingTypeField{}, - "chat_template_kwargs": chatTemplateKwargs{}, - "something-unknown": chatTemplateKwargs{}, // legacy default-mode fallback - } - for value, want := range cases { - got := parseThinkingOverride(map[string]string{ExtraConfigThinkingControl: value}) - assert.IsType(t, want, got, "value=%q", value) +// TestLegacyThinkingControlOverride pins the compatibility path: a stored +// thinking_control still forces the wire format it names, so an operator who +// pinned one against a misdetected provider keeps the behavior they chose. +func TestLegacyThinkingControlOverride(t *testing.T) { + cases := []struct { + control string + want string + }{ + {"enable_thinking", "enable_thinking"}, + {"thinking_type", "thinking.type"}, + {"chat_template_kwargs", "chat_template_kwargs.enable_thinking"}, + {"none", ThinkingControlNone}, } - assert.Nil(t, parseThinkingOverride(nil)) - assert.Nil(t, parseThinkingOverride(map[string]string{})) - assert.Nil(t, parseThinkingOverride(map[string]string{ExtraConfigThinkingControl: ""})) + for _, tc := range cases { + t.Run(tc.control, func(t *testing.T) { + // Aliyun's own default is the boolean, so anything else proves the + // override took effect rather than the vendor default surviving. + got := EffectiveThinkingControl(&ChatConfig{ + Provider: "aliyun", + ModelName: "qwen3-32b", + ExtraConfig: map[string]string{ExtraConfigThinkingControl: tc.control}, + }) + assert.Equal(t, tc.want, got) + }) + } } -func TestEffectiveThinkingControl(t *testing.T) { - assert.Equal(t, "enable_thinking", EffectiveThinkingControl(&ChatConfig{ - Provider: "aliyun", - ModelName: "qwen3-32b", - })) - assert.Equal(t, "chat_template_kwargs", EffectiveThinkingControl(&ChatConfig{ - Provider: "generic", - ModelName: "qwen3", - ExtraConfig: map[string]string{ExtraConfigThinkingControl: "chat_template_kwargs"}, - })) - assert.Equal(t, "none", EffectiveThinkingControl(&ChatConfig{ - Provider: "generic", - ModelName: "qwen3", - ExtraConfig: map[string]string{ExtraConfigThinkingControl: "none"}, - })) +// TestSupportsThinking covers the predicate the UI asks before offering a +// toggle at all. +func TestSupportsThinking(t *testing.T) { + assert.True(t, SupportsThinking(&ChatConfig{Provider: "aliyun", ModelName: "qwen3-32b"})) + assert.True(t, SupportsThinking(&ChatConfig{Provider: "volcengine", ModelName: "doubao-seed-1-6"})) + assert.False(t, SupportsThinking(&ChatConfig{Provider: "openai", ModelName: "gpt-4o"})) + assert.False(t, SupportsThinking(nil)) } diff --git a/internal/models/llm/spi/message.go b/internal/models/llm/spi/message.go index f582a18db34..910e7b812e5 100644 --- a/internal/models/llm/spi/message.go +++ b/internal/models/llm/spi/message.go @@ -179,12 +179,15 @@ func (o *Options) ParamValues() map[ParamID]Value { // EffectiveMaxTokens reports the output ceiling, accepting either spelling the // caller may have used. +// +// MaxCompletionTokens wins when both are set: it is the more specific field, +// and a caller who filled it in was addressing a model that requires it. func (o *Options) EffectiveMaxTokens() int { if o == nil { return 0 } - if o.MaxTokens > 0 { - return o.MaxTokens + if o.MaxCompletionTokens > 0 { + return o.MaxCompletionTokens } - return o.MaxCompletionTokens + return o.MaxTokens } From 1f5881080bb7fbdbcb9be9e47f0c5667e10b4d34 Mon Sep 17 00:00:00 2001 From: lyingbug <[email> Date: Fri, 14 Aug 2026 00:14:45 +0000 Subject: [PATCH 4/9] feat(models): serve the capability manifest and drive the UI from it Add GET /models/capabilities, rendering a plugin descriptor as a form schema: groups, fields, widgets, the vendor's own vocabularies, and the wire field each control writes. It is the fourth surface the one declaration drives, after the request body, validation, and the debug report. The frontend stops predicting vendor behavior. utils/thinkingControl.ts carried its own copy of the Go provider rules under a comment asking readers to keep them aligned; it is deleted. What replaces it reads the manifest: - utils/modelCapabilities.ts holds the types and pure readers, with no I/O so it stays testable on its own. - api/model owns fetching and caches per query, since the manifest is a pure function of the query on the server. - The model editor no longer computes a default wire format. Its control becomes an advanced override defaulting to 'follow the plugin', which reports the field the backend actually resolved. - The debug drawer asks the backend whether a model has a thinking toggle instead of inferring it from the provider name. A stored thinking_control still wins everywhere, so no configuration has to change. Co-authored-by: lyingbug --- frontend/src/api/model/index.ts | 68 ++++++++ frontend/src/components/ModelDebugDrawer.vue | 36 ++++- frontend/src/components/ModelEditorDialog.vue | 127 +++++++-------- frontend/src/i18n/locales/en-US.ts | 10 +- frontend/src/i18n/locales/ko-KR.ts | 10 +- frontend/src/i18n/locales/ru-RU.ts | 10 +- frontend/src/i18n/locales/zh-CN.ts | 10 +- frontend/src/utils/modelCapabilities.test.ts | 65 ++++++++ frontend/src/utils/modelCapabilities.ts | 116 ++++++++++++++ frontend/src/utils/thinkingControl.test.ts | 36 ----- frontend/src/utils/thinkingControl.ts | 89 ----------- internal/handler/model_capability.go | 96 +++++++++++ internal/models/llm/spi/schema.go | 150 ++++++++++++++++++ internal/models/llm/vendors/schema_test.go | 136 ++++++++++++++++ internal/router/routes_infra.go | 3 + 15 files changed, 760 insertions(+), 202 deletions(-) create mode 100644 frontend/src/utils/modelCapabilities.test.ts create mode 100644 frontend/src/utils/modelCapabilities.ts delete mode 100644 frontend/src/utils/thinkingControl.test.ts delete mode 100644 frontend/src/utils/thinkingControl.ts create mode 100644 internal/handler/model_capability.go create mode 100644 internal/models/llm/spi/schema.go create mode 100644 internal/models/llm/vendors/schema_test.go diff --git a/frontend/src/api/model/index.ts b/frontend/src/api/model/index.ts index e18d9657af4..848bd2fd498 100644 --- a/frontend/src/api/model/index.ts +++ b/frontend/src/api/model/index.ts @@ -259,3 +259,71 @@ export function getWeKnoraCloudStatus(): Promise { }) }) } + +// --------------------------------------------------------------------------- +// Model capability manifest +// --------------------------------------------------------------------------- + +import type { ModelCapabilities } from '@/utils/modelCapabilities' + +export interface CapabilityQuery { + provider?: string; + model?: string; + baseUrl?: string; + modelType?: string; + protocol?: string; +} + +function capabilityCacheKey(query: CapabilityQuery): string { + return [ + query.provider ?? '', + query.model ?? '', + query.baseUrl ?? '', + query.modelType ?? 'chat', + query.protocol ?? '', + ].join('|'); +} + +// Capabilities are a pure function of the query on the server, so caching them +// avoids a request per keystroke while a user types a model name. +const capabilityCache = new Map>(); + +/** + * Fetch the capability manifest for a provider and model. + * + * Resolves to null when the backend has no plugin for the provider, which is a + * gap in the catalog rather than an error: the model still works through the + * generic transport, and the caller should fall back to its own defaults. + */ +export function fetchModelCapabilities(query: CapabilityQuery): Promise { + if (!query.provider && !query.baseUrl) { + return Promise.resolve(null); + } + + const key = capabilityCacheKey(query); + const cached = capabilityCache.get(key); + if (cached) return cached; + + const params: Record = { model_type: query.modelType ?? 'chat' }; + if (query.provider) params.provider = query.provider; + if (query.model) params.model = query.model; + if (query.baseUrl) params.base_url = query.baseUrl; + if (query.protocol) params.protocol = query.protocol; + + const request = get('/models/capabilities', params) + .then((res: any) => (res?.data ?? null) as ModelCapabilities | null) + .catch(() => { + // A failed lookup must not block the editor; drop it so a later attempt + // can succeed. + capabilityCache.delete(key); + return null; + }); + + capabilityCache.set(key, request); + return request; +} + +/** Clear the capability cache, for an explicit refresh. */ +export function clearCapabilityCache(): void { + capabilityCache.clear(); +} diff --git a/frontend/src/components/ModelDebugDrawer.vue b/frontend/src/components/ModelDebugDrawer.vue index 00fd0c6d8fc..3f8338decbf 100644 --- a/frontend/src/components/ModelDebugDrawer.vue +++ b/frontend/src/components/ModelDebugDrawer.vue @@ -197,9 +197,9 @@ import { MessagePlugin } from 'tdesign-vue-next' import { useI18n } from 'vue-i18n' import { copyWithToast } from '@/utils/clipboard' import SettingDrawer from '@/components/settings/SettingDrawer.vue' -import { debugModel, type ModelConfig, type ModelDebugResult } from '@/api/model' +import { debugModel, fetchModelCapabilities, type ModelConfig, type ModelDebugResult } from '@/api/model' import { fileSizeVerification } from '@/utils' -import { modelSupportsThinking } from '@/utils/thinkingControl' +import { supportsThinking as capabilitiesSupportThinking } from '@/utils/modelCapabilities' const props = defineProps<{ visible: boolean @@ -242,7 +242,37 @@ let runSequence = 0 const selectedModel = computed(() => props.models.find(model => model.id === selectedModelId.value)) const filteredModels = computed(() => props.models.filter(model => model.type === selectedModelType.value)) const isChat = computed(() => selectedModel.value?.type === 'KnowledgeQA') -const supportsThinking = computed(() => selectedModel.value ? modelSupportsThinking(selectedModel.value) : false) +/** + * Ask the backend whether a saved model has a thinking toggle. Only remote + * chat models can have one, and a stored legacy override still wins because + * the backend still honors it. + */ +async function resolveSupportsThinking(model: ModelConfig): Promise { + if (model.type !== 'KnowledgeQA' || model.source !== 'remote') return false + const capabilities = await fetchModelCapabilities({ + provider: model.parameters.provider || '', + model: model.name || '', + baseUrl: model.parameters.base_url || '', + }) + return capabilitiesSupportThinking( + capabilities, + model.parameters.extra_config?.thinking_control, + ) +} + +// Whether the selected model has a thinking toggle is the backend's answer, +// resolved from the model plugin, so the drawer cannot offer a switch the +// model will ignore. It arrives asynchronously, hence a ref rather than a +// computed. +const supportsThinking = ref(false) +watch( + selectedModel, + async (model) => { + supportsThinking.value = model ? await resolveSupportsThinking(model) : false + if (!supportsThinking.value) thinking.value = false + }, + { immediate: true }, +) const needsFile = computed(() => ['VLLM', 'ASR'].includes(selectedModel.value?.type || '')) const documents = computed(() => documentsText.value.split('\n').map(item => item.trim()).filter(Boolean)) const canRun = computed(() => { diff --git a/frontend/src/components/ModelEditorDialog.vue b/frontend/src/components/ModelEditorDialog.vue index 478346614c0..9c51d72fdef 100644 --- a/frontend/src/components/ModelEditorDialog.vue +++ b/frontend/src/components/ModelEditorDialog.vue @@ -359,8 +359,7 @@ v-model="formData.thinkingControl" :key="`thinking-${formData.id}-${formData.thinkingControl}`" :popup-props="{ overlayClassName: 'thinking-control-select-popup' }" - @change="onThinkingControlManualPick" - > + > activeModelType.value === 'chat' && formData.value.source === 'remote', ) -const resolvedThinkingControl = (): ThinkingControlValue => - defaultThinkingControl( - formData.value.provider || '', - formData.value.modelName || '', - ) +/** + * The capability manifest for the provider and model currently in the form. + * + * It is fetched from the backend, which renders it from the same plugin that + * builds the request. The editor no longer predicts what a provider does; it + * asks, and shows the answer. + */ +const capabilities = ref(null) + +/** The request field this model's thinking toggle will use, or 'none'. */ +const resolvedThinkingWireField = computed(() => thinkingControlOf(capabilities.value)) + +const refreshCapabilities = async () => { + if (!showThinkingControlField.value) { + capabilities.value = null + return + } + capabilities.value = await fetchModelCapabilities({ + provider: formData.value.provider || '', + model: formData.value.modelName || '', + baseUrl: formData.value.baseUrl || '', + }) +} -/** 用户是否手动改过思考参数格式(改过则不再自动覆盖,直到换服务商) */ -const thinkingControlManual = ref(false) /** 正在从 modelData 灌入表单,忽略厂商/来源控件的程序化 change 副作用 */ const hydratingForm = ref(false) -const onThinkingControlManualPick = () => { - thinkingControlManual.value = true -} - -const syncThinkingControlToForm = (force = false) => { - if (!showThinkingControlField.value) return - if (!force && !isEdit.value && thinkingControlManual.value) return - formData.value.thinkingControl = resolvedThinkingControl() -} - const applyThinkingControlFromModelData = () => { if (!props.modelData || activeModelType.value !== 'chat' || formData.value.source !== 'remote') return - thinkingControlManual.value = !!props.modelData.thinkingControl - formData.value.thinkingControl = resolveThinkingControl( - props.modelData.thinkingControl, - formData.value.provider || props.modelData.provider || '', - formData.value.modelName || props.modelData.modelName || '', - ) + // An empty value means "follow the plugin", which is the default now that + // the backend derives the wire format from the vendor's documentation. + formData.value.thinkingControl = props.modelData.thinkingControl || '' } +/** + * The override options. The first entry defers to the plugin and is what + * almost every model should use; the rest exist because a deployment may have + * pinned a format against a misdetected provider, and the backend still + * honors those. + */ const thinkingControlOptions = computed(() => { + const auto = { + value: '', + label: t('model.editor.thinkingControl.auto.label'), + hint: resolvedThinkingWireField.value === THINKING_CONTROL_NONE + ? t('model.editor.thinkingControl.autoNone') + : t('model.editor.thinkingControl.autoField', { field: resolvedThinkingWireField.value }), + } const keys = ['none', 'chatTemplateKwargs', 'enableThinking', 'thinkingType'] as const const values = ['none', 'chat_template_kwargs', 'enable_thinking', 'thinking_type'] as const - return keys.map((key, i) => ({ + return [auto, ...keys.map((key, i) => ({ value: values[i], label: t(`model.editor.thinkingControl.${key}.label`), hint: t(`model.editor.thinkingControl.${key}.hint`), - })) + }))] }) // Header icon for the SettingDrawer — uses the same TDesign icon name table @@ -875,7 +891,7 @@ const formData = ref({ isDefault: false, supportsVision: false, maxConcurrency: undefined, - thinkingControl: defaultThinkingControl('generic', ''), + thinkingControl: '', customHeaders: [], appSecret: '', lkeapRegion: 'ap-guangzhou', @@ -1011,7 +1027,6 @@ const selectModelType = async (type: EditorModelType) => { } if (type !== 'chat') { formData.value.supportsVision = false - thinkingControlManual.value = false } remoteChecked.value = false remoteAvailable.value = false @@ -1026,8 +1041,6 @@ const selectModelType = async (type: EditorModelType) => { handleProviderChange(formData.value.provider || 'generic') } if (showThinkingControlField.value && !isEdit.value) { - thinkingControlManual.value = false - syncThinkingControlToForm(true) } } @@ -1087,8 +1100,6 @@ watch(() => props.visible, (val) => { } if (showThinkingControlField.value && !isEdit.value) { - thinkingControlManual.value = false - syncThinkingControlToForm(true) } } finally { nextTick(() => { @@ -1100,7 +1111,6 @@ watch(() => props.visible, (val) => { // 重置表单 const resetForm = () => { - thinkingControlManual.value = false formData.value = { id: generateId(), name: '', // 保留字段但不使用,保存时用 modelName @@ -1116,7 +1126,7 @@ const resetForm = () => { isDefault: false, supportsVision: false, maxConcurrency: undefined, - thinkingControl: defaultThinkingControl('generic', ''), + thinkingControl: '', customHeaders: [], appSecret: '', lkeapRegion: 'ap-guangzhou', @@ -1159,38 +1169,25 @@ const handleProviderChange = (value: string) => { if (hydratingForm.value) return if (activeModelType.value !== 'chat' || formData.value.source !== 'remote') return if (!isEdit.value) { - thinkingControlManual.value = false - syncThinkingControlToForm(true) return } // 编辑时仅用户主动换厂商才跟随默认 - thinkingControlManual.value = false - syncThinkingControlToForm(true) } +// Refresh the capability manifest whenever the fields identifying the model +// change. The editor no longer recomputes a default here: the plugin owns the +// wire format, and the form only reports it. watch( - () => [formData.value.source, formData.value.provider, formData.value.modelName] as const, - ([source, provider, modelName], [prevSource, prevProvider, prevModelName]) => { - if (hydratingForm.value || isEdit.value) return - if (activeModelType.value !== 'chat' || source !== 'remote') return - if (source === prevSource && provider === prevProvider && modelName === prevModelName) return - - const providerChanged = provider !== prevProvider - - if (providerChanged) { - thinkingControlManual.value = false - syncThinkingControlToForm(true) - return - } - if (!thinkingControlManual.value) { - syncThinkingControlToForm(true) - return - } - const prevDefault = defaultThinkingControl(prevProvider || '', prevModelName || '') - if (formData.value.thinkingControl === prevDefault) { - syncThinkingControlToForm(true) - } + () => [ + formData.value.source, + formData.value.provider, + formData.value.modelName, + formData.value.baseUrl, + ] as const, + () => { + void refreshCapabilities() }, + { immediate: true }, ) // 监听来源变化,重置校验状态(已合并到下面的 watch) @@ -1675,8 +1672,6 @@ watch(() => formData.value.source, () => { && formData.value.source === 'remote' && activeModelType.value === 'chat' ) { - thinkingControlManual.value = false - syncThinkingControlToForm(true) } }) diff --git a/frontend/src/i18n/locales/en-US.ts b/frontend/src/i18n/locales/en-US.ts index 37e870ae2ef..b100088bf72 100755 --- a/frontend/src/i18n/locales/en-US.ts +++ b/frontend/src/i18n/locales/en-US.ts @@ -3939,9 +3939,15 @@ export default { maxConcurrencyLabel: 'Background concurrency limit', maxConcurrencyPlaceholder: '0 = use global default', maxConcurrencyDesc: 'Caps concurrent background (ingestion/enrichment) calls to this model, shared per model across all replicas. 0 or empty falls back to the global default; interactive chat is never affected.', - thinkingControlLabel: 'Thinking mode request format', - thinkingControlDesc: 'Controls how the agent’s “Thinking mode” on/off switch is written to the API. We pre-select based on vendor/model when possible; change it to match your API docs. With “Do not send”, the agent Thinking mode switch has no effect.', + thinkingControlLabel: 'Thinking parameter format (advanced)', + thinkingControlDesc: "The backend's model plugin picks this from the vendor documentation. Override it only when the detected format does not match your deployment. Choosing \"do not send\" disables the agent thinking switch.", thinkingControl: { + auto: { + label: 'Follow the model plugin (recommended)', + hint: 'The backend picks the field from the vendor documentation' + }, + autoField: 'Currently writes: {field}', + autoNone: 'This model has no thinking toggle', none: { label: 'Do not send thinking fields', hint: 'Agent “Thinking mode” switch has no effect; thinking parameters are not sent in requests' diff --git a/frontend/src/i18n/locales/ko-KR.ts b/frontend/src/i18n/locales/ko-KR.ts index 6fbf69d987a..3cbf22451d3 100755 --- a/frontend/src/i18n/locales/ko-KR.ts +++ b/frontend/src/i18n/locales/ko-KR.ts @@ -2312,8 +2312,8 @@ export default { maxConcurrencyLabel: '백그라운드 동시 실행 상한', maxConcurrencyPlaceholder: '0이면 전역 기본값 사용', maxConcurrencyDesc: '문서 인덱싱/보강 등 백그라운드 작업이 이 모델을 호출하는 동시 실행 수를 제한합니다(모델별로 모든 복제본이 공유). 0 또는 비워 두면 전역 기본값을 사용하며, 대화형 채팅에는 영향을 주지 않습니다.', - thinkingControlLabel: '사고 모드 매개변수 형식', - thinkingControlDesc: '에이전트 「사고 모드」 켜기/끄기 시 API에 어떻게 기록할지 결정합니다. 벤더/모델에 따라 미리 선택되며, 실제 API와 다르면 문서에 맞게 수정하세요. 「전송 안 함」을 선택하면 에이전트 「사고 모드」 스위치가 효과가 없습니다.', + thinkingControlLabel: '사고 모드 파라미터 형식 (고급)', + thinkingControlDesc: '백엔드의 모델 플러그인이 공급업체 공식 문서에 따라 결정합니다. 자동 인식이 실제와 다를 때만 수동으로 지정하세요.', dimensionHint: '모델이 선택되었습니다. "차원 감지" 버튼을 클릭하여 벡터 차원을 자동으로 가져옵니다', loadModelListFailed: '모델 목록 로드 실패', listRefreshed: '목록이 새로고침되었습니다', @@ -2443,6 +2443,12 @@ export default { baseUrlInvalid: 'Base URL 형식이 올바르지 않습니다. 유효한 URL을 입력해주세요' }, thinkingControl: { + auto: { + label: '모델 플러그인 따르기 (권장)', + hint: '백엔드가 공급업체 문서에 따라 필드를 선택합니다' + }, + autoField: '현재 전송 필드: {field}', + autoNone: '이 모델에는 사고 전환 스위치가 없습니다', thinkingType: { label: 'thinking.type', hint: 'Volcengine Ark; Tencent LKEAP (DeepSeek V3 등, LKEAP 기본값; R1은 「전송 안 함」)' diff --git a/frontend/src/i18n/locales/ru-RU.ts b/frontend/src/i18n/locales/ru-RU.ts index 228b28df555..7dbfc56d798 100755 --- a/frontend/src/i18n/locales/ru-RU.ts +++ b/frontend/src/i18n/locales/ru-RU.ts @@ -2312,8 +2312,8 @@ export default { maxConcurrencyLabel: 'Лимит фоновой параллельности', maxConcurrencyPlaceholder: '0 — использовать глобальное значение', maxConcurrencyDesc: 'Ограничивает число одновременных фоновых вызовов (индексация/обогащение) к этой модели, общее для модели по всем репликам. 0 или пусто — используется глобальное значение по умолчанию; интерактивный чат не затрагивается.', - thinkingControlLabel: 'Формат параметров режима размышления', - thinkingControlDesc: 'Определяет, как переключатель «Режим размышления» агента записывается в API. При возможности выбирается по поставщику/модели; при несоответствии измените по документации API. При выборе «Не отправлять» переключатель «Режим размышления» агента не действует.', + thinkingControlLabel: 'Формат параметра размышлений (расширенный)', + thinkingControlDesc: 'Плагин модели на бэкенде выбирает его по документации поставщика. Переопределяйте только если автоопределение не совпадает с вашим развёртыванием.', dimensionHint: 'Модель выбрана. Нажмите «Определить размерность», чтобы автоматически получить значение.', loadModelListFailed: 'Не удалось загрузить список моделей', listRefreshed: 'Список обновлён', @@ -2443,6 +2443,12 @@ export default { baseUrlInvalid: 'Недопустимый Base URL, введите корректный адрес' }, thinkingControl: { + auto: { + label: 'Следовать плагину модели (рекомендуется)', + hint: 'Бэкенд выбирает поле по документации поставщика' + }, + autoField: 'Сейчас отправляется: {field}', + autoNone: 'У этой модели нет переключателя размышлений', thinkingType: { label: 'thinking.type', hint: 'Volcengine Ark; Tencent LKEAP (DeepSeek V3 и др.; по умолчанию для LKEAP; для R1 — «Не отправлять»)' diff --git a/frontend/src/i18n/locales/zh-CN.ts b/frontend/src/i18n/locales/zh-CN.ts index d3d9a1b9af0..4c7601b9b7c 100755 --- a/frontend/src/i18n/locales/zh-CN.ts +++ b/frontend/src/i18n/locales/zh-CN.ts @@ -2314,8 +2314,8 @@ export default { maxConcurrencyLabel: '后台并发上限', maxConcurrencyPlaceholder: '0 表示使用全局默认', maxConcurrencyDesc: '限制文档入库/富化等后台任务对该模型的并发调用数(按模型全副本共享)。0 或留空表示沿用全局默认;不影响交互式对话。', - thinkingControlLabel: '思考模式参数格式', - thinkingControlDesc: '决定智能体「思考模式」开/关时如何写入 API。已尝试按厂商/模型预选,若与实际情况不符请按 API 文档手动修改;选「不写入」时,智能体「思考模式」开关不生效。', + thinkingControlLabel: '思考模式参数格式(高级)', + thinkingControlDesc: '默认由后端的模型插件按厂商官方文档决定;只有当自动识别与实际情况不符时才需要手动指定。选「不写入」时,智能体「思考模式」开关不生效。', dimensionHint: '模型已选择,点击"检测维度"按钮自动获取向量维度', loadModelListFailed: '加载模型列表失败', listRefreshed: '列表已刷新', @@ -2445,6 +2445,12 @@ export default { baseUrlInvalid: 'Base URL 格式不正确,请输入有效的 URL' }, thinkingControl: { + auto: { + label: '跟随模型插件(推荐)', + hint: '由后端按厂商官方文档决定写入哪个字段' + }, + autoField: '当前将写入:{field}', + autoNone: '该模型没有思考开关', thinkingType: { label: 'thinking.type', hint: '火山引擎 Ark;腾讯云 LKEAP(DeepSeek V3 等,选 LKEAP 时默认此项;R1 请改「不写入」)' diff --git a/frontend/src/utils/modelCapabilities.test.ts b/frontend/src/utils/modelCapabilities.test.ts new file mode 100644 index 00000000000..ba5a2c0e8d9 --- /dev/null +++ b/frontend/src/utils/modelCapabilities.test.ts @@ -0,0 +1,65 @@ +import assert from 'node:assert/strict' +import test from 'node:test' + +import { + findField, + supportsThinking, + thinkingControlOf, + THINKING_CONTROL_NONE, + type ModelCapabilities, +} from './modelCapabilities.ts' + +// The provider table that used to live here is gone: it duplicated the Go +// provider adapters, had to be kept aligned by hand, and drifted. The backend +// now serves the same facts it acts on, so what is left to test is that this +// module reads the manifest correctly rather than that it predicts the backend. + +function manifest(overrides: Partial = {}): ModelCapabilities { + return { + vendor: 'aliyun', + protocol: 'openai-chat', + supports_thinking: true, + groups: [{ + key: 'thinking', + fields: [{ + id: 'thinking.mode', + kind: 'enum', + widget: 'select', + wire_field: 'enable_thinking', + options: [{ value: 'on' }, { value: 'off' }], + }], + }], + ...overrides, + } +} + +test('thinkingControlOf reports the wire field the backend resolved', () => { + assert.equal(thinkingControlOf(manifest()), 'enable_thinking') +}) + +test('a model without a thinking control reports none', () => { + const withoutThinking = manifest({ supports_thinking: false, groups: [] }) + assert.equal(thinkingControlOf(withoutThinking), THINKING_CONTROL_NONE) +}) + +test('a missing manifest reports none rather than throwing', () => { + assert.equal(thinkingControlOf(null), THINKING_CONTROL_NONE) +}) + +test('findField locates a field across groups', () => { + const twoGroups = manifest({ + groups: [ + { key: 'sampling', fields: [{ id: 'temperature', kind: 'float', widget: 'slider' }] }, + ...manifest().groups, + ], + }) + assert.equal(findField(twoGroups, 'thinking.mode')?.wire_field, 'enable_thinking') + assert.equal(findField(twoGroups, 'temperature')?.widget, 'slider') + assert.equal(findField(twoGroups, 'nope'), undefined) +}) + +test('a stored legacy override still decides, because the backend honors it', () => { + assert.equal(supportsThinking(manifest({ supports_thinking: false }), 'thinking_type'), true) + assert.equal(supportsThinking(manifest(), 'none'), false) + assert.equal(supportsThinking(manifest()), true) +}) diff --git a/frontend/src/utils/modelCapabilities.ts b/frontend/src/utils/modelCapabilities.ts new file mode 100644 index 00000000000..8f8c3a6dd9b --- /dev/null +++ b/frontend/src/utils/modelCapabilities.ts @@ -0,0 +1,116 @@ +/** + * Model capability manifest, served by GET /models/capabilities. + * + * The backend renders it from the same plugin descriptor the request path + * uses, so a control this form offers is a control that will actually reach + * the vendor. This replaces the provider heuristics that used to live in the + * frontend and had to be kept aligned with the Go code by hand — a comment + * asking future readers to do that is not a mechanism, and the two drifted. + * + * This module is deliberately free of I/O: fetching lives in the model API + * client, so the reading helpers stay testable on their own. + */ + +export type ValueKind = 'bool' | 'int' | 'float' | 'enum' +export type Widget = 'switch' | 'select' | 'number' | 'slider' + +export interface EnumOption { + value: string + label_key?: string + help_key?: string +} + +export interface CapabilityValue { + kind: ValueKind + bool?: boolean + num?: number + str?: string +} + +export interface FieldSchema { + id: string + kind: ValueKind + widget: Widget + label_key?: string + help_key?: string + options?: EnumOption[] + min?: number + max?: number + default?: CapabilityValue + /** The request field this control writes, shown as a hint in the editor. */ + wire_field?: string + doc_url?: string +} + +export interface GroupSchema { + key: string + fields: FieldSchema[] +} + +export interface ModelCapabilities { + vendor: string + display_name?: string + protocol: string + protocols?: string[] + groups: GroupSchema[] + supports_thinking: boolean + reasoning_replay?: string + doc_url?: string +} + +/** The neutral parameter ids the editor and the debug drawer refer to. */ +export const PARAM_THINKING_MODE = 'thinking.mode' +export const PARAM_THINKING_EFFORT = 'thinking.effort' +export const PARAM_THINKING_BUDGET = 'thinking.budget' + +/** Reported when a model has no thinking toggle at all. */ +export const THINKING_CONTROL_NONE = 'none' + +/** + * The legacy override values a model may still have stored in + * `extra_config.thinking_control`. The backend continues to honor them, so the + * editor keeps offering them for deployments that pinned one. + */ +export const LEGACY_THINKING_CONTROLS = [ + 'none', + 'chat_template_kwargs', + 'enable_thinking', + 'thinking_type', +] as const + +/** Find one field in a manifest, across groups. */ +export function findField( + capabilities: ModelCapabilities | null, + id: string, +): FieldSchema | undefined { + if (!capabilities) return undefined + for (const group of capabilities.groups ?? []) { + const field = group.fields.find(f => f.id === id) + if (field) return field + } + return undefined +} + +/** + * The request field carrying this model's thinking toggle, or 'none'. + * + * The value is a wire field such as 'enable_thinking' or 'thinking.type' + * rather than a category name, because the field is what the old setting was + * always standing in for. + */ +export function thinkingControlOf(capabilities: ModelCapabilities | null): string { + return findField(capabilities, PARAM_THINKING_MODE)?.wire_field ?? THINKING_CONTROL_NONE +} + +/** + * Whether a model exposes a thinking toggle, honoring a stored legacy override + * because the backend still does. + */ +export function supportsThinking( + capabilities: ModelCapabilities | null, + storedControl?: string, +): boolean { + const stored = storedControl?.trim().toLowerCase() + if (stored) return stored !== THINKING_CONTROL_NONE + return capabilities?.supports_thinking ?? false +} diff --git a/frontend/src/utils/thinkingControl.test.ts b/frontend/src/utils/thinkingControl.test.ts deleted file mode 100644 index 6502bc0d4e9..00000000000 --- a/frontend/src/utils/thinkingControl.test.ts +++ /dev/null @@ -1,36 +0,0 @@ -import assert from 'node:assert/strict' -import test from 'node:test' - -import { defaultThinkingControl } from './thinkingControl.ts' - -// Cases mirror internal/models/chat/provider_test.go TestResolveProvider thinking defaults. -test('defaultThinkingControl matches backend provider adapters', () => { - const cases: Array<[string, string, ReturnType]> = [ - ['generic', 'anything', 'chat_template_kwargs'], - ['nvidia', 'anything', 'chat_template_kwargs'], - ['volcengine', 'doubao', 'thinking_type'], - ['aliyun', 'qwen3-32b', 'enable_thinking'], - ['aliyun', 'qwen-plus', 'enable_thinking'], - ['aliyun', 'gpt-4', 'none'], - ['lkeap', '', 'thinking_type'], - ['lkeap', 'deepseek-v3.1', 'thinking_type'], - ['lkeap', 'deepseek-r1', 'none'], - ['openai', 'gpt-4o', 'none'], - ['openai', 'gpt-5', 'none'], - ['azure_openai', 'gpt-4', 'none'], - ['deepseek', 'deepseek-chat', 'none'], - ['zhipu', 'glm-4', 'none'], - ['gemini', 'gemini-2.0', 'none'], - ['siliconflow', 'qwen3-8b', 'none'], - ['hunyuan', 'hunyuan-turbo', 'none'], - ['moonshot', 'moonshot-v1-8k', 'none'], - ['weknoracloud', 'anything', 'none'], - ] - for (const [provider, model, want] of cases) { - assert.equal( - defaultThinkingControl(provider, model), - want, - `${provider}/${model}`, - ) - } -}) diff --git a/frontend/src/utils/thinkingControl.ts b/frontend/src/utils/thinkingControl.ts deleted file mode 100644 index 122bced7a56..00000000000 --- a/frontend/src/utils/thinkingControl.ts +++ /dev/null @@ -1,89 +0,0 @@ -/** Mirrors backend internal/models/provider.IsQwenThinkingModel */ -export function isQwenThinkingModel(modelName: string): boolean { - const lower = modelName.trim().toLowerCase() - return ( - lower.startsWith('qwen3') - || lower.startsWith('qwen-plus') - || lower.startsWith('qwen-max') - || lower.startsWith('qwen-turbo') - ) -} - -/** Mirrors backend internal/models/provider.IsLKEAPDeepSeekR1Model */ -export function isLkeapDeepSeekR1Model(modelName: string): boolean { - return modelName.toLowerCase().includes('deepseek-r1') -} - -export type ThinkingControlValue = - | 'none' - | 'chat_template_kwargs' - | 'enable_thinking' - | 'thinking_type' - -const THINKING_CONTROL_VALUES: ThinkingControlValue[] = [ - 'none', - 'chat_template_kwargs', - 'enable_thinking', - 'thinking_type', -] - -/** - * Default thinking_control for provider+model. - * Must stay aligned with chat.resolveProvider(...).Thinking() in provider.go. - */ -export function defaultThinkingControl( - provider: string, - modelName = '', -): ThinkingControlValue { - const p = provider.trim().toLowerCase() - const model = modelName.trim() - - switch (p) { - case 'aliyun': - return isQwenThinkingModel(model) ? 'enable_thinking' : 'none' - case 'lkeap': - // R1 系列后端不发 thinking 参数;其余(含未填模型名)按 LKEAP 的 thinking.type 格式预选 - if (model && isLkeapDeepSeekR1Model(model)) return 'none' - return 'thinking_type' - case 'generic': - case 'nvidia': - return 'chat_template_kwargs' - case 'volcengine': - return 'thinking_type' - default: - // openai, azure_openai, anthropic, zhipu, deepseek, gemini, siliconflow, - // hunyuan, moonshot, openrouter, weknoracloud, … → baseProvider / noThinking - return 'none' - } -} - -/** Resolve stored extra_config or fall back to the provider default. */ -export function resolveThinkingControl( - saved: string | undefined, - provider: string, - modelName = '', -): ThinkingControlValue { - const v = saved?.trim().toLowerCase() - if (THINKING_CONTROL_VALUES.includes(v as ThinkingControlValue)) { - return v as ThinkingControlValue - } - return defaultThinkingControl(provider, modelName) -} - -/** Whether the debug drawer should expose Think on/off for this saved model. */ -export function modelSupportsThinking(model: { - type: string - source: string - name: string - parameters: { - provider?: string - extra_config?: { thinking_control?: string } - } -}): boolean { - if (model.type !== 'KnowledgeQA' || model.source !== 'remote') return false - return resolveThinkingControl( - model.parameters.extra_config?.thinking_control, - model.parameters.provider || '', - model.name || '', - ) !== 'none' -} diff --git a/internal/handler/model_capability.go b/internal/handler/model_capability.go new file mode 100644 index 00000000000..94966b61c6e --- /dev/null +++ b/internal/handler/model_capability.go @@ -0,0 +1,96 @@ +package handler + +import ( + "net/http" + "strings" + + "github.com/Tencent/WeKnora/internal/logger" + "github.com/Tencent/WeKnora/internal/models/llm/spi" + _ "github.com/Tencent/WeKnora/internal/models/llm/vendors" // register the built-in model plugins + "github.com/Tencent/WeKnora/internal/models/provider" + secutils "github.com/Tencent/WeKnora/internal/utils" + "github.com/gin-gonic/gin" +) + +// GetModelCapabilities godoc +// @Summary 获取模型能力清单 +// @Description 返回指定厂商+模型的参数能力清单,用于前端动态渲染模型配置表单 +// @Tags 模型管理 +// @Accept json +// @Produce json +// @Param provider query string true "厂商标识,如 openai / aliyun / anthropic" +// @Param model query string false "模型名称,用于命中模型级插件" +// @Param model_type query string false "模型类型,默认 chat" +// @Param protocol query string false "协议,厂商支持多协议时使用" +// @Success 200 {object} map[string]interface{} "能力清单" +// @Security Bearer +// @Security ApiKeyAuth +// @Router /models/capabilities [get] +// +// The response is rendered from the same descriptor the request path uses, so +// the form can only offer controls that will actually reach the vendor. That +// is the point of the endpoint: the frontend used to predict this with its own +// copy of the provider rules, which had to be kept in sync by hand. +func (h *ModelHandler) GetModelCapabilities(c *gin.Context) { + ctx := c.Request.Context() + + vendor := strings.TrimSpace(c.Query("provider")) + model := strings.TrimSpace(c.Query("model")) + baseURL := strings.TrimSpace(c.Query("base_url")) + if vendor == "" && baseURL != "" { + // A user who has pasted a base URL but not chosen a provider still + // gets the right form, using the same detection the request path uses. + vendor = string(provider.DetectProvider(baseURL)) + } + if vendor == "" { + c.JSON(http.StatusBadRequest, gin.H{ + "success": false, + "error": "provider is required", + }) + return + } + + kind := modelKindFromQuery(c.Query("model_type")) + logger.Infof(ctx, "Model capabilities requested: provider=%s model=%s kind=%s", + secutils.SanitizeForLog(vendor), secutils.SanitizeForLog(model), kind) + + desc, ok := spi.Resolve(spi.Query{ + Vendor: vendor, + Kind: kind, + Model: model, + Protocol: spi.ProtocolID(strings.TrimSpace(c.Query("protocol"))), + }) + if !ok { + // A vendor with no plugin is a gap in the catalog, not an error: the + // model still works through the generic transport, and the form should + // fall back to its built-in defaults rather than break. + c.JSON(http.StatusOK, gin.H{ + "success": true, + "data": nil, + }) + return + } + + schema := desc.Schema() + schema.Protocols = spi.Protocols(vendor, kind) + c.JSON(http.StatusOK, gin.H{ + "success": true, + "data": schema, + }) +} + +// modelKindFromQuery maps the frontend's model-type vocabulary onto the seam's. +func modelKindFromQuery(modelType string) spi.ModelKind { + switch strings.ToLower(strings.TrimSpace(modelType)) { + case "embedding": + return spi.KindEmbedding + case "rerank": + return spi.KindRerank + case "vllm", "vision": + return spi.KindVision + case "asr": + return spi.KindASR + default: + return spi.KindChat + } +} diff --git a/internal/models/llm/spi/schema.go b/internal/models/llm/spi/schema.go new file mode 100644 index 00000000000..7ca1da0c829 --- /dev/null +++ b/internal/models/llm/spi/schema.go @@ -0,0 +1,150 @@ +package spi + +import "sort" + +// This file renders a descriptor as a form schema. +// +// It is the fourth surface the declaration drives, after the request body, +// validation, and the debug report. Serving the form from the same source +// removes the failure this seam was built to end: a frontend that predicted a +// vendor's behavior with its own copy of the rules, and drifted. + +// FieldSchema describes one control a form should render. +type FieldSchema struct { + // ID is the neutral parameter identity the form submits back. + ID ParamID `json:"id"` + // Kind is the value domain. + Kind ValueKind `json:"kind"` + // Widget is the control to render. + Widget Widget `json:"widget"` + // LabelKey and HelpKey are i18n keys the frontend resolves. The backend + // sends no display text, so language stays a frontend concern. + LabelKey string `json:"label_key,omitempty"` + HelpKey string `json:"help_key,omitempty"` + // Options is the vendor's vocabulary for an enum field. + Options []EnumOption `json:"options,omitempty"` + // Min and Max bound a numeric field. + Min *float64 `json:"min,omitempty"` + Max *float64 `json:"max,omitempty"` + // Default is the value sent when the user leaves the field alone. + Default *Value `json:"default,omitempty"` + // WireField names the request field this control writes, which is what + // the model editor shows in place of the old free-text wire-format + // setting — and unlike that setting, it is read from the request path + // rather than chosen independently of it. + WireField string `json:"wire_field,omitempty"` + // DocURL links the vendor documentation behind the field. + DocURL string `json:"doc_url,omitempty"` +} + +// GroupSchema is one section of the form. +type GroupSchema struct { + // Key names the section, e.g. "thinking". + Key string `json:"key"` + // Fields are the controls in declaration order. + Fields []FieldSchema `json:"fields"` +} + +// FormSchema is everything a form needs to render a model's settings. +type FormSchema struct { + // Vendor is the resolved plugin. + Vendor string `json:"vendor"` + // DisplayName is the plugin's human-readable name. + DisplayName string `json:"display_name,omitempty"` + // Protocol is the wire protocol this model will use. + Protocol ProtocolID `json:"protocol"` + // Protocols lists the alternatives, for vendors offering more than one. A + // single entry means there is nothing to choose and the form should not + // offer a selector. + Protocols []ProtocolID `json:"protocols,omitempty"` + // Groups are the form sections. + Groups []GroupSchema `json:"groups"` + // SupportsThinking reports whether the model has a reasoning toggle, which + // is the one question the chat UI asks outside the form. + SupportsThinking bool `json:"supports_thinking"` + // ReasoningReplay reports whether prior reasoning must be replayed. + ReasoningReplay ReasoningReplay `json:"reasoning_replay,omitempty"` + // DocURL links the vendor documentation this plugin follows. + DocURL string `json:"doc_url,omitempty"` +} + +// groupOrder fixes the section order so the form does not reshuffle when a +// vendor happens to declare its parameters in a different sequence. +var groupOrder = map[string]int{ + "thinking": 10, + "sampling": 20, + "limits": 30, +} + +// Schema renders a descriptor as a form schema, omitting the parameters a form +// must not offer: a hidden one, and anything pinned or forbidden, whose value +// the user cannot influence. +func (d Descriptor) Schema() FormSchema { + schema := FormSchema{ + Vendor: d.Vendor, + DisplayName: d.DisplayName, + Protocol: d.Protocol, + ReasoningReplay: d.EffectiveReplay(), + DocURL: d.DocURL, + Groups: []GroupSchema{}, + } + + byGroup := map[string][]FieldSchema{} + order := map[ParamID]int{} + for _, p := range d.Params { + if p.ID == ParamThinkingMode && p.EffectiveSupport() != SupportForbidden { + schema.SupportsThinking = true + } + if p.UI.Hidden || p.EffectiveSupport() != SupportUser { + continue + } + field := FieldSchema{ + ID: p.ID, + Kind: p.Kind, + Widget: p.EffectiveWidget(), + LabelKey: p.UI.LabelKey, + HelpKey: p.UI.HelpKey, + Options: p.Enum, + Min: p.Min, + Max: p.Max, + Default: p.Default, + DocURL: p.DocURL, + } + if p.Encode != nil { + field.WireField = p.Encode.ID() + } + group := p.UI.Group + if group == "" { + group = "advanced" + } + order[p.ID] = p.UI.Order + byGroup[group] = append(byGroup[group], field) + } + + keys := make([]string, 0, len(byGroup)) + for key := range byGroup { + keys = append(keys, key) + } + sort.SliceStable(keys, func(i, j int) bool { + oi, oki := groupOrder[keys[i]] + oj, okj := groupOrder[keys[j]] + if oki != okj { + return oki + } + if oi != oj { + return oi < oj + } + return keys[i] < keys[j] + }) + + for _, key := range keys { + fields := byGroup[key] + // Declaration order is the tiebreak, so a vendor that adds a field + // does not reorder the ones already there. + sort.SliceStable(fields, func(i, j int) bool { + return order[fields[i].ID] < order[fields[j].ID] + }) + schema.Groups = append(schema.Groups, GroupSchema{Key: key, Fields: fields}) + } + return schema +} diff --git a/internal/models/llm/vendors/schema_test.go b/internal/models/llm/vendors/schema_test.go new file mode 100644 index 00000000000..5ad1bccaadc --- /dev/null +++ b/internal/models/llm/vendors/schema_test.go @@ -0,0 +1,136 @@ +package vendors + +import ( + "testing" + + "github.com/Tencent/WeKnora/internal/models/llm/spi" +) + +// The form schema is the fourth surface the descriptors drive. These tests +// assert that it reports the same facts the request path acts on, because the +// bug this seam replaces was precisely a form that disagreed with the wire. + +func schemaFor(t *testing.T, vendor, model string) spi.FormSchema { + t.Helper() + desc, ok := spi.Resolve(spi.Query{Vendor: vendor, Kind: spi.KindChat, Model: model}) + if !ok { + t.Fatalf("no descriptor for %s/%s", vendor, model) + } + return desc.Schema() +} + +func findField(schema spi.FormSchema, id spi.ParamID) (spi.FieldSchema, string, bool) { + for _, group := range schema.Groups { + for _, field := range group.Fields { + if field.ID == id { + return field, group.Key, true + } + } + } + return spi.FieldSchema{}, "", false +} + +// The form must name the wire field a control writes, and it must be the field +// the encoder actually uses rather than a category chosen separately. +func TestSchemaReportsTheWireField(t *testing.T) { + cases := []struct { + vendor, model string + wantWire string + }{ + {"aliyun", "qwen3-max", "enable_thinking"}, + {"deepseek", "deepseek-chat", "thinking.type"}, + {"volcengine", "doubao-seed-1-6", "thinking.type"}, + {"generic", "qwen3-32b", "chat_template_kwargs.enable_thinking"}, + {"anthropic", "claude-sonnet-4-6", "thinking.type"}, + } + + for _, tc := range cases { + t.Run(tc.vendor+"/"+tc.model, func(t *testing.T) { + schema := schemaFor(t, tc.vendor, tc.model) + if !schema.SupportsThinking { + t.Fatalf("%s should report a thinking toggle", tc.vendor) + } + field, group, ok := findField(schema, spi.ParamThinkingMode) + if !ok { + t.Fatalf("no thinking control in the form") + } + if field.WireField != tc.wantWire { + t.Errorf("wire field = %q, want %q", field.WireField, tc.wantWire) + } + if group != "thinking" { + t.Errorf("thinking control landed in group %q", group) + } + if field.Widget != spi.WidgetSelect { + t.Errorf("widget = %q, want select", field.Widget) + } + }) + } +} + +// Only Volcengine documents a third thinking state, so only its form may offer +// one. A shared three-value control would send `auto` to vendors that reject it. +func TestSchemaOffersAutoOnlyWhereDocumented(t *testing.T) { + ark, _, ok := findField(schemaFor(t, "volcengine", "doubao-seed-1-6"), spi.ParamThinkingMode) + if !ok { + t.Fatal("volcengine should expose a thinking control") + } + if len(ark.Options) != 3 { + t.Errorf("volcengine should offer three modes, got %d", len(ark.Options)) + } + + deepseek, _, ok := findField(schemaFor(t, "deepseek", "deepseek-chat"), spi.ParamThinkingMode) + if !ok { + t.Fatal("deepseek should expose a thinking control") + } + if len(deepseek.Options) != 2 { + t.Errorf("deepseek should offer two modes, got %d", len(deepseek.Options)) + } +} + +// A parameter the user cannot influence must not appear as a control: a +// forbidden one would do nothing, and a pinned one would lie about its value. +func TestSchemaHidesParametersTheUserCannotSet(t *testing.T) { + reasoning := schemaFor(t, "openai", "o3-mini") + for _, hidden := range []spi.ParamID{spi.ParamTemperature, spi.ParamTopP} { + if _, _, ok := findField(reasoning, hidden); ok { + t.Errorf("%s is forbidden on reasoning models and must not be offered", hidden) + } + } + if _, _, ok := findField(reasoning, spi.ParamMaxTokens); !ok { + t.Error("the output ceiling is still settable and should be offered") + } + + moonshot := schemaFor(t, "moonshot", "moonshot-v1-8k") + if _, _, ok := findField(moonshot, spi.ParamTemperature); ok { + t.Error("a pinned temperature must not be offered as a control") + } +} + +// Effort ladders differ per vendor, and the form must show the vendor's own. +func TestSchemaCarriesEachVendorsEffortLadder(t *testing.T) { + deepseek, _, ok := findField(schemaFor(t, "deepseek", "deepseek-chat"), spi.ParamThinkingEffort) + if !ok { + t.Fatal("deepseek publishes an effort ladder") + } + if len(deepseek.Options) != 3 { + t.Errorf("deepseek ladder has %d rungs, want 3", len(deepseek.Options)) + } + + if _, _, ok := findField(schemaFor(t, "zhipu", "glm-4.6"), spi.ParamThinkingEffort); ok { + t.Error("GLM-4.6 has no effort ladder and must not be offered one") + } + if _, _, ok := findField(schemaFor(t, "zhipu", "glm-5.2"), spi.ParamThinkingEffort); !ok { + t.Error("GLM-5.2 publishes an effort ladder") + } +} + +// A vendor offering two protocols must say so, and one offering a single +// protocol must not make the user choose. +func TestProtocolsReportedPerVendor(t *testing.T) { + if got := spi.Protocols("openai", spi.KindChat); len(got) != 2 { + t.Errorf("openai protocols = %v, want Chat Completions and Responses", got) + } + if got := spi.Protocols("deepseek", spi.KindChat); len(got) != 1 { + t.Errorf("deepseek protocols = %v, want one", got) + } +} diff --git a/internal/router/routes_infra.go b/internal/router/routes_infra.go index 3631273b213..196db7fc28c 100644 --- a/internal/router/routes_infra.go +++ b/internal/router/routes_infra.go @@ -22,6 +22,9 @@ func RegisterModelRoutes( { // 获取模型厂商列表 — Viewer+ models.GET("/providers", g.Viewer(), handler.ListModelProviders) + // Capability manifest for one provider+model, used by the model editor + // to render the form from the same declaration the request path uses. + models.GET("/capabilities", g.Viewer(), handler.GetModelCapabilities) // 创建模型 — Admin+ models.POST("", g.Admin(), handler.CreateModel) // 获取模型列表 — Viewer+ From 4898a8e66107b01e67fd5d2c3c3eb30e28235a01 Mon Sep 17 00:00:00 2001 From: lyingbug <[email> Date: Fri, 14 Aug 2026 00:16:48 +0000 Subject: [PATCH 5/9] feat(models): add a three-state thinking mode to the neutral options A boolean cannot reach the mode several vendors document, where the model decides whether to reason per request: Volcengine's `auto` and Anthropic's adaptive thinking were unreachable from the application even though both plugins encode them. Options.ThinkingMode takes 'on', 'off', or 'auto' and wins over the boolean, which stays for the callers that only need two states. Co-authored-by: lyingbug --- internal/models/llm/spi/message.go | 34 +++++++++++++++++++++++++----- 1 file changed, 29 insertions(+), 5 deletions(-) diff --git a/internal/models/llm/spi/message.go b/internal/models/llm/spi/message.go index 910e7b812e5..7be034de448 100644 --- a/internal/models/llm/spi/message.go +++ b/internal/models/llm/spi/message.go @@ -113,6 +113,12 @@ type Options struct { PresencePenalty float64 `json:"presence_penalty"` // Thinking is the neutral reasoning toggle; nil defers to the model. Thinking *bool `json:"thinking"` + // ThinkingMode is the three-state reasoning control: "on", "off", or + // "auto". It exists because a boolean cannot reach the mode several + // vendors document, where the model decides per request — Volcengine's + // `auto` and Anthropic's adaptive thinking. It takes precedence over + // Thinking, which remains for the callers that only need on/off. + ThinkingMode string `json:"thinking_mode,omitempty"` // ThinkingEffort and ThinkingBudget are the depth controls, empty or zero // when the caller has no opinion. A value a vendor does not accept is // reported in the plan rather than sent. @@ -155,11 +161,7 @@ func (o *Options) ParamValues() map[ParamID]Value { if maxTokens := o.EffectiveMaxTokens(); maxTokens > 0 { values[ParamMaxTokens] = IntValue(maxTokens) } - if o.Thinking != nil { - mode := ThinkingOff - if *o.Thinking { - mode = ThinkingOn - } + if mode := o.EffectiveThinkingMode(); mode != "" { values[ParamThinkingMode] = EnumValue(mode) } if o.ThinkingEffort != "" { @@ -177,6 +179,28 @@ func (o *Options) ParamValues() map[ParamID]Value { return values } +// EffectiveThinkingMode reports the requested reasoning mode, or empty when +// the caller had no opinion and the model's own default should stand. +// +// The explicit three-state field wins over the boolean, so a caller that +// upgrades to it does not have to also clear the older one. +func (o *Options) EffectiveThinkingMode() string { + if o == nil { + return "" + } + switch o.ThinkingMode { + case ThinkingOn, ThinkingOff, ThinkingAuto: + return o.ThinkingMode + } + if o.Thinking == nil { + return "" + } + if *o.Thinking { + return ThinkingOn + } + return ThinkingOff +} + // EffectiveMaxTokens reports the output ceiling, accepting either spelling the // caller may have used. // From 6ad1f94d9ae53e68bc5a67597af316c22f19a68f Mon Sep 17 00:00:00 2001 From: lyingbug <[email> Date: Fri, 14 Aug 2026 00:25:00 +0000 Subject: [PATCH 6/9] feat(plugin): add the domain-neutral plugin kernel MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Four registries had grown independently — web search providers, document parser engines, datasource connectors, and model vendors — each with its own answer to the same five questions: what is this, what configuration does it need, is that configuration valid, is it usable right now, and how does a UI render a form for it. Four answers diverged in practice. Some registries validated configuration and some did not; only the parser registry reported availability with a reason; the model registry's rules were duplicated in the frontend by hand. The kernel supplies one answer to each, and knows nothing about any domain: - Manifest: identity and capability tags, serializable so a plugin behind an RPC is described the same way as one compiled in. - Schema: typed configuration that validates a value and renders its form from one declaration, replacing bespoke structs and magic string maps. Secrets never travel back out. - Health: usability as a reported fact with a reason, not a promise. - Registry[T]: generic per domain for compile-time safety, with reversible registration and total validation at registration time. - Catalog: the non-generic view across domains, so one endpoint can answer what a deployment can do, and external plugins discovered over RPC join it. Exercised through an invented domain rather than a real one: if the tests needed to know what a model is, the kernel would not be neutral. Co-authored-by: lyingbug --- internal/plugin/plugin.go | 183 ++++++++++ internal/plugin/plugin_test.go | 391 +++++++++++++++++++++ internal/plugin/registry.go | 301 +++++++++++++++++ internal/plugin/schema.go | 601 +++++++++++++++++++++++++++++++++ 4 files changed, 1476 insertions(+) create mode 100644 internal/plugin/plugin.go create mode 100644 internal/plugin/plugin_test.go create mode 100644 internal/plugin/registry.go create mode 100644 internal/plugin/schema.go diff --git a/internal/plugin/plugin.go b/internal/plugin/plugin.go new file mode 100644 index 00000000000..635e940d5db --- /dev/null +++ b/internal/plugin/plugin.go @@ -0,0 +1,183 @@ +// Package plugin is WeKnora's plugin kernel: the identity, configuration, +// health, and registration machinery every extensible subsystem shares. +// +// The kernel is deliberately domain-neutral. It does not know what a model is, +// what a search provider is, or what a parser engine is. A domain defines its +// own capability interface and registers plugins that implement it; the kernel +// supplies only what all of them turned out to need independently. +// +// That "independently" is the justification. Before this package existed the +// codebase had four separate registries — web search providers, document +// parser engines, datasource connectors, and model vendors — each with its own +// answer to the same five questions: +// +// - What is this thing called, and what can it do? +// - What configuration does it need, and is a given configuration valid? +// - Is it usable right now, and if not, why not? +// - How do I build a working instance from a configuration? +// - How does a UI render a form for it? +// +// Four answers to five questions is four times the surface to keep correct, +// and in practice they diverged: some registries validated configuration, some +// did not; one reported availability with a reason, the rest left callers to +// guess; the model registry's rules were duplicated in the frontend by hand. +// One kernel makes those answers uniform, and makes a new pluggable subsystem +// a matter of declaring a capability interface rather than writing a fifth +// registry. +package plugin + +import ( + "context" + "fmt" + "strings" +) + +// Kind namespaces a capability domain. It is a dotted string so a domain can +// carve out sub-domains without a second registry: "llm.chat" and +// "llm.embedding" are distinct kinds sharing a prefix. +type Kind string + +// String reports the kind for logging and diagnostics. +func (k Kind) String() string { return string(k) } + +// Domain reports the portion of the kind before the first dot, so a UI can +// group "llm.chat" and "llm.embedding" without knowing either. +func (k Kind) Domain() string { + if idx := strings.IndexByte(string(k), '.'); idx >= 0 { + return string(k)[:idx] + } + return string(k) +} + +// Manifest is the identity and capability declaration every plugin publishes. +// +// It is serializable on purpose: the same structure describes an in-tree Go +// plugin and a plugin running behind an RPC, which is what lets a subsystem +// discover remote implementations without a second description format. The +// document parser already merges locally registered engines with engines +// discovered over RPC, and had to invent its own merge because there was no +// shared manifest to merge. +type Manifest struct { + // Kind is the capability domain this plugin serves. + Kind Kind `json:"kind"` + // ID is the stable identifier, unique within a kind. It is explicit + // configuration, never inferred: a plugin chosen by sniffing a URL is a + // plugin nobody can reliably select, and the model layer's URL detection + // is the cautionary example. + ID string `json:"id"` + // DisplayName is a human-readable name for diagnostics and logs. A UI + // prefers SummaryKey so it can translate. + DisplayName string `json:"display_name,omitempty"` + // SummaryKey is an i18n key describing the plugin. The kernel carries no + // display text: language belongs to the frontend. + SummaryKey string `json:"summary_key,omitempty"` + // DocURL links the upstream documentation this plugin implements, so a + // reviewer can check a claim rather than trust it. + DocURL string `json:"doc_url,omitempty"` + // Tags are capability markers the owning domain interprets — supported + // file extensions for a parser, supported features for a search provider. + // The kernel treats them as opaque so a domain can add one without + // changing this package. + Tags []string `json:"tags,omitempty"` + // Deprecated, when non-empty, explains what to use instead. A deprecated + // plugin still resolves, so existing configurations keep working. + Deprecated string `json:"deprecated,omitempty"` + // External marks a plugin whose implementation lives outside this process, + // discovered from a remote catalog. It has a manifest and a schema but no + // Go factory, and the owning domain decides how to drive it. + External bool `json:"external,omitempty"` +} + +// HasTag reports whether the manifest carries a capability tag. +func (m Manifest) HasTag(tag string) bool { + for _, t := range m.Tags { + if strings.EqualFold(t, tag) { + return true + } + } + return false +} + +// validate checks a manifest at registration time. +func (m Manifest) validate() error { + if strings.TrimSpace(string(m.Kind)) == "" { + return fmt.Errorf("plugin manifest: kind is required") + } + if strings.TrimSpace(m.ID) == "" { + return fmt.Errorf("plugin manifest %s: id is required", m.Kind) + } + if strings.ContainsAny(m.ID, " \t/\\") { + return fmt.Errorf("plugin manifest %s/%s: id must not contain spaces or slashes", m.Kind, m.ID) + } + return nil +} + +// HealthState is how usable a plugin is under a given configuration. +type HealthState string + +const ( + // Ready means the plugin is fully usable. + Ready HealthState = "ready" + // Degraded means it works with a documented limitation. A caller that + // needs the full capability must treat this as a refusal rather than + // rounding it up to ready. + Degraded HealthState = "degraded" + // Unavailable means it cannot serve requests as configured. + Unavailable HealthState = "unavailable" +) + +// Health is a plugin's own report on whether it can work right now. +// +// Reporting it as a fact, with a reason, is the difference between a settings +// page that explains why an engine is greyed out and one that silently omits +// it. The document parser registry already learned this and returns +// "(available, reason)"; making it part of the kernel means every subsystem +// gets the same honesty instead of rediscovering the need. +type Health struct { + // State is the usability verdict. + State HealthState `json:"state"` + // ReasonKey is an i18n key explaining a non-ready state. + ReasonKey string `json:"reason_key,omitempty"` + // Detail carries specifics for an operator: the endpoint that refused, + // the credential that is missing. It is not translated. + Detail string `json:"detail,omitempty"` +} + +// OK reports whether the plugin can serve requests at all. +func (h Health) OK() bool { return h.State == Ready || h.State == Degraded } + +// Healthy returns a ready verdict. +func Healthy() Health { return Health{State: Ready} } + +// Unhealthy returns an unavailable verdict with a reason. +func Unhealthy(reasonKey, detail string) Health { + return Health{State: Unavailable, ReasonKey: reasonKey, Detail: detail} +} + +// Limited returns a degraded verdict with a reason. +func Limited(reasonKey, detail string) Health { + return Health{State: Degraded, ReasonKey: reasonKey, Detail: detail} +} + +// Plugin is one implementation of a domain capability T. +// +// The three methods are the contract every subsystem needed anyway. Probe is +// separate from New because a settings page must be able to ask "would this +// work?" without building a live instance, and because the answer depends on +// the configuration rather than on the plugin alone. +type Plugin[T any] interface { + // Manifest declares identity and capabilities. + Manifest() Manifest + // Schema declares the configuration this plugin accepts. It drives + // validation and the settings form from one declaration, so a form cannot + // offer a field the plugin will reject. + Schema() Schema + // Probe reports whether the plugin can work with a configuration, without + // building an instance. Implementations that cannot check cheaply should + // report Ready and let New fail. + Probe(ctx context.Context, cfg Config) Health + // New builds a configured instance. It receives a configuration the kernel + // has already validated against Schema, so implementations do not repeat + // the domain checks. + New(ctx context.Context, cfg Config) (T, error) +} diff --git a/internal/plugin/plugin_test.go b/internal/plugin/plugin_test.go new file mode 100644 index 00000000000..8fc2ec3e875 --- /dev/null +++ b/internal/plugin/plugin_test.go @@ -0,0 +1,391 @@ +package plugin_test + +import ( + "context" + "errors" + "testing" + + "github.com/Tencent/WeKnora/internal/plugin" +) + +// The kernel is exercised through an invented domain rather than a real one. +// If these tests needed to know what a model or a search provider is, the +// kernel would not be domain-neutral and the whole premise would be wrong. + +// Greeter is a made-up domain capability. +type Greeter interface { + Greet(name string) string +} + +type greeter struct{ prefix string } + +func (g greeter) Greet(name string) string { return g.prefix + " " + name } + +const kindGreeter = plugin.Kind("demo.greeter") + +// politePlugin is a well-formed plugin: a required secret, a bounded number, +// a closed vocabulary, and a default. +type politePlugin struct { + health plugin.Health + newErr error +} + +func (politePlugin) Manifest() plugin.Manifest { + return plugin.Manifest{ + Kind: kindGreeter, + ID: "polite", + DisplayName: "Polite greeter", + Tags: []string{"formal"}, + } +} + +func (politePlugin) Schema() plugin.Schema { + return plugin.Schema{Fields: []plugin.Field{ + {ID: "token", Kind: plugin.KindString, Required: true, Secret: true}, + { + ID: "style", Kind: plugin.KindEnum, + Enum: []plugin.EnumOption{{Value: "hello"}, {Value: "greetings"}}, + Default: plugin.Ptr(plugin.EnumValue("hello")), + }, + { + ID: "volume", Kind: plugin.KindInt, + Min: plugin.Float64(1), Max: plugin.Float64(10), + Default: plugin.Ptr(plugin.IntValue(3)), + }, + {ID: "internal", Kind: plugin.KindBool, UI: plugin.FieldUI{Hidden: true}}, + }} +} + +func (p politePlugin) Probe(context.Context, plugin.Config) plugin.Health { + if p.health.State == "" { + return plugin.Healthy() + } + return p.health +} + +func (p politePlugin) New(_ context.Context, cfg plugin.Config) (Greeter, error) { + if p.newErr != nil { + return nil, p.newErr + } + return greeter{prefix: cfg.String("style")}, nil +} + +func newRegistry(t *testing.T) *plugin.Registry[Greeter] { + t.Helper() + return plugin.NewRegistryWithCatalog[Greeter](kindGreeter, plugin.NewCatalog()) +} + +func TestOpenValidatesBeforeBuilding(t *testing.T) { + reg := newRegistry(t) + if _, err := reg.Register(politePlugin{}); err != nil { + t.Fatalf("register: %v", err) + } + + got, err := reg.Open(context.Background(), "polite", map[string]any{ + "token": "secret-value", + "style": "greetings", + }) + if err != nil { + t.Fatalf("open: %v", err) + } + if greeting := got.Greet("world"); greeting != "greetings world" { + t.Errorf("greeting = %q", greeting) + } +} + +// A plugin must never see a configuration its schema would reject, so the +// failure has to happen before New is called. +func TestOpenRejectsAnInvalidConfiguration(t *testing.T) { + reg := newRegistry(t) + _, _ = reg.Register(politePlugin{newErr: errors.New("New must not be reached")}) + + cases := map[string]map[string]any{ + "missing required secret": {"style": "hello"}, + "value outside the vocabulary": { + "token": "t", "style": "howdy", + }, + "number below the minimum": { + "token": "t", "volume": 0, + }, + "unparseable number": { + "token": "t", "volume": "loud", + }, + } + for name, raw := range cases { + t.Run(name, func(t *testing.T) { + if _, err := reg.Open(context.Background(), "polite", raw); err == nil { + t.Fatal("expected the configuration to be rejected") + } + }) + } +} + +// Configuration arrives from JSON bodies and stored string maps, which both +// lose the distinction between 8080 and "8080". Being lenient about the +// representation is what lets stored configurations survive without migration. +func TestCoercionAcceptsTheRepresentationsConfigurationActuallyArrivesIn(t *testing.T) { + reg := newRegistry(t) + _, _ = reg.Register(politePlugin{}) + + for _, volume := range []any{5, "5", float64(5)} { + cfg, err := politePlugin{}.Schema().Validate(map[string]any{"token": "t", "volume": volume}) + if err != nil { + t.Fatalf("volume %#v: %v", volume, err) + } + if cfg.Int("volume") != 5 { + t.Errorf("volume %#v became %d", volume, cfg.Int("volume")) + } + } + _ = reg +} + +func TestDefaultsApplyWhenOmitted(t *testing.T) { + cfg, err := politePlugin{}.Schema().Validate(map[string]any{"token": "t"}) + if err != nil { + t.Fatalf("validate: %v", err) + } + if cfg.String("style") != "hello" { + t.Errorf("style default = %q", cfg.String("style")) + } + if cfg.Int("volume") != 3 { + t.Errorf("volume default = %d", cfg.Int("volume")) + } +} + +// A secret must not travel back out of the kernel, or an introspectable +// plugin becomes a credential leak. +func TestSecretsAreNeverEchoedBack(t *testing.T) { + cfg, err := politePlugin{}.Schema().Validate(map[string]any{"token": "super-secret"}) + if err != nil { + t.Fatalf("validate: %v", err) + } + redacted := cfg.Redacted() + if _, present := redacted["token"]; present { + t.Error("the secret survived redaction") + } + if redacted["style"] != "hello" { + t.Errorf("redaction dropped a non-secret field: %#v", redacted) + } + + for _, group := range (politePlugin{}).Schema().Form() { + for _, field := range group.Fields { + if field.ID == "token" { + if field.EffectiveWidget() != plugin.WidgetPassword { + t.Errorf("a secret should render masked, got %q", field.EffectiveWidget()) + } + if field.Default != nil { + t.Error("a secret must not carry a default into the form") + } + } + } + } +} + +// The form is what a settings page renders, so it must omit anything an +// operator cannot set. +func TestFormOmitsWhatAnOperatorCannotSet(t *testing.T) { + schema := plugin.Schema{Fields: []plugin.Field{ + {ID: "visible", Kind: plugin.KindString}, + {ID: "hidden", Kind: plugin.KindString, UI: plugin.FieldUI{Hidden: true}}, + { + ID: "pinned", Kind: plugin.KindString, + Support: plugin.SupportPinned, Pin: plugin.Ptr(plugin.StringValue("fixed")), + }, + {ID: "forbidden", Kind: plugin.KindString, Support: plugin.SupportForbidden}, + }} + + var rendered []string + for _, group := range schema.Form() { + for _, field := range group.Fields { + rendered = append(rendered, field.ID) + } + } + if len(rendered) != 1 || rendered[0] != "visible" { + t.Errorf("form rendered %v, want only the visible field", rendered) + } + + // A pinned field still reaches the plugin; it just is not a control. + cfg, err := schema.Validate(map[string]any{"pinned": "ignored"}) + if err != nil { + t.Fatalf("validate: %v", err) + } + if cfg.String("pinned") != "fixed" { + t.Errorf("pinned value = %q, want the pin", cfg.String("pinned")) + } + if cfg.Has("forbidden") { + t.Error("a forbidden field must not reach the plugin") + } +} + +// Malformed declarations are a programming error, and startup is the right +// place to find out. +func TestMalformedDeclarationsAreRejectedAtRegistration(t *testing.T) { + cases := map[string]plugin.Schema{ + "enum without options": {Fields: []plugin.Field{ + {ID: "mode", Kind: plugin.KindEnum}, + }}, + "pin outside the domain": {Fields: []plugin.Field{ + { + ID: "size", Kind: plugin.KindInt, Max: plugin.Float64(10), + Support: plugin.SupportPinned, Pin: plugin.Ptr(plugin.IntValue(99)), + }, + }}, + "required field carrying a default": {Fields: []plugin.Field{ + { + ID: "name", Kind: plugin.KindString, Required: true, + Default: plugin.Ptr(plugin.StringValue("x")), + }, + }}, + "duplicate field": {Fields: []plugin.Field{ + {ID: "a", Kind: plugin.KindString}, + {ID: "a", Kind: plugin.KindString}, + }}, + } + + for name, schema := range cases { + t.Run(name, func(t *testing.T) { + reg := newRegistry(t) + if _, err := reg.Register(brokenPlugin{schema: schema}); err == nil { + t.Fatal("expected registration to be rejected") + } + }) + } +} + +type brokenPlugin struct{ schema plugin.Schema } + +func (brokenPlugin) Manifest() plugin.Manifest { + return plugin.Manifest{Kind: kindGreeter, ID: "broken"} +} +func (b brokenPlugin) Schema() plugin.Schema { return b.schema } +func (brokenPlugin) Probe(context.Context, plugin.Config) plugin.Health { return plugin.Healthy() } +func (brokenPlugin) New(context.Context, plugin.Config) (Greeter, error) { + return nil, errors.New("unreachable") +} + +// Registration is reversible so tests do not leak into one another and a +// plugin can outlive less than the process. +func TestRegistrationIsReversible(t *testing.T) { + catalog := plugin.NewCatalog() + reg := plugin.NewRegistryWithCatalog[Greeter](kindGreeter, catalog) + + undo, err := reg.Register(politePlugin{}) + if err != nil { + t.Fatalf("register: %v", err) + } + if len(catalog.Entries(kindGreeter)) != 1 { + t.Fatal("the plugin should appear in the catalog") + } + + undo() + if _, ok := reg.Lookup("polite"); ok { + t.Error("the plugin should be gone from the registry") + } + if len(catalog.Entries(kindGreeter)) != 0 { + t.Error("the plugin should be gone from the catalog") + } +} + +func TestDuplicateIdsAreRejected(t *testing.T) { + reg := newRegistry(t) + if _, err := reg.Register(politePlugin{}); err != nil { + t.Fatalf("register: %v", err) + } + if _, err := reg.Register(politePlugin{}); err == nil { + t.Error("registering the same id twice should fail") + } +} + +// Probe answers "would this work?" without building anything, and an invalid +// configuration is an unavailable verdict rather than an error, because that +// is what a settings page needs to display. +func TestProbeReportsRatherThanThrows(t *testing.T) { + reg := newRegistry(t) + _, _ = reg.Register(politePlugin{health: plugin.Limited("demo.noNetwork", "offline")}) + + health := reg.Probe(context.Background(), "polite", map[string]any{"token": "t"}) + if health.State != plugin.Degraded || !health.OK() { + t.Errorf("health = %+v, want a degraded but usable verdict", health) + } + + invalid := reg.Probe(context.Background(), "polite", map[string]any{}) + if invalid.State != plugin.Unavailable { + t.Errorf("an invalid configuration should probe unavailable, got %+v", invalid) + } + + missing := reg.Probe(context.Background(), "nope", nil) + if missing.State != plugin.Unavailable { + t.Errorf("an unregistered id should probe unavailable, got %+v", missing) + } +} + +// One catalog across domains is what lets a single endpoint answer "what can +// this deployment do?" without a listing endpoint per subsystem. +func TestCatalogSpansDomains(t *testing.T) { + catalog := plugin.NewCatalog() + greeters := plugin.NewRegistryWithCatalog[Greeter](kindGreeter, catalog) + if _, err := greeters.Register(politePlugin{}); err != nil { + t.Fatalf("register: %v", err) + } + + const kindCounter = plugin.Kind("demo.counter") + counters := plugin.NewRegistryWithCatalog[Greeter](kindCounter, catalog) + if _, err := counters.Register(otherDomainPlugin{}); err != nil { + t.Fatalf("register: %v", err) + } + + if got := catalog.Kinds(); len(got) != 2 { + t.Errorf("catalog kinds = %v, want both domains", got) + } + if got := len(catalog.Entries(kindGreeter)); got != 1 { + t.Errorf("filtering by kind returned %d entries", got) + } + if got := len(catalog.Entries("")); got != 2 { + t.Errorf("unfiltered catalog returned %d entries", got) + } +} + +type otherDomainPlugin struct{} + +func (otherDomainPlugin) Manifest() plugin.Manifest { + return plugin.Manifest{Kind: plugin.Kind("demo.counter"), ID: "tally"} +} +func (otherDomainPlugin) Schema() plugin.Schema { return plugin.Schema{} } +func (otherDomainPlugin) Probe(context.Context, plugin.Config) plugin.Health { + return plugin.Healthy() +} +func (otherDomainPlugin) New(context.Context, plugin.Config) (Greeter, error) { + return greeter{prefix: "count"}, nil +} + +// A plugin registered under the wrong registry is a wiring mistake worth +// catching immediately. +func TestKindMismatchIsRejected(t *testing.T) { + reg := newRegistry(t) + if _, err := reg.Register(otherDomainPlugin{}); err == nil { + t.Error("registering a counter into the greeter registry should fail") + } +} + +// An external plugin has a manifest and a schema but no in-process +// implementation, which is how a domain lists engines discovered over RPC in +// the same catalog as local ones. +func TestExternalPluginsJoinTheCatalog(t *testing.T) { + catalog := plugin.NewCatalog() + undo, err := catalog.PublishExternal( + plugin.Manifest{Kind: kindGreeter, ID: "remote-greeter"}, + plugin.Schema{Fields: []plugin.Field{{ID: "endpoint", Kind: plugin.KindString, Required: true}}}, + ) + if err != nil { + t.Fatalf("publish: %v", err) + } + + entries := catalog.Entries(kindGreeter) + if len(entries) != 1 || !entries[0].Manifest.External { + t.Fatalf("entries = %+v, want one external entry", entries) + } + undo() + if len(catalog.Entries(kindGreeter)) != 0 { + t.Error("withdrawing an external plugin should remove it") + } +} diff --git a/internal/plugin/registry.go b/internal/plugin/registry.go new file mode 100644 index 00000000000..58a870b3927 --- /dev/null +++ b/internal/plugin/registry.go @@ -0,0 +1,301 @@ +package plugin + +import ( + "context" + "fmt" + "sort" + "strings" + "sync" +) + +// Registration is a reversible handle. Registering returns the function that +// undoes it, which keeps tests from leaking state into one another and leaves +// room for plugins whose lifetime is shorter than the process — a tenant +// enabling an integration, a remote catalog being refreshed. +type Registration func() + +// Registry holds the plugins implementing one domain capability. +// +// It is generic so each domain keeps compile-time type safety: a +// Registry[WebSearchProvider] cannot hand back a chat model. The shared +// Catalog below gives the non-generic view a settings page needs, so +// genericity costs nothing at the UI boundary. +type Registry[T any] struct { + kind Kind + mu sync.RWMutex + entries []registryEntry[T] + nextSeq int + catalog *Catalog +} + +type registryEntry[T any] struct { + plugin Plugin[T] + seq int +} + +// NewRegistry returns a registry for one capability kind, publishing its +// manifests into the shared catalog. +func NewRegistry[T any](kind Kind) *Registry[T] { + return &Registry[T]{kind: kind, catalog: DefaultCatalog} +} + +// NewRegistryWithCatalog returns a registry publishing into a specific +// catalog, for tests that must not touch the process-wide one. +func NewRegistryWithCatalog[T any](kind Kind, catalog *Catalog) *Registry[T] { + return &Registry[T]{kind: kind, catalog: catalog} +} + +// Kind reports the capability this registry serves. +func (r *Registry[T]) Kind() Kind { return r.kind } + +// Register validates and adds a plugin, returning the function that removes it. +// +// Validation is total: a plugin whose manifest or schema is malformed is +// rejected rather than half-registered, so a mistake surfaces at startup +// instead of as a broken settings page in production. +func (r *Registry[T]) Register(p Plugin[T]) (Registration, error) { + manifest := p.Manifest() + if manifest.Kind == "" { + manifest.Kind = r.kind + } + if manifest.Kind != r.kind { + return nil, fmt.Errorf("plugin %s declares kind %s but was registered under %s", + manifest.ID, manifest.Kind, r.kind) + } + if err := manifest.validate(); err != nil { + return nil, err + } + where := fmt.Sprintf("plugin %s/%s", manifest.Kind, manifest.ID) + if err := p.Schema().validate(where); err != nil { + return nil, err + } + + r.mu.Lock() + for _, existing := range r.entries { + if strings.EqualFold(existing.plugin.Manifest().ID, manifest.ID) { + r.mu.Unlock() + return nil, fmt.Errorf("%s is already registered", where) + } + } + seq := r.nextSeq + r.nextSeq++ + r.entries = append(r.entries, registryEntry[T]{plugin: p, seq: seq}) + r.mu.Unlock() + + undoCatalog := r.catalog.publish(manifest, p.Schema()) + return func() { + r.remove(seq) + undoCatalog() + }, nil +} + +// MustRegister adds a plugin and panics if it is malformed. Plugin packages +// call it from init, where a declaration is a constant of the program and a +// mistake in one is a programming error rather than a runtime condition. +func (r *Registry[T]) MustRegister(p Plugin[T]) Registration { + undo, err := r.Register(p) + if err != nil { + panic(fmt.Sprintf("plugin: %v", err)) + } + return undo +} + +func (r *Registry[T]) remove(seq int) { + r.mu.Lock() + defer r.mu.Unlock() + for i, e := range r.entries { + if e.seq == seq { + r.entries = append(r.entries[:i], r.entries[i+1:]...) + return + } + } +} + +// Lookup reports the plugin with an id. +func (r *Registry[T]) Lookup(id string) (Plugin[T], bool) { + r.mu.RLock() + defer r.mu.RUnlock() + for _, e := range r.entries { + if strings.EqualFold(e.plugin.Manifest().ID, id) { + return e.plugin, true + } + } + return nil, false +} + +// List reports every registered plugin in registration order. +func (r *Registry[T]) List() []Plugin[T] { + r.mu.RLock() + defer r.mu.RUnlock() + entries := make([]registryEntry[T], len(r.entries)) + copy(entries, r.entries) + sort.SliceStable(entries, func(i, j int) bool { return entries[i].seq < entries[j].seq }) + + out := make([]Plugin[T], 0, len(entries)) + for _, e := range entries { + out = append(out, e.plugin) + } + return out +} + +// Select reports the plugins whose manifest satisfies a predicate, for domains +// that route by capability rather than by id — a parser engine chosen by file +// extension, for instance. +func (r *Registry[T]) Select(match func(Manifest) bool) []Plugin[T] { + var out []Plugin[T] + for _, p := range r.List() { + if match(p.Manifest()) { + out = append(out, p) + } + } + return out +} + +// Open resolves a plugin by id, validates the configuration against its +// schema, and builds an instance. +// +// It is the one path a consumer needs, and going through it is what guarantees +// that no plugin ever receives a configuration its schema would reject. A +// consumer that reaches for Lookup and builds by hand has opted out of that +// guarantee, which is why this is the documented entry point. +func (r *Registry[T]) Open(ctx context.Context, id string, raw map[string]any) (T, error) { + var zero T + p, ok := r.Lookup(id) + if !ok { + return zero, fmt.Errorf("no %s plugin registered with id %q", r.kind, id) + } + cfg, err := p.Schema().Validate(raw) + if err != nil { + return zero, fmt.Errorf("%s/%s: %w", r.kind, id, err) + } + instance, err := p.New(ctx, cfg) + if err != nil { + return zero, fmt.Errorf("%s/%s: %w", r.kind, id, err) + } + return instance, nil +} + +// Probe reports whether a plugin would work with a configuration, without +// building an instance. An invalid configuration is itself an unavailable +// verdict rather than an error, because that is what a settings page wants to +// display next to the field the operator got wrong. +func (r *Registry[T]) Probe(ctx context.Context, id string, raw map[string]any) Health { + p, ok := r.Lookup(id) + if !ok { + return Unhealthy("plugin.notRegistered", fmt.Sprintf("no %s plugin with id %q", r.kind, id)) + } + cfg, err := p.Schema().Validate(raw) + if err != nil { + return Unhealthy("plugin.invalidConfig", err.Error()) + } + return p.Probe(ctx, cfg) +} + +// CatalogEntry is the non-generic view of one registered plugin: everything a +// settings page needs, and nothing that requires knowing the capability type. +type CatalogEntry struct { + Manifest Manifest `json:"manifest"` + Groups []Group `json:"groups"` +} + +// Catalog is the process-wide index of registered plugins across every domain. +// +// It exists so one endpoint can answer "what can this deployment do?" without +// a per-domain listing endpoint, and so a new pluggable subsystem appears in +// that answer by registering rather than by also touching the API layer. +type Catalog struct { + mu sync.RWMutex + entries []catalogEntry + nextSeq int +} + +type catalogEntry struct { + entry CatalogEntry + seq int +} + +// NewCatalog returns an empty catalog. +func NewCatalog() *Catalog { return &Catalog{} } + +// DefaultCatalog is the catalog registries publish into by default. +var DefaultCatalog = NewCatalog() + +func (c *Catalog) publish(manifest Manifest, schema Schema) Registration { + c.mu.Lock() + defer c.mu.Unlock() + seq := c.nextSeq + c.nextSeq++ + c.entries = append(c.entries, catalogEntry{ + entry: CatalogEntry{Manifest: manifest, Groups: schema.Form()}, + seq: seq, + }) + return func() { c.withdraw(seq) } +} + +func (c *Catalog) withdraw(seq int) { + c.mu.Lock() + defer c.mu.Unlock() + for i, e := range c.entries { + if e.seq == seq { + c.entries = append(c.entries[:i], c.entries[i+1:]...) + return + } + } +} + +// PublishExternal records a plugin discovered from a remote catalog, which has +// a manifest and a schema but no in-process implementation. +// +// The document parser already merges locally registered engines with engines +// discovered over RPC and had to hand-roll that merge; expressing an external +// plugin in the same catalog is what makes such a merge a kernel feature +// rather than a per-domain one. +func (c *Catalog) PublishExternal(manifest Manifest, schema Schema) (Registration, error) { + manifest.External = true + if err := manifest.validate(); err != nil { + return nil, err + } + if err := schema.validate(fmt.Sprintf("external plugin %s/%s", manifest.Kind, manifest.ID)); err != nil { + return nil, err + } + return c.publish(manifest, schema), nil +} + +// Entries reports the catalog in registration order, optionally filtered to +// one kind. An empty kind reports everything. +func (c *Catalog) Entries(kind Kind) []CatalogEntry { + c.mu.RLock() + defer c.mu.RUnlock() + + filtered := make([]catalogEntry, 0, len(c.entries)) + for _, e := range c.entries { + if kind == "" || e.entry.Manifest.Kind == kind { + filtered = append(filtered, e) + } + } + sort.SliceStable(filtered, func(i, j int) bool { return filtered[i].seq < filtered[j].seq }) + + out := make([]CatalogEntry, 0, len(filtered)) + for _, e := range filtered { + out = append(out, e.entry) + } + return out +} + +// Kinds reports the capability kinds that have at least one plugin. +func (c *Catalog) Kinds() []Kind { + c.mu.RLock() + defer c.mu.RUnlock() + + seen := map[Kind]struct{}{} + var out []Kind + for _, e := range c.entries { + if _, dup := seen[e.entry.Manifest.Kind]; dup { + continue + } + seen[e.entry.Manifest.Kind] = struct{}{} + out = append(out, e.entry.Manifest.Kind) + } + sort.Slice(out, func(i, j int) bool { return out[i] < out[j] }) + return out +} diff --git a/internal/plugin/schema.go b/internal/plugin/schema.go new file mode 100644 index 00000000000..a9d0c88f908 --- /dev/null +++ b/internal/plugin/schema.go @@ -0,0 +1,601 @@ +package plugin + +import ( + "fmt" + "math" + "sort" + "strconv" + "strings" +) + +// Configuration is the kernel's second job, and the one the previous +// registries handled worst. Web search providers took a typed Go struct that +// no UI could introspect; document parser engines took a +// map[string]string of magic keys checked by hand; the model layer took both, +// plus an extra_config map whose keys were documented only in a Vue file. +// +// A schema replaces all three: one declaration that validates a configuration +// and renders its form, so a settings page cannot offer a field the plugin +// will reject, and a plugin cannot read a key the form never collected. + +// ValueKind is the domain of a configuration value. +type ValueKind string + +const ( + // KindString is free text. + KindString ValueKind = "string" + // KindBool is a flag. + KindBool ValueKind = "bool" + // KindInt is a whole number. + KindInt ValueKind = "int" + // KindFloat is a fractional number. + KindFloat ValueKind = "float" + // KindEnum is a closed vocabulary. The vocabulary belongs to the plugin, + // not to a shared normalization: two vendors' effort ladders genuinely + // differ, and pretending otherwise sends values one of them rejects. + KindEnum ValueKind = "enum" +) + +// Value is a configuration value tagged with its kind. It is a concrete struct +// rather than `any` because validators, encoders, and the form renderer all +// need the shape without a type switch at every access. +type Value struct { + Kind ValueKind `json:"kind"` + Str string `json:"str,omitempty"` + Bool bool `json:"bool,omitempty"` + Num float64 `json:"num,omitempty"` +} + +// StringValue returns a text value. +func StringValue(v string) Value { return Value{Kind: KindString, Str: v} } + +// BoolValue returns a flag value. +func BoolValue(v bool) Value { return Value{Kind: KindBool, Bool: v} } + +// IntValue returns a whole-number value. +func IntValue(v int) Value { return Value{Kind: KindInt, Num: float64(v)} } + +// FloatValue returns a fractional value. +func FloatValue(v float64) Value { return Value{Kind: KindFloat, Num: v} } + +// EnumValue returns a value drawn from a plugin's vocabulary. +func EnumValue(v string) Value { return Value{Kind: KindEnum, Str: v} } + +// Text reports the value as a string. The second result is false when the +// value is neither text nor an enum member. +func (v Value) Text() (string, bool) { + if v.Kind != KindString && v.Kind != KindEnum { + return "", false + } + return v.Str, true +} + +// Int reports the value as a whole number, truncating a float toward zero. +func (v Value) Int() (int, bool) { + if v.Kind != KindInt && v.Kind != KindFloat { + return 0, false + } + return int(math.Trunc(v.Num)), true +} + +// Float reports the value as a float. +func (v Value) Float() (float64, bool) { + if v.Kind != KindInt && v.Kind != KindFloat { + return 0, false + } + return v.Num, true +} + +// Flag reports the value as a boolean. +func (v Value) Flag() (bool, bool) { + if v.Kind != KindBool { + return false, false + } + return v.Bool, true +} + +// JSON reports the value as it should appear in a JSON document. +func (v Value) JSON() any { + switch v.Kind { + case KindString, KindEnum: + return v.Str + case KindBool: + return v.Bool + case KindInt: + return int(math.Trunc(v.Num)) + case KindFloat: + return v.Num + default: + return nil + } +} + +// String renders the value for logs and diagnostics. +func (v Value) String() string { + switch v.Kind { + case KindString, KindEnum: + return v.Str + case KindBool: + return strconv.FormatBool(v.Bool) + case KindInt: + return strconv.Itoa(int(math.Trunc(v.Num))) + case KindFloat: + return strconv.FormatFloat(v.Num, 'g', -1, 64) + default: + return "" + } +} + +// Coerce converts a loosely typed value — as it arrives from JSON, a form, or +// a stored string map — into a value of the requested kind. +// +// It is lenient about representation and strict about domain: "8080" becomes +// an int because HTTP and YAML both lose the distinction, while "maybe" does +// not become a bool. Being lenient here is what lets stored configurations +// survive without a migration. +func Coerce(kind ValueKind, raw any) (Value, error) { + switch typed := raw.(type) { + case Value: + if typed.Kind == kind { + return typed, nil + } + return Coerce(kind, typed.JSON()) + case nil: + return Value{}, fmt.Errorf("value is missing") + case string: + return coerceString(kind, typed) + case bool: + if kind != KindBool { + return Value{}, fmt.Errorf("expected %s, got a boolean", kind) + } + return BoolValue(typed), nil + case float64: + return coerceNumber(kind, typed) + case float32: + return coerceNumber(kind, float64(typed)) + case int: + return coerceNumber(kind, float64(typed)) + case int64: + return coerceNumber(kind, float64(typed)) + default: + return Value{}, fmt.Errorf("cannot read %T as %s", raw, kind) + } +} + +func coerceString(kind ValueKind, raw string) (Value, error) { + trimmed := strings.TrimSpace(raw) + switch kind { + case KindString: + return StringValue(raw), nil + case KindEnum: + if trimmed == "" { + return Value{}, fmt.Errorf("value is empty") + } + return EnumValue(trimmed), nil + case KindBool: + b, err := strconv.ParseBool(trimmed) + if err != nil { + return Value{}, fmt.Errorf("%q is not a boolean", raw) + } + return BoolValue(b), nil + case KindInt: + n, err := strconv.Atoi(trimmed) + if err != nil { + return Value{}, fmt.Errorf("%q is not a whole number", raw) + } + return IntValue(n), nil + case KindFloat: + f, err := strconv.ParseFloat(trimmed, 64) + if err != nil { + return Value{}, fmt.Errorf("%q is not a number", raw) + } + return FloatValue(f), nil + default: + return Value{}, fmt.Errorf("unknown value kind %q", kind) + } +} + +func coerceNumber(kind ValueKind, raw float64) (Value, error) { + switch kind { + case KindInt: + return IntValue(int(math.Trunc(raw))), nil + case KindFloat: + return FloatValue(raw), nil + case KindString: + return StringValue(strconv.FormatFloat(raw, 'g', -1, 64)), nil + default: + return Value{}, fmt.Errorf("expected %s, got a number", kind) + } +} + +// Widget is the control a form should render. It is inferred from a field's +// shape, and overridable where the inference is wrong. +type Widget string + +const ( + // WidgetText renders a single-line text input. + WidgetText Widget = "text" + // WidgetPassword renders a masked input for a secret. + WidgetPassword Widget = "password" + // WidgetSwitch renders a toggle. + WidgetSwitch Widget = "switch" + // WidgetSelect renders a closed vocabulary. + WidgetSelect Widget = "select" + // WidgetNumber renders a numeric input. + WidgetNumber Widget = "number" + // WidgetSlider renders a bounded numeric range. + WidgetSlider Widget = "slider" +) + +// EnumOption is one member of a plugin's vocabulary, carrying the stored value +// with the i18n key a form uses to label it. +type EnumOption struct { + Value string `json:"value"` + LabelKey string `json:"label_key,omitempty"` + HelpKey string `json:"help_key,omitempty"` +} + +// FieldUI is how a form presents a field. It travels with the field rather +// than in a parallel table so a plugin that adds configuration gets the form +// control for free, and a field that disappears cannot linger in the UI. +type FieldUI struct { + // Hidden keeps a field out of the form while still accepting it in a + // stored configuration. + Hidden bool `json:"hidden,omitempty"` + // Group buckets the field into a form section. + Group string `json:"group,omitempty"` + // LabelKey, HelpKey, and PlaceholderKey are i18n keys. + LabelKey string `json:"label_key,omitempty"` + HelpKey string `json:"help_key,omitempty"` + PlaceholderKey string `json:"placeholder_key,omitempty"` + // Widget overrides the control inferred from the field's shape. + Widget Widget `json:"widget,omitempty"` + // Order sorts fields within a group. + Order int `json:"order,omitempty"` +} + +// Support is how a plugin dispositions a field. +type Support string + +const ( + // SupportUser means the operator sets it. + SupportUser Support = "user" + // SupportPinned means the plugin always uses Pin and the operator cannot + // override it. + SupportPinned Support = "pinned" + // SupportForbidden means the field must never take effect. Declaring it — + // rather than omitting it — is what lets a caller learn that a value it + // supplied was dropped, and why. + SupportForbidden Support = "forbidden" +) + +// Field declares one configuration input. +type Field struct { + // ID is the key under which the value is stored and submitted. + ID string `json:"id"` + // Kind is the value domain. + Kind ValueKind `json:"kind"` + // Support dispositions the field; the zero value is SupportUser. + Support Support `json:"support,omitempty"` + // Required rejects a configuration that omits the field and supplies no + // default. + Required bool `json:"required,omitempty"` + // Secret marks a credential. The kernel keeps secrets out of rendered + // forms and out of any configuration it echoes back, so a plugin cannot + // accidentally leak one by being introspectable. + Secret bool `json:"secret,omitempty"` + // Enum is the vocabulary for an enum field. + Enum []EnumOption `json:"enum,omitempty"` + // Min and Max bound a numeric field. + Min *float64 `json:"min,omitempty"` + Max *float64 `json:"max,omitempty"` + // Default is used when a configuration omits the field. + Default *Value `json:"default,omitempty"` + // Pin is the forced value for a pinned field. + Pin *Value `json:"pin,omitempty"` + // UI is the form presentation. + UI FieldUI `json:"ui"` + // DocURL links the documentation behind this field. + DocURL string `json:"doc_url,omitempty"` +} + +// EffectiveSupport reports the disposition, treating the zero value as user. +func (f Field) EffectiveSupport() Support { + if f.Support == "" { + return SupportUser + } + return f.Support +} + +// EffectiveWidget reports the control to render. +func (f Field) EffectiveWidget() Widget { + if f.UI.Widget != "" { + return f.UI.Widget + } + switch { + case f.Secret: + return WidgetPassword + case f.Kind == KindBool: + return WidgetSwitch + case f.Kind == KindEnum: + return WidgetSelect + case f.Kind == KindFloat && f.Min != nil && f.Max != nil: + return WidgetSlider + case f.Kind == KindInt || f.Kind == KindFloat: + return WidgetNumber + default: + return WidgetText + } +} + +// Accepts reports whether a value falls inside the field's declared domain. +func (f Field) Accepts(v Value) bool { + if v.Kind != f.Kind { + return false + } + switch f.Kind { + case KindEnum: + for _, opt := range f.Enum { + if opt.Value == v.Str { + return true + } + } + return false + case KindInt, KindFloat: + if f.Min != nil && v.Num < *f.Min { + return false + } + if f.Max != nil && v.Num > *f.Max { + return false + } + return true + default: + return true + } +} + +// validate checks a field declaration at registration time. These are the +// mistakes plugin authors actually make, and catching them at startup beats +// discovering them through a malformed form or a rejected request. +func (f Field) validate(where string) error { + if strings.TrimSpace(f.ID) == "" { + return fmt.Errorf("%s: field id is required", where) + } + if f.Kind == "" { + return fmt.Errorf("%s field %s: kind is required", where, f.ID) + } + if f.Kind == KindEnum && len(f.Enum) == 0 { + return fmt.Errorf("%s field %s: an enum needs at least one option", where, f.ID) + } + if f.Kind != KindEnum && len(f.Enum) > 0 { + return fmt.Errorf("%s field %s: only an enum may declare options", where, f.ID) + } + if f.EffectiveSupport() == SupportPinned && f.Pin == nil { + return fmt.Errorf("%s field %s: a pinned field needs a pin value", where, f.ID) + } + if f.Pin != nil && !f.Accepts(*f.Pin) { + return fmt.Errorf("%s field %s: pin %s is outside the declared domain", where, f.ID, f.Pin) + } + if f.Default != nil && !f.Accepts(*f.Default) { + return fmt.Errorf("%s field %s: default %s is outside the declared domain", where, f.ID, f.Default) + } + if f.Min != nil && f.Max != nil && *f.Min > *f.Max { + return fmt.Errorf("%s field %s: min %v exceeds max %v", where, f.ID, *f.Min, *f.Max) + } + if f.Required && f.Default != nil { + return fmt.Errorf("%s field %s: a required field must not also carry a default", where, f.ID) + } + return nil +} + +// Schema is a plugin's configuration contract. +type Schema struct { + // Fields are the inputs, in the order a form should present them. + Fields []Field +} + +// Field reports a declared field by id. +func (s Schema) Field(id string) (Field, bool) { + for _, f := range s.Fields { + if f.ID == id { + return f, true + } + } + return Field{}, false +} + +// Validate turns a loosely typed configuration into a validated one. +// +// Precedence is pin, then the supplied value, then the default. A value the +// domain rejects is an error rather than a silent fallback: a configuration +// that quietly means something other than what was written is worse than one +// that refuses to load. +func (s Schema) Validate(raw map[string]any) (Config, error) { + values := make(map[string]Value, len(s.Fields)) + var problems []string + + declared := make(map[string]struct{}, len(s.Fields)) + for _, f := range s.Fields { + declared[f.ID] = struct{}{} + + if f.EffectiveSupport() == SupportPinned { + values[f.ID] = *f.Pin + continue + } + if f.EffectiveSupport() == SupportForbidden { + continue + } + + supplied, ok := raw[f.ID] + if !ok || supplied == nil || supplied == "" { + if f.Default != nil { + values[f.ID] = *f.Default + } else if f.Required { + problems = append(problems, fmt.Sprintf("%s is required", f.ID)) + } + continue + } + + value, err := Coerce(f.Kind, supplied) + if err != nil { + problems = append(problems, fmt.Sprintf("%s: %v", f.ID, err)) + continue + } + if !f.Accepts(value) { + problems = append(problems, fmt.Sprintf("%s: %s is not an accepted value", f.ID, value)) + continue + } + values[f.ID] = value + } + + if len(problems) > 0 { + return Config{}, fmt.Errorf("invalid configuration: %s", strings.Join(problems, "; ")) + } + return Config{values: values, schema: s}, nil +} + +// Group is one section of a rendered form. +type Group struct { + Key string `json:"key"` + Fields []Field `json:"fields"` +} + +// Form renders the schema for a UI, omitting what an operator cannot set: a +// hidden field, and anything pinned or forbidden. Secrets are rendered as +// masked inputs but never carry a default, so a stored credential is not +// echoed back into a form. +func (s Schema) Form() []Group { + byGroup := map[string][]Field{} + var order []string + + for _, f := range s.Fields { + if f.UI.Hidden || f.EffectiveSupport() != SupportUser { + continue + } + rendered := f + if rendered.Secret { + rendered.Default = nil + } + key := f.UI.Group + if key == "" { + key = "general" + } + if _, seen := byGroup[key]; !seen { + order = append(order, key) + } + byGroup[key] = append(byGroup[key], rendered) + } + + groups := make([]Group, 0, len(order)) + for _, key := range order { + fields := byGroup[key] + sort.SliceStable(fields, func(i, j int) bool { + return fields[i].UI.Order < fields[j].UI.Order + }) + groups = append(groups, Group{Key: key, Fields: fields}) + } + return groups +} + +// validate checks the whole schema at registration time. +func (s Schema) validate(where string) error { + seen := make(map[string]struct{}, len(s.Fields)) + for _, f := range s.Fields { + if _, dup := seen[f.ID]; dup { + return fmt.Errorf("%s: duplicate field %s", where, f.ID) + } + seen[f.ID] = struct{}{} + if err := f.validate(where); err != nil { + return err + } + } + return nil +} + +// Config is a configuration that has passed its schema. +// +// A plugin receives one instead of a raw map, so it never repeats validation +// and never reads a key the schema did not declare. The accessors report +// zero values for anything absent, because absence has already been decided +// to be acceptable by the time a Config exists. +type Config struct { + values map[string]Value + schema Schema +} + +// NewConfig builds a validated configuration directly, for tests and for +// callers assembling one in code rather than from stored settings. +func NewConfig(schema Schema, raw map[string]any) (Config, error) { + return schema.Validate(raw) +} + +// Has reports whether a value is present. +func (c Config) Has(id string) bool { + _, ok := c.values[id] + return ok +} + +// Value reports a raw value and whether it is present. +func (c Config) Value(id string) (Value, bool) { + v, ok := c.values[id] + return v, ok +} + +// String reports a text or enum value, or "" when absent. +func (c Config) String(id string) string { + v, ok := c.values[id] + if !ok { + return "" + } + text, _ := v.Text() + return text +} + +// Int reports a numeric value, or 0 when absent. +func (c Config) Int(id string) int { + v, ok := c.values[id] + if !ok { + return 0 + } + n, _ := v.Int() + return n +} + +// Float reports a numeric value, or 0 when absent. +func (c Config) Float(id string) float64 { + v, ok := c.values[id] + if !ok { + return 0 + } + f, _ := v.Float() + return f +} + +// Bool reports a flag, or false when absent. +func (c Config) Bool(id string) bool { + v, ok := c.values[id] + if !ok { + return false + } + b, _ := v.Flag() + return b +} + +// Redacted reports the configuration with secrets removed, for logging and for +// any response that echoes a configuration back. +func (c Config) Redacted() map[string]any { + out := make(map[string]any, len(c.values)) + for id, v := range c.values { + if field, ok := c.schema.Field(id); ok && field.Secret { + continue + } + out[id] = v.JSON() + } + return out +} + +// Float64 returns a pointer to f, for populating bounds inline. +func Float64(f float64) *float64 { return &f } + +// Ptr returns a pointer to v, for populating defaults and pins inline. +func Ptr(v Value) *Value { return &v } From f4241c90a65b2b294d0fd5c1104f9524b16bd23f Mon Sep 17 00:00:00 2001 From: lyingbug <[email> Date: Fri, 14 Aug 2026 00:31:56 +0000 Subject: [PATCH 7/9] refactor(websearch): move web search onto the plugin kernel The first domain migrated, and the proof the kernel is not model-specific. What the domain had: a registry mapping an id to a factory, credential checks hand-written inside each constructor, provider options as undocumented keys in an ExtraConfig map, and a wiring list in the dependency container that had to be edited to add a provider. What it has now: - Providers self-register from init, so the container's wiring list and the old registry are deleted rather than replaced. - Each declares its inputs, so the kernel validates before a constructor runs and every provider reports a missing credential the same way. - The options three providers parsed out of ExtraConfig by hand are declared fields with their documented vocabularies, so a wrong Zhipu search engine is refused locally instead of sent upstream. ExtraConfig is now a transport detail between the adapter and an untouched constructor. - Probe reports why a configuration will not work, reusing the provider checks a schema cannot express, and the answers distinguish the failures. - API keys render masked and never travel back through the catalog. Provider implementations are untouched; only registration changed. Co-authored-by: lyingbug --- internal/application/service/web_search.go | 13 +- internal/container/container.go | 18 -- internal/handler/web_search_provider.go | 10 +- internal/infrastructure/web_search/plugin.go | 284 ++++++++++++++++++ .../infrastructure/web_search/plugin_test.go | 158 ++++++++++ internal/infrastructure/web_search/plugins.go | 101 +++++++ .../infrastructure/web_search/registry.go | 45 --- 7 files changed, 553 insertions(+), 76 deletions(-) create mode 100644 internal/infrastructure/web_search/plugin.go create mode 100644 internal/infrastructure/web_search/plugin_test.go create mode 100644 internal/infrastructure/web_search/plugins.go delete mode 100644 internal/infrastructure/web_search/registry.go diff --git a/internal/application/service/web_search.go b/internal/application/service/web_search.go index 1a8047747ac..99e82f26391 100644 --- a/internal/application/service/web_search.go +++ b/internal/application/service/web_search.go @@ -17,18 +17,18 @@ import ( // WebSearchService provides web search functionality. // It resolves provider configurations from the database and creates provider -// instances on-demand via the infrastructure registry. +// instances on-demand from the web-search plugin registry. type WebSearchService struct { - registry *infra_web_search.Registry providerRepo interfaces.WebSearchProviderRepository timeout int } // NewWebSearchService creates a new web search service. -// The registry holds provider type factories; the providerRepo loads tenant-specific configurations. +// Providers self-register with the plugin kernel, which validates a stored +// configuration against the provider's schema before building it; the +// providerRepo loads those tenant-specific configurations. func NewWebSearchService( cfg *config.Config, - registry *infra_web_search.Registry, providerRepo interfaces.WebSearchProviderRepository, ) (interfaces.WebSearchService, error) { timeout := 10 // default timeout in seconds @@ -37,7 +37,6 @@ func NewWebSearchService( } return &WebSearchService{ - registry: registry, providerRepo: providerRepo, timeout: timeout, }, nil @@ -106,7 +105,7 @@ func (s *WebSearchService) resolveProvider( } params := mergeProxyFromWebSearchConfig(entity.Parameters, cfg) - provider, err := s.registry.CreateProvider(string(entity.Provider), params) + provider, err := infra_web_search.Open(ctx, string(entity.Provider), params) if err != nil { return nil, fmt.Errorf("failed to create provider %s (%s): %w", entity.Name, entity.Provider, err) } @@ -119,7 +118,7 @@ func (s *WebSearchService) resolveProvider( params := mergeProxyFromWebSearchConfig(types.WebSearchProviderParameters{ APIKey: cfg.APIKey, }, cfg) - provider, err := s.registry.CreateProvider(cfg.Provider, params) + provider, err := infra_web_search.Open(ctx, cfg.Provider, params) if err != nil { return nil, fmt.Errorf("web search provider %s is not available: %w", cfg.Provider, err) } diff --git a/internal/container/container.go b/internal/container/container.go index 52954a17ab0..8c236cc44e9 100644 --- a/internal/container/container.go +++ b/internal/container/container.go @@ -77,7 +77,6 @@ import ( "github.com/Tencent/WeKnora/internal/im/wecom" "github.com/Tencent/WeKnora/internal/im/yunzhijia" "github.com/Tencent/WeKnora/internal/infrastructure/docparser" - infra_web_search "github.com/Tencent/WeKnora/internal/infrastructure/web_search" "github.com/Tencent/WeKnora/internal/logger" "github.com/Tencent/WeKnora/internal/mcp" "github.com/Tencent/WeKnora/internal/models/chat" @@ -247,9 +246,6 @@ func BuildContainer(container *dig.Container) *dig.Container { must(container.Provide(service.NewEmbedChannelService)) // Web search service (needed by AgentService) - logger.Debugf(ctx, "[Container] Registering web search registry and providers...") - must(container.Provide(infra_web_search.NewRegistry)) - must(container.Invoke(registerWebSearchProviders)) must(container.Provide(repository.NewWebSearchProviderRepository)) must(container.Provide(repository.NewVectorStoreRepository)) must(container.Provide(repository.NewStorageBackendRepository)) @@ -1593,20 +1589,6 @@ func NewDuckDB() (*sql.DB, error) { // registerWebSearchProviders registers all web search provider types to the registry. // Each provider type is registered with its factory function that accepts parameters. // Provider instances are created on-demand when tenants configure them. -func registerWebSearchProviders(registry *infra_web_search.Registry) { - registry.Register("duckduckgo", infra_web_search.NewDuckDuckGoProvider) - registry.Register("google", infra_web_search.NewGoogleProvider) - registry.Register("bing", infra_web_search.NewBingProvider) - registry.Register("tavily", infra_web_search.NewTavilyProvider) - registry.Register("ollama", infra_web_search.NewOllamaProvider) - registry.Register("baidu", infra_web_search.NewBaiduProvider) - registry.Register("searxng", infra_web_search.NewSearxngProvider) - registry.Register("keenable", infra_web_search.NewKeenableProvider) - registry.Register("zhipu", infra_web_search.NewZhipuProvider) - registry.Register("exa", infra_web_search.NewExaProvider) - registry.Register("metaso", infra_web_search.NewMetasoProvider) -} - // registerIMService registers adapter factories, loads enabled channels, and // wires the process-lifetime shutdown hook. Each platform's factory lives in // its own subpackage to keep this file focused on wiring. diff --git a/internal/handler/web_search_provider.go b/internal/handler/web_search_provider.go index c571371913c..d38d774260b 100644 --- a/internal/handler/web_search_provider.go +++ b/internal/handler/web_search_provider.go @@ -17,18 +17,16 @@ import ( // WebSearchProviderHandler handles HTTP requests for web search provider CRUD type WebSearchProviderHandler struct { - repo interfaces.WebSearchProviderRepository - service interfaces.WebSearchProviderService - registry *infra_web_search.Registry + repo interfaces.WebSearchProviderRepository + service interfaces.WebSearchProviderService } // NewWebSearchProviderHandler creates a new handler func NewWebSearchProviderHandler( repo interfaces.WebSearchProviderRepository, service interfaces.WebSearchProviderService, - registry *infra_web_search.Registry, ) *WebSearchProviderHandler { - return &WebSearchProviderHandler{repo: repo, service: service, registry: registry} + return &WebSearchProviderHandler{repo: repo, service: service} } // --- request DTOs --- @@ -411,7 +409,7 @@ func (h *WebSearchProviderHandler) TestProviderRaw(c *gin.Context) { // via /test instead. func (h *WebSearchProviderHandler) doTestSearch(ctx context.Context, providerType string, params types.WebSearchProviderParameters) error { logger.Infof(ctx, "[WebSearch][Test] testing provider type=%s", providerType) - searchProvider, err := h.registry.CreateProvider(providerType, params) + searchProvider, err := infra_web_search.Open(ctx, providerType, params) if err != nil { logger.Warnf(ctx, "[WebSearch][Test] failed to create provider: %v", err) return fmt.Errorf("failed to create provider: %w", err) diff --git a/internal/infrastructure/web_search/plugin.go b/internal/infrastructure/web_search/plugin.go new file mode 100644 index 00000000000..9fcc95e3fee --- /dev/null +++ b/internal/infrastructure/web_search/plugin.go @@ -0,0 +1,284 @@ +package web_search + +import ( + "context" + + "github.com/Tencent/WeKnora/internal/plugin" + "github.com/Tencent/WeKnora/internal/types" + "github.com/Tencent/WeKnora/internal/types/interfaces" +) + +// Web search on the plugin kernel. +// +// This domain is what the kernel looked like before it existed: a registry +// mapping an id to a factory, with no schema, no health, and no metadata. Each +// provider validated its own configuration inside its constructor with a +// hand-written "API key is required", the settings form was built from a list +// maintained separately, and adding a provider meant editing the dependency +// container. +// +// Declaring the same facts instead gives all three surfaces at once: the +// kernel validates before a constructor runs, renders the form, and reports +// health. What is left per provider is the declaration and the constructor +// that was already there. + +// Kind is the web-search capability. +const Kind = plugin.Kind("websearch") + +// Plugins is the registry of web search providers. +var Plugins = plugin.NewRegistry[interfaces.WebSearchProvider](Kind) + +// Configuration field ids. They match the stored parameter names so an +// existing configuration loads unchanged. +const ( + FieldAPIKey = "api_key" + FieldEngineID = "engine_id" + FieldBaseURL = "base_url" + FieldProxyURL = "proxy_url" + + // Provider-specific options. They were previously undocumented keys inside + // the ExtraConfig map, readable only by finding the constructor that + // parsed them; declaring them makes each one validated and renderable. + FieldSearchEngine = "search_engine" + FieldContentSize = "content_size" + FieldScope = "scope" + FieldIncludeText = "include_text" +) + +// Factory builds a provider from the stored parameters. It is the signature +// the existing constructors already have, so migrating a provider is a +// declaration rather than a rewrite. +type Factory func(params types.WebSearchProviderParameters) (interfaces.WebSearchProvider, error) + +// Definition is what one provider declares. +type Definition struct { + // ID is the stored provider type, e.g. "tavily". + ID string + // DisplayName is the human-readable name. + DisplayName string + // Fields are the provider-specific inputs. The shared proxy field is + // appended automatically, since every provider accepts one. + Fields []plugin.Field + // Tags mark optional capabilities for callers that route by feature. + Tags []string + // DocURL links the provider's API documentation. + DocURL string + // New builds the provider. + New Factory + // Validate optionally performs the provider-specific checks a schema + // cannot express, such as verifying a self-hosted URL is absolute. It runs + // during Probe, so a settings page can report the problem before saving. + Validate func(params types.WebSearchProviderParameters) error +} + +// Register declares a provider and adds it to the registry. Providers call it +// from init in their own file, so adding one no longer means editing a central +// wiring list. +func Register(def Definition) plugin.Registration { + return Plugins.MustRegister(&providerPlugin{def: def}) +} + +// APIKeyField declares the credential nearly every provider needs. +func APIKeyField(required bool) plugin.Field { + return plugin.Field{ + ID: FieldAPIKey, + Kind: plugin.KindString, + Required: required, + Secret: true, + UI: plugin.FieldUI{ + Group: "credentials", + LabelKey: "websearch.field.apiKey.label", + Order: 10, + }, + } +} + +// BaseURLField declares the endpoint a self-hosted provider needs. +func BaseURLField(required bool) plugin.Field { + return plugin.Field{ + ID: FieldBaseURL, + Kind: plugin.KindString, + Required: required, + UI: plugin.FieldUI{ + Group: "endpoint", + LabelKey: "websearch.field.baseUrl.label", + HelpKey: "websearch.field.baseUrl.help", + Order: 10, + }, + } +} + +// EngineIDField declares the Custom Search Engine id Google requires. +func EngineIDField() plugin.Field { + return plugin.Field{ + ID: FieldEngineID, + Kind: plugin.KindString, + Required: true, + UI: plugin.FieldUI{ + Group: "credentials", + LabelKey: "websearch.field.engineId.label", + Order: 20, + }, + } +} + +// OptionField declares a provider-specific option with its vocabulary. +func OptionField(id string, order int, defaultValue string, values ...string) plugin.Field { + options := make([]plugin.EnumOption, 0, len(values)) + for _, value := range values { + options = append(options, plugin.EnumOption{ + Value: value, + LabelKey: "websearch.option." + id + "." + value, + }) + } + field := plugin.Field{ + ID: id, + Kind: plugin.KindEnum, + Enum: options, + UI: plugin.FieldUI{ + Group: "options", + LabelKey: "websearch.field." + id + ".label", + Order: order, + }, + } + if defaultValue != "" { + field.Default = plugin.Ptr(plugin.EnumValue(defaultValue)) + } + return field +} + +// FlagField declares a provider-specific boolean option. +func FlagField(id string, order int) plugin.Field { + return plugin.Field{ + ID: id, + Kind: plugin.KindBool, + UI: plugin.FieldUI{ + Group: "options", + LabelKey: "websearch.field." + id + ".label", + Order: order, + }, + } +} + +// proxyField is appended to every provider: outbound traffic may need a proxy +// regardless of which API is behind it. +func proxyField() plugin.Field { + return plugin.Field{ + ID: FieldProxyURL, + Kind: plugin.KindString, + UI: plugin.FieldUI{ + Group: "endpoint", + LabelKey: "websearch.field.proxyUrl.label", + HelpKey: "websearch.field.proxyUrl.help", + Order: 90, + }, + } +} + +// providerPlugin adapts a Definition to the kernel's plugin contract. +type providerPlugin struct{ def Definition } + +func (p *providerPlugin) Manifest() plugin.Manifest { + return plugin.Manifest{ + Kind: Kind, + ID: p.def.ID, + DisplayName: p.def.DisplayName, + SummaryKey: "websearch.provider." + p.def.ID + ".summary", + DocURL: p.def.DocURL, + Tags: p.def.Tags, + } +} + +func (p *providerPlugin) Schema() plugin.Schema { + fields := make([]plugin.Field, 0, len(p.def.Fields)+1) + fields = append(fields, p.def.Fields...) + fields = append(fields, proxyField()) + return plugin.Schema{Fields: fields} +} + +func (p *providerPlugin) Probe(_ context.Context, cfg plugin.Config) plugin.Health { + if p.def.Validate == nil { + return plugin.Healthy() + } + if err := p.def.Validate(ParamsFromConfig(cfg)); err != nil { + return plugin.Unhealthy("websearch.invalidConfig", err.Error()) + } + return plugin.Healthy() +} + +func (p *providerPlugin) New(_ context.Context, cfg plugin.Config) (interfaces.WebSearchProvider, error) { + return p.def.New(ParamsFromConfig(cfg)) +} + +// ParamsFromConfig converts a validated configuration into the stored +// parameter struct the provider constructors already accept. +// +// Provider-specific options travel back through ExtraConfig because that is +// what the constructors read. The difference from before is that they are +// declared fields now: validated against their vocabulary and rendered as +// form controls, instead of undocumented map keys a caller had to know about. +// ExtraConfig has become a transport detail between this adapter and an +// untouched constructor rather than an open passthrough. +func ParamsFromConfig(cfg plugin.Config) types.WebSearchProviderParameters { + params := types.WebSearchProviderParameters{ + APIKey: cfg.String(FieldAPIKey), + EngineID: cfg.String(FieldEngineID), + BaseURL: cfg.String(FieldBaseURL), + ProxyURL: cfg.String(FieldProxyURL), + } + for _, id := range optionFieldIDs { + if !cfg.Has(id) { + continue + } + value, _ := cfg.Value(id) + if params.ExtraConfig == nil { + params.ExtraConfig = map[string]string{} + } + params.ExtraConfig[id] = value.String() + } + return params +} + +// optionFieldIDs are the provider-specific option fields that reach a +// constructor through ExtraConfig. +var optionFieldIDs = []string{ + FieldSearchEngine, FieldContentSize, FieldScope, FieldIncludeText, +} + +// RawFromParams converts stored parameters into the loosely typed map the +// kernel validates, so an existing configuration reaches the new path without +// a data migration. +func RawFromParams(params types.WebSearchProviderParameters) map[string]any { + raw := map[string]any{ + FieldAPIKey: params.APIKey, + FieldEngineID: params.EngineID, + FieldBaseURL: params.BaseURL, + FieldProxyURL: params.ProxyURL, + } + for key, value := range params.ExtraConfig { + if _, declared := raw[key]; !declared { + raw[key] = value + } + } + return raw +} + +// Open builds a provider of the given type from stored parameters. It is the +// replacement for the old registry's CreateProvider, and unlike it the +// configuration is validated against the provider's schema first. +func Open(ctx context.Context, providerType string, params types.WebSearchProviderParameters) (interfaces.WebSearchProvider, error) { + return Plugins.Open(ctx, providerType, RawFromParams(params)) +} + +// Probe reports whether a provider would work with the given parameters, +// without building it. +func Probe(ctx context.Context, providerType string, params types.WebSearchProviderParameters) plugin.Health { + return Plugins.Probe(ctx, providerType, RawFromParams(params)) +} + +// Field and Params are local aliases so a provider declaration reads without +// repeating two package qualifiers on every line. +type ( + Field = plugin.Field + Params = types.WebSearchProviderParameters +) diff --git a/internal/infrastructure/web_search/plugin_test.go b/internal/infrastructure/web_search/plugin_test.go new file mode 100644 index 00000000000..f1a47b93c21 --- /dev/null +++ b/internal/infrastructure/web_search/plugin_test.go @@ -0,0 +1,158 @@ +package web_search + +import ( + "context" + "testing" + + "github.com/Tencent/WeKnora/internal/plugin" + "github.com/Tencent/WeKnora/internal/types" +) + +// These tests are about what the migration bought, not about the search +// providers themselves. Before it, a missing credential surfaced as an error +// from inside a constructor, provider-specific options were undocumented map +// keys, and the settings form was a separate list someone had to remember to +// update. + +// A configuration is validated before a constructor runs, so a provider never +// has to repeat the check and a caller gets the same error shape everywhere. +func TestMissingCredentialIsRejectedBeforeTheProviderIsBuilt(t *testing.T) { + _, err := Open(context.Background(), "tavily", types.WebSearchProviderParameters{}) + if err == nil { + t.Fatal("a provider requiring an API key should refuse an empty configuration") + } + + if _, err := Open(context.Background(), "tavily", types.WebSearchProviderParameters{ + APIKey: "tvly-test", + }); err != nil { + t.Fatalf("a complete configuration should build: %v", err) + } +} + +// DuckDuckGo needs no credential, and the schema is what says so. +func TestAFreeProviderNeedsNoCredential(t *testing.T) { + if _, err := Open(context.Background(), "duckduckgo", types.WebSearchProviderParameters{}); err != nil { + t.Fatalf("duckduckgo should build without configuration: %v", err) + } +} + +// Google needs two inputs, and the second one used to be discoverable only by +// reading the constructor. +func TestGoogleDeclaresBothInputs(t *testing.T) { + if _, err := Open(context.Background(), "google", types.WebSearchProviderParameters{ + APIKey: "key", + }); err == nil { + t.Fatal("google should refuse a configuration without an engine id") + } + if _, err := Open(context.Background(), "google", types.WebSearchProviderParameters{ + APIKey: "key", EngineID: "cx", + }); err != nil { + t.Fatalf("google should build with both inputs: %v", err) + } +} + +// Provider-specific options were map keys with no vocabulary. Declaring them +// means a wrong value is refused locally rather than sent upstream. +func TestProviderOptionsAreValidatedAgainstTheirVocabulary(t *testing.T) { + valid := types.WebSearchProviderParameters{ + APIKey: "key", + ExtraConfig: map[string]string{FieldSearchEngine: "search_pro"}, + } + if _, err := Open(context.Background(), "zhipu", valid); err != nil { + t.Fatalf("a documented search engine should be accepted: %v", err) + } + + invalid := types.WebSearchProviderParameters{ + APIKey: "key", + ExtraConfig: map[string]string{FieldSearchEngine: "search_turbo"}, + } + if _, err := Open(context.Background(), "zhipu", invalid); err == nil { + t.Fatal("an undocumented search engine should be refused") + } +} + +// An option the caller omits still reaches the provider as its documented +// default, which is what the constructor used to do by hand. +func TestOptionDefaultsSurvive(t *testing.T) { + p, ok := Plugins.Lookup("zhipu") + if !ok { + t.Fatal("zhipu should be registered") + } + cfg, err := p.Schema().Validate(map[string]any{FieldAPIKey: "key"}) + if err != nil { + t.Fatalf("validate: %v", err) + } + params := ParamsFromConfig(cfg) + if got := params.ExtraConfig[FieldSearchEngine]; got != "search_std" { + t.Errorf("search_engine default = %q, want search_std", got) + } + if got := params.ExtraConfig[FieldContentSize]; got != "medium" { + t.Errorf("content_size default = %q, want medium", got) + } +} + +// Probe answers "would this work?" without building anything, and reuses the +// provider-specific check a schema cannot express. +// +// A probe may do real work — SearXNG's validator resolves the host to reject +// an SSRF target — which is exactly why it is separate from schema validation +// rather than folded into it. The assertions here stay on the two answers that +// need no network, because a probe that reaches the network is the caller's +// choice to make, not something a unit test should depend on. +func TestProbeReportsWhySelfHostedConfigurationFails(t *testing.T) { + missing := Probe(context.Background(), "searxng", types.WebSearchProviderParameters{}) + if missing.State != plugin.Unavailable { + t.Errorf("an empty configuration should probe unavailable, got %+v", missing) + } + + malformed := Probe(context.Background(), "searxng", types.WebSearchProviderParameters{ + BaseURL: "not-a-url", + }) + if malformed.State != plugin.Unavailable || malformed.Detail == "" { + t.Errorf("a relative URL should probe unavailable with a reason, got %+v", malformed) + } + + // The point of a reported reason is that it distinguishes the failures. + // "You left it blank" and "that is not a URL" are different problems, and + // a settings page needs to say which one it is. + if missing.Detail == malformed.Detail { + t.Errorf("both failures reported the same reason: %q", missing.Detail) + } +} + +// The credential must not travel back out through the catalog, which is what +// a settings page reads. +func TestCredentialsAreNotExposedThroughTheCatalog(t *testing.T) { + for _, entry := range plugin.DefaultCatalog.Entries(Kind) { + for _, group := range entry.Groups { + for _, field := range group.Fields { + if field.ID != FieldAPIKey { + continue + } + if field.EffectiveWidget() != plugin.WidgetPassword { + t.Errorf("%s renders its key as %q", entry.Manifest.ID, field.EffectiveWidget()) + } + if field.Default != nil { + t.Errorf("%s leaks a key default into the form", entry.Manifest.ID) + } + } + } + } +} + +// Every provider that used to be wired in the dependency container must still +// be reachable; registration moved, the catalog did not shrink. +func TestEveryProviderIsStillRegistered(t *testing.T) { + want := []string{ + "duckduckgo", "google", "bing", "tavily", "ollama", + "baidu", "searxng", "keenable", "zhipu", "exa", "metaso", + } + for _, id := range want { + if _, ok := Plugins.Lookup(id); !ok { + t.Errorf("provider %s is no longer registered", id) + } + } + if got := len(Plugins.List()); got != len(want) { + t.Errorf("registry holds %d providers, want %d", got, len(want)) + } +} diff --git a/internal/infrastructure/web_search/plugins.go b/internal/infrastructure/web_search/plugins.go new file mode 100644 index 00000000000..e6e2d175e10 --- /dev/null +++ b/internal/infrastructure/web_search/plugins.go @@ -0,0 +1,101 @@ +package web_search + +// Provider declarations. +// +// Each entry replaces a line in the dependency container's wiring list plus a +// hand-written credential check inside the constructor. They live together +// here while the domain has a fixed set of providers; a provider that grows +// its own package moves its declaration alongside its implementation, and +// nothing else changes because registration is by init rather than by a +// central list. + +func init() { + // --- Free providers: no credential at all. --- + + Register(Definition{ + ID: "duckduckgo", + DisplayName: "DuckDuckGo", + New: NewDuckDuckGoProvider, + }) + + // --- Self-hosted: an endpoint instead of a credential. --- + + Register(Definition{ + ID: "searxng", + DisplayName: "SearXNG", + Fields: []Field{BaseURLField(true)}, + Tags: []string{"self-hosted"}, + New: NewSearxngProvider, + // A schema can require the URL but not check that it is absolute and + // http(s); the provider already had that check, so Probe reuses it. + Validate: func(params Params) error { return ValidateSearxngBaseURL(params.BaseURL) }, + }) + + // --- Key-only providers. --- + + for _, p := range []struct { + id, name, doc string + new Factory + }{ + {"bing", "Bing Search", "https://learn.microsoft.com/bing/search-apis", NewBingProvider}, + {"tavily", "Tavily", "https://docs.tavily.com", NewTavilyProvider}, + {"baidu", "Baidu Search", "", NewBaiduProvider}, + {"keenable", "Keenable", "", NewKeenableProvider}, + {"ollama", "Ollama Web Search", "https://docs.ollama.com", NewOllamaProvider}, + } { + Register(Definition{ + ID: p.id, + DisplayName: p.name, + DocURL: p.doc, + Fields: []Field{APIKeyField(true)}, + New: p.new, + }) + } + + // --- Providers with their own options. --- + + Register(Definition{ + ID: "google", + DisplayName: "Google Programmable Search", + DocURL: "https://developers.google.com/custom-search/v1/overview", + Fields: []Field{APIKeyField(true), EngineIDField()}, + New: NewGoogleProvider, + }) + + Register(Definition{ + ID: "exa", + DisplayName: "Exa", + DocURL: "https://docs.exa.ai", + Fields: []Field{ + APIKeyField(true), + FlagField(FieldIncludeText, 10), + }, + New: NewExaProvider, + }) + + Register(Definition{ + ID: "zhipu", + DisplayName: "Zhipu Web Search", + DocURL: "https://docs.bigmodel.cn/cn/guide/tools/web-search", + Fields: []Field{ + APIKeyField(true), + OptionField(FieldSearchEngine, 10, "search_std", + "search_std", "search_pro", "search_pro_sogou", "search_pro_quark"), + OptionField(FieldContentSize, 20, "medium", "medium", "high"), + }, + New: NewZhipuProvider, + Validate: ValidateZhipuParameters, + }) + + Register(Definition{ + ID: "metaso", + DisplayName: "Metaso AI Search", + Fields: []Field{ + APIKeyField(true), + OptionField(FieldScope, 10, "webpage", + "webpage", "document", "scholar", "podcast", "video", "image"), + }, + New: NewMetasoProvider, + Validate: ValidateMetasoParameters, + }) +} diff --git a/internal/infrastructure/web_search/registry.go b/internal/infrastructure/web_search/registry.go deleted file mode 100644 index 16e2739595a..00000000000 --- a/internal/infrastructure/web_search/registry.go +++ /dev/null @@ -1,45 +0,0 @@ -package web_search - -import ( - "fmt" - "sync" - - "github.com/Tencent/WeKnora/internal/types" - "github.com/Tencent/WeKnora/internal/types/interfaces" -) - -// ProviderFactory creates a new web search provider instance from parameters. -type ProviderFactory func(params types.WebSearchProviderParameters) (interfaces.WebSearchProvider, error) - -// Registry manages web search provider type registrations. -// It maps provider type IDs (e.g., "bing", "google") to their factory functions. -// Instances are created on-demand with tenant-specific parameters. -type Registry struct { - factories map[string]ProviderFactory - mu sync.RWMutex -} - -// NewRegistry creates a new web search provider registry -func NewRegistry() *Registry { - return &Registry{ - factories: make(map[string]ProviderFactory), - } -} - -// Register registers a provider type factory by ID -func (r *Registry) Register(id string, factory ProviderFactory) { - r.mu.Lock() - defer r.mu.Unlock() - r.factories[id] = factory -} - -// CreateProvider creates a provider instance by type with the given parameters. -func (r *Registry) CreateProvider(providerType string, params types.WebSearchProviderParameters) (interfaces.WebSearchProvider, error) { - r.mu.RLock() - factory, ok := r.factories[providerType] - r.mu.RUnlock() - if !ok { - return nil, fmt.Errorf("web search provider type %s not registered", providerType) - } - return factory(params) -} From fc173e840ad9a55bf854abc3a60bb2d93e1ee747 Mon Sep 17 00:00:00 2001 From: lyingbug <[email> Date: Fri, 14 Aug 2026 00:33:06 +0000 Subject: [PATCH 8/9] docs: add the plugin development guide MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Covers contributing a plugin to an existing domain, turning a subsystem into a pluggable domain, and the out-of-tree paths (embedded registration and remote plugins published into the catalog). States the three constraints a plugin author has to follow — declare facts rather than branch on names, treat the schema as the only configuration contract, and report availability honestly — and explains why the backend sends structure and i18n keys but never display text. Co-authored-by: lyingbug --- ...00\345\217\221\346\214\207\345\215\227.md" | 151 ++++++++++++++++++ 1 file changed, 151 insertions(+) create mode 100644 "docs/\346\217\222\344\273\266\345\274\200\345\217\221\346\214\207\345\215\227.md" diff --git "a/docs/\346\217\222\344\273\266\345\274\200\345\217\221\346\214\207\345\215\227.md" "b/docs/\346\217\222\344\273\266\345\274\200\345\217\221\346\214\207\345\215\227.md" new file mode 100644 index 00000000000..0ba74bf74d7 --- /dev/null +++ "b/docs/\346\217\222\344\273\266\345\274\200\345\217\221\346\214\207\345\215\227.md" @@ -0,0 +1,151 @@ +# 插件开发指南 + +WeKnora 的可扩展子系统统一建立在插件内核 `internal/plugin` 之上。本文说明如何为已有领域贡献一个插件、如何把一个新子系统改造成可插拔的领域,以及进程外插件的接入方式。 + +## 内核解决什么问题 + +在内核出现之前,仓库里已经长出了四套互不相通的注册表——联网搜索、文档解析引擎、数据源连接器、模型厂商。它们各自回答同一组问题,而且答案不一致: + +| 问题 | 内核之前 | 内核之后 | +|---|---|---| +| 这是什么,能干什么 | 有的只有 id,有的有 Name/Description | `Manifest` | +| 需要什么配置 | 有的用 Go struct,有的用 `map[string]string` 魔法键 | `Schema` | +| 配置是否合法 | 各自在构造函数里手写 if | 内核在构造前校验 | +| 现在能不能用,不能用为什么 | 只有文档解析回答了,其余靠猜 | `Health` | +| 前端表单怎么渲染 | 单独维护一份列表,容易漂移 | `Schema.Form()` | + +四份答案意味着四倍的维护面,而且实际上已经漂移了:模型层的规则被手工复制到了前端。内核把这五个答案统一,代价是每个插件多写一段声明。 + +## 核心概念 + +```go +type Plugin[T any] interface { + Manifest() Manifest // 身份与能力 + Schema() Schema // 配置契约 + Probe(ctx, Config) Health // 现在能不能用 + New(ctx, Config) (T, error) // 造一个实例 +} +``` + +`T` 是**领域自己定义的能力接口**,内核不认识它。这样 `Registry[WebSearchProvider]` 不可能返回一个聊天模型,而跨领域的统一列表由非泛型的 `Catalog` 提供。 + +三条设计约束,写插件时请遵守: + +1. **只声明事实,不分支于名字。** 任何需要 `switch pluginID` 的东西,都应该变成描述符里的一个字段。 +2. **配置的唯一契约是 Schema。** 不要读 Schema 没声明的键;需要新选项就加字段,它会自动获得校验和表单控件。 +3. **可用性是上报的事实,不是承诺。** `Probe` 返回 `Degraded` 时,需要完整能力的调用方必须视为拒绝,而不是四舍五入成可用。 + +## 给已有领域加一个插件 + +以联网搜索为例。完整例子见 `internal/infrastructure/web_search/plugins.go`。 + +### 1. 实现能力接口 + +领域已经定义好了接口(这里是 `interfaces.WebSearchProvider`),照常实现即可。构造函数保持原样——**不需要**在里面校验配置,那是 Schema 的事。 + +```go +func NewAcmeProvider(params types.WebSearchProviderParameters) (interfaces.WebSearchProvider, error) { + client, err := NewSearchHTTPClient(15*time.Second, params.ProxyURL) + if err != nil { + return nil, err + } + return &AcmeProvider{client: client, apiKey: params.APIKey}, nil +} +``` + +### 2. 声明它 + +```go +func init() { + Register(Definition{ + ID: "acme", + DisplayName: "Acme Search", + DocURL: "https://acme.example/docs/api", + Fields: []Field{ + APIKeyField(true), + OptionField(FieldScope, 10, "web", "web", "news"), + }, + New: NewAcmeProvider, + }) +} +``` + +到此为止你已经获得:配置校验、设置页表单、`/plugins` 目录里的一条记录、以及统一的错误形态。**不需要**去改依赖注入容器,也不需要去改前端。 + +### 3. 需要 Schema 表达不了的检查时 + +Schema 能表达"必填""在这个词表里""在这个数值区间"。表达不了的(比如"这个 URL 必须是绝对的 http(s) 且不指向内网")放进 `Validate`,它会在 `Probe` 时运行: + +```go +Validate: func(params Params) error { return ValidateAcmeEndpoint(params.BaseURL) }, +``` + +注意 `Probe` 可能做真实 I/O(SearXNG 的校验会做 DNS 解析)。这正是 `Probe` 与 Schema 校验分开的原因:保存表单前想快速校验就只跑 Schema,想知道"真的能连上吗"才跑 Probe。 + +### 4. 写测试 + +测行为,不测实现。参考 `internal/infrastructure/web_search/plugin_test.go`:缺凭据要在构造函数之前被拒、选项的错误取值要在本地被拒而不是发给上游、密钥不能通过目录漏出去。 + +## 把一个新子系统改造成可插拔领域 + +需要三样东西,加起来通常不到 100 行。 + +### 1. 定义能力接口和 Kind + +```go +// Kind 用点号分域,UI 可以按前缀分组而不必认识每一个。 +const Kind = plugin.Kind("chunking") + +type Chunker interface { + Chunk(ctx context.Context, doc []byte) ([]Chunk, error) +} + +var Plugins = plugin.NewRegistry[Chunker](Kind) +``` + +### 2. 给领域一个便捷的声明入口 + +不是必须的,但能让插件声明读起来干净。web_search 的 `Definition` + `Register` 就是这一层:它把「已有的构造函数」适配成内核的 `Plugin[T]`,顺便自动附上每个插件都需要的公共字段(比如代理地址)。 + +### 3. 让消费方走 `Open` + +```go +chunker, err := chunking.Plugins.Open(ctx, cfg.Strategy, cfg.Params) +``` + +`Open` 会先按 Schema 校验再构造。绕过它自己 `Lookup` + 手工构造,等于放弃了这个保证——所以 `Open` 是文档化的入口。 + +### 迁移既有子系统的注意事项 + +- **存储格式不用动。** 写一对 `RawFromParams` / `ParamsFromConfig` 的桥接函数,旧配置无需数据迁移就能进入新路径。web_search 就是这么做的。 +- **实现不用动。** 只改注册方式,构造函数原样保留,回归面最小。 +- **迁移完要删掉旧注册表**,否则两套并存比一套混乱更糟。 + +## 进程外插件 + +Go 没有好用的动态加载,所以第三方插件走两条路: + +**嵌入式**:把 WeKnora 当库用的场景,直接调 `Registry.Register` 注册自己的实现,无需 fork。注册是可逆的,返回的函数就是注销。 + +**远程插件**:`Manifest` 和 `Schema` 都是可序列化的,所以一个跑在别处的插件可以发布同样的清单,由 `Catalog.PublishExternal` 登记进目录: + +```go +undo, err := plugin.DefaultCatalog.PublishExternal(manifest, schema) +``` + +外部插件在目录里带 `external: true` 标记,有清单和表单但没有进程内工厂,由所属领域决定怎么驱动它(HTTP、gRPC、子进程都可以)。文档解析领域已经在通过 RPC 发现远程引擎并和本地引擎合并,那套手写的合并逻辑正是这个机制要取代的。 + +## 前端如何消费 + +不要在前端复刻后端规则。`Catalog` 已经把每个插件的表单结构(分组、字段、控件类型、词表、i18n key)暴露出来,前端按结构渲染即可。 + +后端只发**结构和 i18n key**,不发显示文本——语言归前端管。这条约束的由来是一个真实教训:模型层的厂商规则曾被复制到 `frontend/src/utils/thinkingControl.ts`,注释里写着"必须与后端保持一致",然后它们就不一致了。一句请求后人保持同步的注释不是机制。 + +## 参考实现 + +| 位置 | 看什么 | +|---|---| +| `internal/plugin/` | 内核本身;测试用虚构领域,证明它不认识任何具体业务 | +| `internal/infrastructure/web_search/plugin.go` | 领域适配层:Definition、公共字段、存储桥接 | +| `internal/infrastructure/web_search/plugins.go` | 插件声明集合 | +| `internal/infrastructure/web_search/plugin_test.go` | 该测什么 | From 53253f338f67a81689dec142975415d29815cd03 Mon Sep 17 00:00:00 2001 From: lyingbug <[email> Date: Fri, 14 Aug 2026 00:34:21 +0000 Subject: [PATCH 9/9] docs: record the model-domain plugin design and why the first attempt failed The first attempt built a clean seam and then bolted it onto the legacy model layer, so none of the old mess disappeared: URL sniffing still chose the vendor, ChatConfig still carried vendor-specific credential fields, extra_config still held magic keys, and the SDK/raw dual transport stayed. It also put the reusable parts in the wrong place. Configuration, values, and form rendering ended up inside the LLM package where no other subsystem could use them; they belong in the kernel, and now do. The document states the target shape, the explicit ModelSpec that replaces ChatConfig, the one place URL sniffing survives as a migration fallback, the list of what must be deleted rather than left coexisting, and a five-step order where each step is independently revertible. Co-authored-by: lyingbug --- ...66\345\214\226\350\256\276\350\256\241.md" | 116 ++++++++++++++++++ 1 file changed, 116 insertions(+) create mode 100644 "docs/\346\250\241\345\236\213\346\217\222\344\273\266\345\214\226\350\256\276\350\256\241.md" diff --git "a/docs/\346\250\241\345\236\213\346\217\222\344\273\266\345\214\226\350\256\276\350\256\241.md" "b/docs/\346\250\241\345\236\213\346\217\222\344\273\266\345\214\226\350\256\276\350\256\241.md" new file mode 100644 index 00000000000..c5870e08281 --- /dev/null +++ "b/docs/\346\250\241\345\236\213\346\217\222\344\273\266\345\214\226\350\256\276\350\256\241.md" @@ -0,0 +1,116 @@ +# 模型管理插件化设计 + +本文是模型域迁移到插件内核(`internal/plugin`)的设计。它记录**为什么第一版不合格**、目标形态、以及要删掉什么。 + +## 第一版为什么不合格 + +第一版做了一个干净的模型能力缝,然后把它螺栓到了旧结构上。结果是旧的混乱一个都没消失,只是多了一层新抽象盖在上面: + +| 旧的混乱 | 第一版的处理 | 问题 | +|---|---|---| +| `provider.DetectProvider(baseURL)` 按 URL 嗅探厂商 | 继续调用它 | 用户无法可靠指定厂商;一个 URL 改写就换了一套参数语义 | +| `ChatConfig` 上帝结构,混进了 WeKnoraCloud 专用的 AppID/AppSecret | 继续用它 | 每加一个需要特殊凭据的厂商,这个结构就长一个字段 | +| `extra_config map[string]string` 魔法键 | 继续用它,还新增了 `protocol` 键 | 键名只在代码里,没有契约 | +| `RemoteAPIChat` 的 SDK/raw 双路径 | 保留,用报文 diff 决定走哪条 | 一个请求两种可能的传输,行为难以推理 | +| `providerAdapter` 传输层特判 | 保留 | 和插件描述符两套厂商概念并存 | +| `ModelSource` 枚举与真实路由无关 | 未处理 | 误导性字段仍在 | + +而且那层新抽象**只服务模型**:配置、值系统、表单渲染这些明显通用的东西被放在了 `internal/models/llm/spi` 里,别的子系统用不上。这是位置放错,不是抽象错。 + +## 目标形态 + +三层,关键是**内核不知道什么是模型**。 + +``` +internal/plugin/ 领域中立的内核(已完成) + Manifest / Schema / Health / Registry[T] / Catalog + +internal/models/llm/ LLM 领域,建在内核上 + wire/ Encoder、Draft、Plan —— 只有 HTTP 模型 API 才有的"参数→报文字段"映射 + protocol/ 三个标准协议驱动(已完成,无需改动) + vendors/ 厂商插件声明 + runtime/ 自己的客户端:解析插件 → 构造报文 → 发请求 → 解码 + +internal/models/chat/ 只剩一层薄适配器,把旧接口桥到新运行时 +``` + +### 领域能力接口 + +```go +const KindChat = plugin.Kind("llm.chat") + +// Chat 是领域自己的能力接口,内核不认识它。 +var ChatPlugins = plugin.NewRegistry[Chat](KindChat) +``` + +`llm.embedding`、`llm.rerank`、`llm.vision`、`llm.asr` 是同构的兄弟 Kind,各自一个注册表。这直接解决了 embedding / rerank 现在那两个 giant switch。 + +### 配置:用内核的 Schema,不再自建 + +第一版在 `llm/spi` 里自己写了一套 `Value` / `Param` / `FieldSchema`。这些**全部删除**,改用 `plugin.Value` / `plugin.Field` / `plugin.Schema`。LLM 层只保留内核表达不了的东西: + +```go +// Param 是一个内核字段,外加"它写到报文的哪个位置"。 +// 这是 LLM 域独有的概念:内核管配置,wire 管报文。 +type Param struct { + plugin.Field + Encode Encoder // enable_thinking / thinking.type / chat_template_kwargs / ... +} +``` + +这样 `thinking.mode` 既是一个有值域、有表单控件的配置字段(内核负责),又知道自己落到哪个 JSON 字段(LLM 域负责)。 + +### 模型规格:显式,不嗅探 + +`ChatConfig` 被替换为: + +```go +type ModelSpec struct { + ModelID string // WeKnora 内部 id + PluginID string // 显式的厂商插件 id,绝不从 URL 推断 + Model string // 发给厂商的模型名 + Protocol plugin.Kind // 可选,厂商支持多协议时才需要 + Config map[string]any // 交给插件 Schema 校验 +} +``` + +- **没有 `Source` 枚举**:本地 Ollama 就是一个 plugin id 为 `ollama` 的插件。 +- **没有 AppID/AppSecret 专用字段**:需要应用级凭据的厂商在自己的 Schema 里声明两个 `Secret` 字段。 +- **没有 `extra_config`**:所有配置都是 Schema 声明过的字段。 +- **没有 URL 嗅探**:`PluginID` 为空时不猜,报错要求用户选择。 + +### 迁移期的兼容 + +存量数据里 `PluginID` 是空的、配置在 `parameters` 和 `extra_config` 里。处理方式和 web_search 一样:写一对桥接函数,**不做数据迁移**。 + +```go +// 只在读取存量记录时调用一次;写入路径永远写显式 PluginID。 +func SpecFromLegacyModel(m *types.Model) ModelSpec +``` + +其中对空 `PluginID` 的记录,**保留一次性的 URL 嗅探作为兜底**,并在日志和 API 响应里标记为"推断得到",提示用户去设置页确认。嗅探从此只存在于这一个函数里,而不是散布在请求路径上。 + +## 要删掉什么 + +迁移完成后这些应当消失,而不是共存: + +- `internal/models/chat/provider.go` 的 `providerAdapter` 体系 —— 它剩下的四项传输行为(签名、非派生 endpoint、消息改写、工具调用元数据)改为插件描述符字段或插件自己的 `New` 实现。 +- `internal/models/chat/remote_api.go` 的 SDK/raw 双路径 —— 新运行时始终自己序列化,只有一条传输路径。 +- `internal/models/llm/spi` 里的 `Value` / `Param` / `ParamUI` / `FieldSchema` / `Registry` —— 由内核提供。 +- `ChatConfig`、`ExtraConfigThinkingControl`、`ExtraConfigProtocol`。 +- `provider.DetectProvider` 在请求路径上的所有调用点(只保留在存量桥接里)。 +- embedding / rerank / vlm / asr 的 `switch` 工厂。 + +## 迁移顺序 + +每一步都可独立合入、可独立回滚: + +1. **wire 层归位**:`llm/spi` 的配置概念换成内核的,只留 `Encoder` / `Draft` / `Plan`。厂商声明和黄金报文测试基本不动。 +2. **chat 领域上内核**:定义 `KindChat` 与 `ChatPlugins`,厂商描述符改为 `plugin.Plugin[Chat]`,新运行时实现三协议的完整路径。 +3. **切换调用方**:`chat.NewChat` 变成薄适配器,内部走新运行时;旧的 `RemoteAPIChat` 与 `providerAdapter` 删除。 +4. **其余四种模型类型**:embedding / rerank / vision / asr 各自一个 Kind,消灭对应的 switch。 +5. **前端**:模型编辑表单改为按 `Catalog` 的表单结构通用渲染,不再为模型单独写一套。 + +## 为什么不一次做完 + +第 3 步会改动模型调用链上的每一个消费方(RAG 问答、Agent、抽取、摘要、向量化)。分步做是为了让每一步的回归面都可控,而不是一次性替换后再去查是哪里坏了。