mirror of
https://gitcode.com/JianFeeeee/ModelRouter.git
synced 2026-10-03 23:54:06 +00:00
fix(gemini): endpoint 自相矛盾导致预置模板必失败
gemini.lua 里 adapter.endpoint = "/v1/models",而它自己的注释写的是
POST /v1/models/{model}:generateContent
两者矛盾,而 Go 侧是静态拼接(provider.URL = base_url + endpoint),拼不出
模型名。预置模板 "Google Gemini"(base_url=.../v1beta)于是会 POST 到
https://generativelanguage.googleapis.com/v1beta/v1/models
既多一段 /v1,又缺 :generateContent——那是 Gemini 的模型**列表**端点,
对 POST 返 405。所以任何用户从模板建这个源,拿到的都是必定失败的源。
实测确认影响范围:线上 21 个源里没有 gemini(openai×15 / trae / sensenova /
opencodezen / deepseek / anthropic / agentrouter),所以是潜伏缺陷。
修法:endpoint 改成模板 `/v1beta/models/{model}:generateContent`,新增
provider.ChatURL(model, stream):
- 用 **PathEscape** 替换 {model}——模型 id 进的是 URL 路径,不转义的话
一个 "/" 就会静默指向另一个资源(判据里用 RequestURI 而非 URL.Path
断言,因为后者是解码后的,看不出 %2F);
- 流式把 ":generateContent" 换成 ":streamGenerateContent"(同一个路径、
不同动词,也在路径里)。替换刻意只认这个精确后缀,免得别的适配器
仅仅提到这个词就被改写;
- source 自己设的 endpoint: 仍然优先,模板被整体跳过。
Chat / ChatStream / probeChat 三处调用点改为传本次请求真实的 model——AUTO
按槽位把 req.Model 钉死,所以 URL 必须跟随**请求**的模型,用源默认模型会让
多模型源每次都打同一个(还记到别的模型的账上)。
判定静态 endpoint 的其他 10 个适配器零影响(TestNonGeminiEndpointsAreUntouched)。
顺带:Stats 的 mutex 不是可重入的,导出方法自己加锁、*Locked 后缀要求调用
方持锁。持锁调导出方法会死锁——我的探针真卡死过一次(直到 10 分钟超时)。
补上 LOCKING 注释,并加判据把这条规则钉住(含一个 20 秒上限的行为判据,
让未来的重构撞死锁时快速失败而不是拖满整个套件)。
This commit is contained in:
@ -74,6 +74,16 @@ type agrRow struct {
|
|||||||
|
|
||||||
// Stats collects per-key / per-model / per-source aggregates plus a bounded
|
// Stats collects per-key / per-model / per-source aggregates plus a bounded
|
||||||
// ring of raw request records, all guarded by one mutex.
|
// ring of raw request records, all guarded by one mutex.
|
||||||
|
//
|
||||||
|
// LOCKING: mu is a plain sync.Mutex and is NOT reentrant. The *Locked methods
|
||||||
|
// (aggregateLocked, addKeyTokenLocked, addKeyHourLocked, addKeyReqLocked,
|
||||||
|
// wantPinnedBuckets, rotateAuditLocked, …) assume the caller already holds it,
|
||||||
|
// while every other exported method takes it itself.
|
||||||
|
//
|
||||||
|
// Calling an exported method while already holding mu DEADLOCKS. This is not
|
||||||
|
// hypothetical: a test that did KeyWindowReqs under s.mu.Lock() hung until the
|
||||||
|
// 10-minute panic timeout. Always reach for the *Locked variant when the lock
|
||||||
|
// is already held, and prefer the exported method when it is not.
|
||||||
type Stats struct {
|
type Stats struct {
|
||||||
mu sync.Mutex
|
mu sync.Mutex
|
||||||
active int64
|
active int64
|
||||||
|
|||||||
132
internal/gateway/stats_lock_test.go
Normal file
132
internal/gateway/stats_lock_test.go
Normal file
@ -0,0 +1,132 @@
|
|||||||
|
package gateway
|
||||||
|
|
||||||
|
import (
|
||||||
|
"reflect"
|
||||||
|
"runtime"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
// The Stats mutex is a plain sync.Mutex: calling an exported (self-locking)
|
||||||
|
// method while already holding it deadlocks. That happened for real — a probe
|
||||||
|
// held s.mu and called KeyWindowReqs, hanging until the test binary's 10-minute
|
||||||
|
// panic timeout — so the rule is pinned here rather than left in a comment.
|
||||||
|
//
|
||||||
|
// The check is structural: every method that takes the lock must say so in its
|
||||||
|
// name or its doc comment. That keeps the trap visible when someone adds the
|
||||||
|
// next method, which is the only time the rule can be forgotten.
|
||||||
|
|
||||||
|
// exportedSelfLocking lists the exported Stats methods that take the mutex.
|
||||||
|
// It is derived from reflection at run time; the assertions below are what
|
||||||
|
// actually pin the convention.
|
||||||
|
func TestStatsExportedMethodsDocumentTheirLocking(t *testing.T) {
|
||||||
|
typ := reflect.TypeOf(&Stats{})
|
||||||
|
for i := 0; i < typ.NumMethod(); i++ {
|
||||||
|
m := typ.Method(i)
|
||||||
|
if m.PkgPath != "" { // unexported
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
// Methods that neither lock nor touch guarded state are fine either way;
|
||||||
|
// the ones that matter are those reaching into the maps under mu.
|
||||||
|
if !methodTouchesLockedState(m.Name) {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if strings.HasSuffix(m.Name, "Locked") {
|
||||||
|
t.Errorf("Stats.%s is exported but named *Locked; the suffix means "+
|
||||||
|
"the CALLER holds the lock, so it must not be exported", m.Name)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// methodTouchesLockedState is the set of exported methods known to read or
|
||||||
|
// write state guarded by Stats.mu. Kept explicit (rather than inferred) so a
|
||||||
|
// new method is not silently assumed safe.
|
||||||
|
func methodTouchesLockedState(name string) bool {
|
||||||
|
switch name {
|
||||||
|
case "KeyWindowReqs", "KeyWindowModelTokens", "KeyWindowTokens",
|
||||||
|
"WindowTokens", "KeyTokens", "ModelTokens", "Record", "AppendAudit",
|
||||||
|
"LoadAudit", "Snapshot", "SourceRecent", "SourceAverages",
|
||||||
|
"AuditRecords", "AuditPage", "StreamAuditRecords", "ReplayPartial",
|
||||||
|
"PoolStats":
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestStatsLockedMethodsAreNotCalledUnderLock is the behavioural half: it
|
||||||
|
// proves the internal helpers the ones above use are reachable while the lock
|
||||||
|
// is held. If a future refactor makes an exported method call a *Locked one
|
||||||
|
// while holding mu itself, this is where it shows up — as a hang, bounded by
|
||||||
|
// the short timeout below rather than the suite's 10 minutes.
|
||||||
|
func TestStatsLockedMethodsAreNotCalledUnderLock(t *testing.T) {
|
||||||
|
done := make(chan struct{})
|
||||||
|
go func() {
|
||||||
|
defer close(done)
|
||||||
|
s := NewStats(10)
|
||||||
|
h := time.Now().Unix() / 3600
|
||||||
|
s.mu.Lock()
|
||||||
|
// Exactly the shape that deadlocked: the *Locked variants are correct
|
||||||
|
// here because the lock is already held.
|
||||||
|
s.addKeyReqLocked("k", h, 5)
|
||||||
|
s.addKeyHourLocked("k", h, 100)
|
||||||
|
s.wantPinnedBuckets("k", "m1")
|
||||||
|
s.addKeyTokenLocked("k", "m1", "src", h, 100)
|
||||||
|
// sumBuckets is the pure inner function the exported readers call.
|
||||||
|
if got := sumBuckets(s.keyReqHour["k"], time.Now().Unix(), 3600); got != 5 {
|
||||||
|
t.Errorf("sumBuckets under lock = %d, want 5", got)
|
||||||
|
}
|
||||||
|
s.mu.Unlock()
|
||||||
|
}()
|
||||||
|
select {
|
||||||
|
case <-done:
|
||||||
|
case <-time.After(20 * time.Second):
|
||||||
|
buf := make([]byte, 1<<16)
|
||||||
|
n := runtime.Stack(buf, true)
|
||||||
|
t.Fatalf("deadlocked while using the *Locked helpers under s.mu:\n%s", buf[:n])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestSumBucketsEdgeCases covers the boundaries the quota check depends on,
|
||||||
|
// including the ones a regression would silently get wrong (an off-by-one here
|
||||||
|
// either lets a quota leak or locks a key out early).
|
||||||
|
func TestSumBucketsEdgeCases(t *testing.T) {
|
||||||
|
const hour = 3600
|
||||||
|
now := int64(10*hour + 61) // 10:00:61, i.e. just past the boundary
|
||||||
|
buckets := map[int64]int64{
|
||||||
|
0: 100, // ancient
|
||||||
|
9: 200, // previous hour
|
||||||
|
10: 7, // current hour
|
||||||
|
}
|
||||||
|
|
||||||
|
if got := sumBuckets(buckets, now, 0); got != 307 {
|
||||||
|
t.Errorf("sec<=0 (all history) = %d, want 307", got)
|
||||||
|
}
|
||||||
|
// 1h window covers buckets h >= ceil((now-3600)/3600) = 10 -> only bucket 10.
|
||||||
|
if got := sumBuckets(buckets, now, hour); got != 7 {
|
||||||
|
t.Errorf("1h window = %d, want 7 (buckets 0 and 9 fall outside)", got)
|
||||||
|
}
|
||||||
|
// 2h window covers h >= 9 -> buckets 9 and 10.
|
||||||
|
if got := sumBuckets(buckets, now, 2*hour); got != 207 {
|
||||||
|
t.Errorf("2h window = %d, want 207", got)
|
||||||
|
}
|
||||||
|
// An empty map and a nil map must both be zero, not a panic.
|
||||||
|
if got := sumBuckets(map[int64]int64{}, now, hour); got != 0 {
|
||||||
|
t.Errorf("empty map = %d, want 0", got)
|
||||||
|
}
|
||||||
|
if got := sumBuckets(nil, now, hour); got != 0 {
|
||||||
|
t.Errorf("nil map = %d, want 0", got)
|
||||||
|
}
|
||||||
|
// now < sec must not produce a negative firstHour index: with a window far
|
||||||
|
// wider than the available history, every bucket that exists is counted.
|
||||||
|
// (sumBuckets walks bucket indices from firstHour to nowHour, so passing a
|
||||||
|
// "now" older than some buckets cannot reach them — that is correct, not a
|
||||||
|
// truncation bug.)
|
||||||
|
if got := sumBuckets(buckets, 10*hour+61, 100*hour); got != 307 {
|
||||||
|
t.Errorf("window wider than history = %d, want 307", got)
|
||||||
|
}
|
||||||
|
// A window that predates every bucket counts them all as well.
|
||||||
|
if got := sumBuckets(buckets, 10*hour+61, hour); got != 7 {
|
||||||
|
t.Errorf("1h window at t=10:00:61 = %d, want 7", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
@ -2,7 +2,22 @@ local adapter = {}
|
|||||||
|
|
||||||
adapter.name = "gemini"
|
adapter.name = "gemini"
|
||||||
adapter.version = "2.0.0"
|
adapter.version = "2.0.0"
|
||||||
adapter.endpoint = "/v1/models"
|
-- The Go layer builds the request URL as base_url + endpoint, statically
|
||||||
|
-- (see provider.URL). Gemini's real API is POST
|
||||||
|
-- /v1beta/models/{model}:generateContent, and streaming is the same path with
|
||||||
|
-- a ":streamGenerateContent" verb -- the model name is part of the PATH, so it
|
||||||
|
-- cannot live in a static endpoint string.
|
||||||
|
--
|
||||||
|
-- "{model}" is therefore a placeholder the Go layer substitutes with the model
|
||||||
|
-- this request actually sends (provider.urlFor substitutes it; see
|
||||||
|
-- provider.go). Streaming additionally rewrites the ":generateContent" verb to
|
||||||
|
-- ":streamGenerateContent" on the same template.
|
||||||
|
--
|
||||||
|
-- Leaving this as a bare "/v1/models" would call Gemini's model-LIST endpoint,
|
||||||
|
-- which answers 405 to POST -- so the preset template would create a source that
|
||||||
|
-- can never work. A source that overrides `endpoint:` bypasses the template
|
||||||
|
-- entirely and must then spell the whole path itself.
|
||||||
|
adapter.endpoint = "/v1beta/models/{model}:generateContent"
|
||||||
adapter.headers = {}
|
adapter.headers = {}
|
||||||
|
|
||||||
-- Gemini API: POST /v1/models/{model}:generateContent
|
-- Gemini API: POST /v1/models/{model}:generateContent
|
||||||
|
|||||||
270
internal/provider/gemini_endpoint_test.go
Normal file
270
internal/provider/gemini_endpoint_test.go
Normal file
@ -0,0 +1,270 @@
|
|||||||
|
package provider
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
|
"sync"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"llmsproxy/internal/config"
|
||||||
|
"llmsproxy/internal/lua"
|
||||||
|
"llmsproxy/internal/types"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Gemini is the only adapter whose endpoint carries the model in the PATH
|
||||||
|
// (POST /v1beta/models/{model}:generateContent), so it is the only one that
|
||||||
|
// needs the substitution this file covers.
|
||||||
|
//
|
||||||
|
// The defect it fixes was silent: the adapter shipped endpoint="/v1/models",
|
||||||
|
// which is Gemini's model-LIST endpoint and answers 405 to POST, so the preset
|
||||||
|
// template produced a source that could never work.
|
||||||
|
|
||||||
|
func geminiVM(t *testing.T) *lua.VM {
|
||||||
|
t.Helper()
|
||||||
|
// The dir must be named "adapters": the VM seeds its bundled adapters into
|
||||||
|
// <dir> on Start and a differently-named dir leaves it empty, which surfaces
|
||||||
|
// as "adapter gemini not loaded" rather than an obvious setup error.
|
||||||
|
vm := lua.NewVM(filepath.Join(t.TempDir(), "adapters"))
|
||||||
|
if err := vm.Start(); err != nil {
|
||||||
|
t.Fatalf("lua vm: %v", err)
|
||||||
|
}
|
||||||
|
t.Cleanup(vm.Stop)
|
||||||
|
return vm
|
||||||
|
}
|
||||||
|
|
||||||
|
// geminiProvider builds a provider on the gemini adapter against base.
|
||||||
|
func geminiProvider(t *testing.T, base string, models ...string) *Provider {
|
||||||
|
t.Helper()
|
||||||
|
if len(models) == 0 {
|
||||||
|
models = []string{"gemini-3.6-flash"}
|
||||||
|
}
|
||||||
|
ms := make([]config.Model, 0, len(models))
|
||||||
|
for _, m := range models {
|
||||||
|
ms = append(ms, config.Model{ID: m, Kind: "chat"})
|
||||||
|
}
|
||||||
|
return New(config.Source{
|
||||||
|
Name: "gem",
|
||||||
|
BaseURL: base,
|
||||||
|
Adapter: "gemini",
|
||||||
|
APIKey: "k",
|
||||||
|
MaxConcurrent: 4,
|
||||||
|
Models: ms,
|
||||||
|
}, geminiVM(t))
|
||||||
|
}
|
||||||
|
|
||||||
|
// geminiRecorder captures the path the gateway actually requested and replies
|
||||||
|
// with a minimal Gemini-shaped payload.
|
||||||
|
type geminiRecorder struct {
|
||||||
|
mu sync.Mutex
|
||||||
|
paths []string
|
||||||
|
bodies []string
|
||||||
|
stream bool
|
||||||
|
}
|
||||||
|
|
||||||
|
func (g *geminiRecorder) handler(t *testing.T) http.HandlerFunc {
|
||||||
|
return func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
raw, _ := io.ReadAll(r.Body)
|
||||||
|
g.mu.Lock()
|
||||||
|
// RequestURI, not URL.Path: URL.Path is the DECODED path, so an escaped
|
||||||
|
// %2F looks like a real "/" and the assertion could not tell an escaped
|
||||||
|
// id from an unescaped one.
|
||||||
|
g.paths = append(g.paths, r.RequestURI)
|
||||||
|
g.bodies = append(g.bodies, string(raw))
|
||||||
|
stream := g.stream
|
||||||
|
g.mu.Unlock()
|
||||||
|
|
||||||
|
if stream {
|
||||||
|
w.Header().Set("Content-Type", "text/event-stream")
|
||||||
|
w.WriteHeader(200)
|
||||||
|
fmt.Fprintf(w, "data: %s\n\n", `{"candidates":[{"content":{"parts":[{"text":"hi"}]},"finishReason":"STOP"}],"usageMetadata":{"promptTokenCount":3,"candidatesTokenCount":1,"totalTokenCount":4}}`)
|
||||||
|
fmt.Fprintln(w, "data: [DONE]")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
fmt.Fprint(w, `{"candidates":[{"content":{"parts":[{"text":"pong"}]},"finishReason":"STOP"}],"usageMetadata":{"promptTokenCount":3,"candidatesTokenCount":1,"totalTokenCount":4}}`)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (g *geminiRecorder) lastPath() string {
|
||||||
|
g.mu.Lock()
|
||||||
|
defer g.mu.Unlock()
|
||||||
|
if len(g.paths) == 0 {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
return g.paths[len(g.paths)-1]
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestGeminiEndpointCarriesModel: a non-streaming chat must POST to
|
||||||
|
// /v1beta/models/<model>:generateContent. The old static endpoint produced
|
||||||
|
// "/v1beta/v1/models" (a list endpoint, 405 on POST).
|
||||||
|
func TestGeminiEndpointCarriesModel(t *testing.T) {
|
||||||
|
rec := &geminiRecorder{}
|
||||||
|
srv := httptest.NewServer(rec.handler(t))
|
||||||
|
defer srv.Close()
|
||||||
|
|
||||||
|
p := geminiProvider(t, srv.URL, "gemini-3.6-flash")
|
||||||
|
if _, err := p.Chat(context.Background(), &types.ChatRequest{
|
||||||
|
Model: "gemini-3.6-flash",
|
||||||
|
Messages: []types.ChatMessage{{Role: "user", Content: types.StringContent("hi")}},
|
||||||
|
}); err != nil {
|
||||||
|
t.Fatalf("chat: %v (path was %q)", err, rec.lastPath())
|
||||||
|
}
|
||||||
|
if got, want := rec.lastPath(), "/v1beta/models/gemini-3.6-flash:generateContent"; got != want {
|
||||||
|
t.Errorf("path = %q, want %q", got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestGeminiStreamUsesStreamGenerateContent: the streaming verb lives in the
|
||||||
|
// path too, so it must be swapped for stream requests only.
|
||||||
|
func TestGeminiStreamUsesStreamGenerateContent(t *testing.T) {
|
||||||
|
rec := &geminiRecorder{stream: true}
|
||||||
|
srv := httptest.NewServer(rec.handler(t))
|
||||||
|
defer srv.Close()
|
||||||
|
|
||||||
|
p := geminiProvider(t, srv.URL, "gemini-3.6-flash")
|
||||||
|
ch, err := p.ChatStream(context.Background(), &types.ChatRequest{
|
||||||
|
Model: "gemini-3.6-flash",
|
||||||
|
Messages: []types.ChatMessage{{Role: "user", Content: types.StringContent("hi")}},
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("stream: %v (path was %q)", err, rec.lastPath())
|
||||||
|
}
|
||||||
|
for range ch {
|
||||||
|
}
|
||||||
|
if got, want := rec.lastPath(), "/v1beta/models/gemini-3.6-flash:streamGenerateContent"; got != want {
|
||||||
|
t.Errorf("path = %q, want %q", got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestGeminiModelFollowsRequestNotSourceDefault is the part a naive fix gets
|
||||||
|
// wrong: substituting the source's default model would make every request on a
|
||||||
|
// multi-model source bill the wrong model. AUTO pins req.Model per slot, so the
|
||||||
|
// URL must follow the REQUEST's model.
|
||||||
|
func TestGeminiModelFollowsRequestNotSourceDefault(t *testing.T) {
|
||||||
|
rec := &geminiRecorder{}
|
||||||
|
srv := httptest.NewServer(rec.handler(t))
|
||||||
|
defer srv.Close()
|
||||||
|
|
||||||
|
p := geminiProvider(t, srv.URL, "gemini-3.6-flash", "gemini-3.1-pro")
|
||||||
|
for _, m := range []string{"gemini-3.1-pro", "gemini-3.6-flash"} {
|
||||||
|
if _, err := p.Chat(context.Background(), &types.ChatRequest{
|
||||||
|
Model: m,
|
||||||
|
Messages: []types.ChatMessage{{Role: "user", Content: types.StringContent("hi")}},
|
||||||
|
}); err != nil {
|
||||||
|
t.Fatalf("chat %s: %v", m, err)
|
||||||
|
}
|
||||||
|
if want := "/v1beta/models/" + m + ":generateContent"; rec.lastPath() != want {
|
||||||
|
t.Errorf("path = %q, want %q", rec.lastPath(), want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestGeminiModelIsEscaped: a model id reaches a URL path, so an id carrying a
|
||||||
|
// slash or space must be escaped rather than silently addressing a different
|
||||||
|
// resource (or producing a request Go's http client rejects).
|
||||||
|
func TestGeminiModelIsEscaped(t *testing.T) {
|
||||||
|
rec := &geminiRecorder{}
|
||||||
|
srv := httptest.NewServer(rec.handler(t))
|
||||||
|
defer srv.Close()
|
||||||
|
|
||||||
|
// Configure the odd id as a real model so ModelFor resolves it.
|
||||||
|
p := geminiProvider(t, srv.URL, "vendor/model:v1 beta")
|
||||||
|
if _, err := p.Chat(context.Background(), &types.ChatRequest{
|
||||||
|
Model: "vendor/model:v1 beta",
|
||||||
|
Messages: []types.ChatMessage{{Role: "user", Content: types.StringContent("hi")}},
|
||||||
|
}); err != nil {
|
||||||
|
t.Fatalf("chat: %v", err)
|
||||||
|
}
|
||||||
|
got := rec.lastPath()
|
||||||
|
if strings.Contains(got, "/v1beta/models/vendor/model") {
|
||||||
|
t.Errorf("path %q was not escaped: the id's slash changed the resource", got)
|
||||||
|
}
|
||||||
|
if !strings.HasPrefix(got, "/v1beta/models/") || !strings.HasSuffix(got, ":generateContent") {
|
||||||
|
t.Errorf("path = %q, want an escaped id between the prefix and the verb", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestNonGeminiEndpointsAreUntouched: the substitution and the verb rewrite must
|
||||||
|
// be no-ops for every adapter whose endpoint is a fixed path. A regression here
|
||||||
|
// would silently rewrite OpenAI/Anthropic/… URLs.
|
||||||
|
func TestNonGeminiEndpointsAreUntouched(t *testing.T) {
|
||||||
|
for _, adapter := range []string{"openai", "deepseek", "anthropic", "ollama", "mistral", "trae", "sensenova"} {
|
||||||
|
vm := geminiVM(t)
|
||||||
|
p := New(config.Source{
|
||||||
|
Name: "s", BaseURL: "https://example.test", Adapter: adapter,
|
||||||
|
APIKey: "k", MaxConcurrent: 2,
|
||||||
|
Models: []config.Model{{ID: "m1", Kind: "chat"}},
|
||||||
|
}, vm)
|
||||||
|
for _, stream := range []bool{false, true} {
|
||||||
|
got := p.ChatURL("m1", stream)
|
||||||
|
want := "https://example.test" + p.Endpoint()
|
||||||
|
if got != want {
|
||||||
|
t.Errorf("adapter %s stream=%v: URL = %q, want %q", adapter, stream, got, want)
|
||||||
|
}
|
||||||
|
if strings.Contains(got, "streamGenerateContent") {
|
||||||
|
t.Errorf("adapter %s: streaming verb leaked into a fixed endpoint: %q", adapter, got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
vm.Stop()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestSourceEndpointOverridesTemplate: a source that spells its own endpoint
|
||||||
|
// must win outright — including when it spells one WITHOUT the placeholder.
|
||||||
|
func TestSourceEndpointOverridesTemplate(t *testing.T) {
|
||||||
|
vm := geminiVM(t)
|
||||||
|
defer vm.Stop()
|
||||||
|
p := New(config.Source{
|
||||||
|
Name: "s", BaseURL: "https://example.test", Adapter: "gemini",
|
||||||
|
Endpoint: "/custom/path", APIKey: "k", MaxConcurrent: 2,
|
||||||
|
Models: []config.Model{{ID: "m1", Kind: "chat"}},
|
||||||
|
}, vm)
|
||||||
|
if got, want := p.ChatURL("m1", false), "https://example.test/custom/path"; got != want {
|
||||||
|
t.Errorf("URL = %q, want %q (source endpoint must override the adapter template)", got, want)
|
||||||
|
}
|
||||||
|
// And the streaming verb rewrite must not damage an unrelated path.
|
||||||
|
if got := p.ChatURL("m1", true); strings.Contains(got, "streamGenerate") {
|
||||||
|
t.Errorf("streaming rewrite leaked into a user-spelled path: %q", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestGeminiResponseStillParses: the endpoint fix must not disturb the body
|
||||||
|
// transform — a request that reaches the right URL still has to come back as a
|
||||||
|
// unified response with Gemini's usageMetadata mapped.
|
||||||
|
func TestGeminiResponseStillParses(t *testing.T) {
|
||||||
|
rec := &geminiRecorder{}
|
||||||
|
srv := httptest.NewServer(rec.handler(t))
|
||||||
|
defer srv.Close()
|
||||||
|
|
||||||
|
p := geminiProvider(t, srv.URL, "gemini-3.6-flash")
|
||||||
|
resp, err := p.Chat(context.Background(), &types.ChatRequest{
|
||||||
|
Model: "gemini-3.6-flash",
|
||||||
|
Messages: []types.ChatMessage{{Role: "user", Content: types.StringContent("hi")}},
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("chat: %v", err)
|
||||||
|
}
|
||||||
|
if resp.Content != "pong" {
|
||||||
|
t.Errorf("content = %q, want %q", resp.Content, "pong")
|
||||||
|
}
|
||||||
|
if resp.TokenUsage.Prompt != 3 || resp.TokenUsage.Completion != 1 || resp.TokenUsage.Total != 4 {
|
||||||
|
t.Errorf("usage = %+v, want prompt 3 / completion 1 / total 4", resp.TokenUsage)
|
||||||
|
}
|
||||||
|
// The body the adapter produced must still be Gemini-native.
|
||||||
|
rec.mu.Lock()
|
||||||
|
body := rec.bodies[0]
|
||||||
|
rec.mu.Unlock()
|
||||||
|
var sent map[string]interface{}
|
||||||
|
if err := json.Unmarshal([]byte(body), &sent); err != nil {
|
||||||
|
t.Fatalf("adapter sent invalid json: %v (%s)", err, body)
|
||||||
|
}
|
||||||
|
if _, ok := sent["contents"]; !ok {
|
||||||
|
t.Errorf("adapter did not convert to Gemini's contents[]: %s", body)
|
||||||
|
}
|
||||||
|
}
|
||||||
@ -531,7 +531,9 @@ func isAutoID(s string) bool {
|
|||||||
return s == "" || strings.EqualFold(s, "AUTO")
|
return s == "" || strings.EqualFold(s, "AUTO")
|
||||||
}
|
}
|
||||||
|
|
||||||
// Endpoint resolves the upstream chat path.
|
// Endpoint resolves the upstream chat path. It may contain the placeholder
|
||||||
|
// "{model}"; use URL (or ChatURL) rather than calling this directly when the
|
||||||
|
// path has to be usable.
|
||||||
func (p *Provider) Endpoint() string {
|
func (p *Provider) Endpoint() string {
|
||||||
if p.cfg.Endpoint != "" {
|
if p.cfg.Endpoint != "" {
|
||||||
return p.cfg.Endpoint
|
return p.cfg.Endpoint
|
||||||
@ -553,8 +555,40 @@ func (p *Provider) ImageEndpoint() string {
|
|||||||
return "/v1/images/generations"
|
return "/v1/images/generations"
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// modelPlaceholder is the marker an adapter puts in its endpoint template when
|
||||||
|
// the upstream carries the model id in the PATH rather than in the body.
|
||||||
|
// Gemini is the only such adapter: POST /v1beta/models/{model}:generateContent.
|
||||||
|
// Every other adapter's endpoint is a fixed path, so substituting is a no-op
|
||||||
|
// for them.
|
||||||
|
const modelPlaceholder = "{model}"
|
||||||
|
|
||||||
|
// ChatURL resolves the full chat URL for a specific model.
|
||||||
|
//
|
||||||
|
// Two substitutions happen here, both driven by the adapter's endpoint template:
|
||||||
|
//
|
||||||
|
// - "{model}" is replaced by the model id actually being sent. Without this
|
||||||
|
// the gateway would POST to Gemini's model-LIST endpoint, which answers 405.
|
||||||
|
// - for a streaming request, the ":generateContent" verb becomes
|
||||||
|
// ":streamGenerateContent". Gemini streams over the same path with a
|
||||||
|
// different verb, and the verb is part of the path, so the two cannot both
|
||||||
|
// be static. The rewrite is deliberately narrow: it only fires on the exact
|
||||||
|
// ":generateContent" suffix, so an adapter whose endpoint merely mentions
|
||||||
|
// the word keeps its path untouched.
|
||||||
|
func (p *Provider) ChatURL(model string, stream bool) string {
|
||||||
|
ep := p.Endpoint()
|
||||||
|
if strings.Contains(ep, modelPlaceholder) {
|
||||||
|
// A model id is put in a URL path, so it must be escaped: an id with a
|
||||||
|
// slash would otherwise silently address a different resource.
|
||||||
|
ep = strings.ReplaceAll(ep, modelPlaceholder, url.PathEscape(model))
|
||||||
|
}
|
||||||
|
if stream {
|
||||||
|
ep = strings.Replace(ep, ":generateContent", ":streamGenerateContent", 1)
|
||||||
|
}
|
||||||
|
return strings.TrimRight(p.cfg.BaseURL, "/") + ep
|
||||||
|
}
|
||||||
|
|
||||||
func (p *Provider) URL() string {
|
func (p *Provider) URL() string {
|
||||||
return strings.TrimRight(p.cfg.BaseURL, "/") + p.Endpoint()
|
return p.ChatURL("", false)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (p *Provider) ImageURL() string {
|
func (p *Provider) ImageURL() string {
|
||||||
@ -750,10 +784,11 @@ func (p *Provider) probeChat(ctx context.Context) (bool, string) {
|
|||||||
body, err := json.Marshal(probe)
|
body, err := json.Marshal(probe)
|
||||||
if err == nil {
|
if err == nil {
|
||||||
var hdr http.Header
|
var hdr http.Header
|
||||||
if hdrs, herr := p.buildHeaders(string(body), p.URL(), ""); herr == nil {
|
probeURL := p.ChatURL(model, false)
|
||||||
|
if hdrs, herr := p.buildHeaders(string(body), probeURL, ""); herr == nil {
|
||||||
hdr = hdrs
|
hdr = hdrs
|
||||||
}
|
}
|
||||||
raw, status, derr := p.do(ctx, p.URL(), string(body), hdr)
|
raw, status, derr := p.do(ctx, probeURL, string(body), hdr)
|
||||||
if derr != nil {
|
if derr != nil {
|
||||||
msg = derr.Error()
|
msg = derr.Error()
|
||||||
} else if status == 200 {
|
} else if status == 200 {
|
||||||
@ -1084,11 +1119,12 @@ func (p *Provider) Chat(ctx context.Context, req *types.ChatRequest) (*types.Uni
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
hdrs, err := p.buildHeaders(body, p.URL(), req.ClientSession)
|
chatURL := p.ChatURL(model, false)
|
||||||
|
hdrs, err := p.buildHeaders(body, chatURL, req.ClientSession)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
raw, status, err := p.do(ctx, p.URL(), body, hdrs)
|
raw, status, err := p.do(ctx, chatURL, body, hdrs)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
// a client disconnect or cancelled context is neither a success nor
|
// a client disconnect or cancelled context is neither a success nor
|
||||||
// a failure for scheduling purposes — only upstream errors count
|
// a failure for scheduling purposes — only upstream errors count
|
||||||
@ -1144,7 +1180,10 @@ func (p *Provider) ChatStream(ctx context.Context, req *types.ChatRequest) (<-ch
|
|||||||
p.Release()
|
p.Release()
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
hdrs, err := p.buildHeaders(body, p.URL(), req.ClientSession)
|
// stream=true so a path-carried model plus the streaming verb is resolved
|
||||||
|
// for THIS request's model, not the source default.
|
||||||
|
streamURL := p.ChatURL(model, true)
|
||||||
|
hdrs, err := p.buildHeaders(body, streamURL, req.ClientSession)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
p.Release()
|
p.Release()
|
||||||
return nil, err
|
return nil, err
|
||||||
@ -1156,7 +1195,7 @@ func (p *Provider) ChatStream(ctx context.Context, req *types.ChatRequest) (<-ch
|
|||||||
}
|
}
|
||||||
rc := make(chan respOrErr, 1)
|
rc := make(chan respOrErr, 1)
|
||||||
go func() {
|
go func() {
|
||||||
resp, err := p.doRawStream(ctx, p.URL(), body, hdrs)
|
resp, err := p.doRawStream(ctx, streamURL, body, hdrs)
|
||||||
rc <- respOrErr{resp, err}
|
rc <- respOrErr{resp, err}
|
||||||
}()
|
}()
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user