mirror of
https://wget.la/https://github.com/leookun/cursor-byok
synced 2026-08-17 19:47:10 +08:00
Merge pull request #262 from k23223/feat/model-adapter-model-selection
feat(model-adapter): add remote model selection
This commit is contained in:
@@ -24,6 +24,12 @@ type ModelAdapterTestResult = client.ModelAdapterTestResult
|
||||
// ModelAdapterTestResultsPayload 定义测速结果事件载荷。
|
||||
type ModelAdapterTestResultsPayload = client.ModelAdapterTestResultsPayload
|
||||
|
||||
// ModelAdapterModelsRequest 定义模型列表查询请求。
|
||||
type ModelAdapterModelsRequest = client.ModelAdapterModelsRequest
|
||||
|
||||
// ModelAdapterModelsResult 定义模型列表查询结果。
|
||||
type ModelAdapterModelsResult = client.ModelAdapterModelsResult
|
||||
|
||||
// CursorAccountStatus 是可安全展示给桌面前端的独立 Cursor 账号状态。
|
||||
type CursorAccountStatus = client.CursorAccountStatus
|
||||
|
||||
@@ -119,6 +125,11 @@ func (s *ProxyService) GetModelAdapterTestResults() []ModelAdapterTestResult {
|
||||
return s.core.GetModelAdapterTestResults()
|
||||
}
|
||||
|
||||
// FetchModelAdapterModels 用于从模型服务读取可用模型列表。
|
||||
func (s *ProxyService) FetchModelAdapterModels(input ModelAdapterModelsRequest) (ModelAdapterModelsResult, error) {
|
||||
return s.core.FetchModelAdapterModels(input)
|
||||
}
|
||||
|
||||
// GetDeviceID 用于处理与 GetDeviceID 相关的逻辑。
|
||||
func (s *ProxyService) GetDeviceID() (string, error) {
|
||||
return s.core.GetDeviceID()
|
||||
|
||||
@@ -10,6 +10,7 @@ import (
|
||||
"io"
|
||||
"math"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"sort"
|
||||
"strconv"
|
||||
"strings"
|
||||
@@ -31,8 +32,48 @@ const (
|
||||
modelAdapterTestDefaultMaxTokens = 65_536
|
||||
modelAdapterTestEmptyTextError = "未收到文本输出,无法计算测速结果"
|
||||
modelAdapterTestMaxErrorBodyBytes = 8192
|
||||
modelAdapterListTimeout = 20 * time.Second
|
||||
modelAdapterListMaxBodyBytes = 8 << 20
|
||||
modelAdapterListPageSize = 1000
|
||||
modelAdapterListMaxPages = 50
|
||||
)
|
||||
|
||||
// modelListProviderRule 收敛各家模型列表接口的协议差异,避免判断散落到多个函数。
|
||||
type modelListProviderRule struct {
|
||||
// paths 按优先级排列,逐个尝试直到某个返回可用模型
|
||||
paths []string
|
||||
authHeader string
|
||||
authPrefix string
|
||||
extraHeader map[string]string
|
||||
// paginated 为真时按 limit + after_id 游标翻页,直到 has_more 为 false
|
||||
paginated bool
|
||||
}
|
||||
|
||||
var modelListProviderRules = map[string]modelListProviderRule{
|
||||
"openai": {
|
||||
paths: []string{"/models"},
|
||||
authHeader: "Authorization",
|
||||
authPrefix: "Bearer ",
|
||||
},
|
||||
"anthropic": {
|
||||
paths: []string{"/models"},
|
||||
authHeader: "x-api-key",
|
||||
extraHeader: map[string]string{"anthropic-version": "2023-06-01"},
|
||||
paginated: true,
|
||||
},
|
||||
}
|
||||
|
||||
// modelListVersionSegments 用于判断 base url 是否已带版本前缀,带了就不再补 /v1。
|
||||
var modelListVersionSegments = map[string]bool{
|
||||
"v1": true,
|
||||
"v1beta": true,
|
||||
"v2": true,
|
||||
"beta": true,
|
||||
"openai": true,
|
||||
"compat": true,
|
||||
"compatible": true,
|
||||
}
|
||||
|
||||
type ModelAdapterTestStatus string
|
||||
|
||||
const (
|
||||
@@ -58,6 +99,20 @@ type ModelAdapterTestResult struct {
|
||||
TestedAt string `json:"testedAt"`
|
||||
}
|
||||
|
||||
// ModelAdapterModelsRequest 定义从兼容接口读取模型列表所需的最小配置。
|
||||
type ModelAdapterModelsRequest struct {
|
||||
Type string `json:"type"`
|
||||
BaseURL string `json:"baseURL"`
|
||||
APIKey string `json:"apiKey"`
|
||||
CustomHeadersEnabled bool `json:"customHeadersEnabled"`
|
||||
CustomHeadersJSON string `json:"customHeadersJSON"`
|
||||
}
|
||||
|
||||
// ModelAdapterModelsResult 定义可供前端下拉选择的模型列表。
|
||||
type ModelAdapterModelsResult struct {
|
||||
Models []string `json:"models"`
|
||||
}
|
||||
|
||||
// ModelAdapterTestResultsPayload 用于向前端广播当前测速结果快照。
|
||||
type ModelAdapterTestResultsPayload struct {
|
||||
Results []ModelAdapterTestResult `json:"results"`
|
||||
@@ -108,6 +163,254 @@ func (s *ProxyService) GetModelAdapterTestResults() []ModelAdapterTestResult {
|
||||
return s.snapshotModelAdapterTestResults()
|
||||
}
|
||||
|
||||
func (s *ProxyService) FetchModelAdapterModels(input ModelAdapterModelsRequest) (ModelAdapterModelsResult, error) {
|
||||
_ = s
|
||||
provider := strings.ToLower(strings.TrimSpace(input.Type))
|
||||
baseURL := strings.TrimSpace(input.BaseURL)
|
||||
apiKey := strings.TrimSpace(input.APIKey)
|
||||
rule, supported := modelListProviderRules[provider]
|
||||
if !supported {
|
||||
return ModelAdapterModelsResult{}, errors.New("模型类型仅支持 OpenAI 或 Anthropic")
|
||||
}
|
||||
if baseURL == "" {
|
||||
return ModelAdapterModelsResult{}, errors.New("接口地址不能为空")
|
||||
}
|
||||
if apiKey == "" {
|
||||
return ModelAdapterModelsResult{}, errors.New("访问密钥不能为空")
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), modelAdapterListTimeout)
|
||||
defer cancel()
|
||||
|
||||
var lastErr error
|
||||
for _, endpoint := range buildModelListEndpointCandidates(rule, baseURL) {
|
||||
models, err := fetchModelListEndpoint(ctx, rule, endpoint, apiKey, input)
|
||||
if err == nil {
|
||||
return ModelAdapterModelsResult{Models: models}, nil
|
||||
}
|
||||
lastErr = err
|
||||
}
|
||||
if lastErr != nil {
|
||||
return ModelAdapterModelsResult{}, lastErr
|
||||
}
|
||||
return ModelAdapterModelsResult{}, errors.New("未找到可用的模型列表接口")
|
||||
}
|
||||
|
||||
func buildModelListEndpointCandidates(rule modelListProviderRule, rawBaseURL string) []string {
|
||||
base := strings.TrimRight(strings.TrimSpace(rawBaseURL), "/")
|
||||
for _, suffix := range []string{"/chat/completions", "/responses", "/messages"} {
|
||||
if strings.HasSuffix(strings.ToLower(base), suffix) {
|
||||
base = base[:len(base)-len(suffix)]
|
||||
}
|
||||
}
|
||||
base = strings.TrimRight(base, "/")
|
||||
|
||||
tail := strings.ToLower(base[strings.LastIndex(base, "/")+1:])
|
||||
var candidates []string
|
||||
switch {
|
||||
case tail == "models" || tail == "model":
|
||||
// 用户已经填到模型列表地址本身,直接用
|
||||
candidates = []string{base}
|
||||
case modelListVersionSegments[tail]:
|
||||
candidates = prefixModelListPaths(base, "", rule.paths)
|
||||
default:
|
||||
// base 没带版本段,优先试 /v1,再退回裸路径
|
||||
candidates = append(
|
||||
prefixModelListPaths(base, "/v1", rule.paths),
|
||||
prefixModelListPaths(base, "", rule.paths)...,
|
||||
)
|
||||
}
|
||||
|
||||
seen := map[string]struct{}{}
|
||||
endpoints := make([]string, 0, len(candidates))
|
||||
for _, endpoint := range candidates {
|
||||
if _, err := url.ParseRequestURI(endpoint); err != nil {
|
||||
continue
|
||||
}
|
||||
if _, exists := seen[endpoint]; exists {
|
||||
continue
|
||||
}
|
||||
seen[endpoint] = struct{}{}
|
||||
endpoints = append(endpoints, endpoint)
|
||||
}
|
||||
return endpoints
|
||||
}
|
||||
|
||||
func prefixModelListPaths(base string, version string, paths []string) []string {
|
||||
endpoints := make([]string, 0, len(paths))
|
||||
for _, path := range paths {
|
||||
endpoints = append(endpoints, base+version+path)
|
||||
}
|
||||
return endpoints
|
||||
}
|
||||
|
||||
func fetchModelListEndpoint(ctx context.Context, rule modelListProviderRule, endpoint string, apiKey string, input ModelAdapterModelsRequest) ([]string, error) {
|
||||
collected := []string{}
|
||||
cursor := ""
|
||||
for page := 0; page < modelAdapterListMaxPages; page++ {
|
||||
requestURL := endpoint
|
||||
if rule.paginated {
|
||||
requestURL = appendModelListCursor(endpoint, cursor)
|
||||
}
|
||||
payload, err := requestModelListPayload(ctx, rule, requestURL, apiKey, input)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
collected = append(collected, extractModelIDs(payload)...)
|
||||
if !rule.paginated {
|
||||
break
|
||||
}
|
||||
cursor = nextModelListCursor(payload)
|
||||
if cursor == "" {
|
||||
break
|
||||
}
|
||||
if page == modelAdapterListMaxPages-1 {
|
||||
return nil, fmt.Errorf("模型列表分页超过 %d 页,结果可能不完整", modelAdapterListMaxPages)
|
||||
}
|
||||
}
|
||||
|
||||
models := normalizeFetchedModelIDs(collected)
|
||||
if len(models) == 0 {
|
||||
return nil, errors.New("模型列表响应中没有可用模型")
|
||||
}
|
||||
return models, nil
|
||||
}
|
||||
|
||||
func requestModelListPayload(
|
||||
ctx context.Context,
|
||||
rule modelListProviderRule,
|
||||
requestURL string,
|
||||
apiKey string,
|
||||
input ModelAdapterModelsRequest,
|
||||
) (any, error) {
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, requestURL, nil)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
req.Header.Set(rule.authHeader, rule.authPrefix+apiKey)
|
||||
for key, value := range rule.extraHeader {
|
||||
req.Header.Set(key, value)
|
||||
}
|
||||
req.Header.Set("Accept", "application/json")
|
||||
applyModelListCustomHeaders(req.Header, input)
|
||||
|
||||
resp, err := http.DefaultClient.Do(req)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
body, readErr := io.ReadAll(io.LimitReader(resp.Body, modelAdapterListMaxBodyBytes))
|
||||
if readErr != nil {
|
||||
return nil, readErr
|
||||
}
|
||||
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
|
||||
message := strings.TrimSpace(string(body))
|
||||
if len(message) > modelAdapterTestMaxErrorBodyBytes {
|
||||
message = message[:modelAdapterTestMaxErrorBodyBytes]
|
||||
}
|
||||
if message == "" {
|
||||
message = resp.Status
|
||||
}
|
||||
return nil, fmt.Errorf("读取模型列表失败:%s", message)
|
||||
}
|
||||
|
||||
var payload any
|
||||
if err := json.Unmarshal(body, &payload); err != nil {
|
||||
return nil, fmt.Errorf("模型列表响应不是合法 JSON:%w", err)
|
||||
}
|
||||
return payload, nil
|
||||
}
|
||||
|
||||
func appendModelListCursor(endpoint string, cursor string) string {
|
||||
query := url.Values{}
|
||||
query.Set("limit", strconv.Itoa(modelAdapterListPageSize))
|
||||
if cursor != "" {
|
||||
query.Set("after_id", cursor)
|
||||
}
|
||||
separator := "?"
|
||||
if strings.Contains(endpoint, "?") {
|
||||
separator = "&"
|
||||
}
|
||||
return endpoint + separator + query.Encode()
|
||||
}
|
||||
|
||||
func nextModelListCursor(payload any) string {
|
||||
object, ok := payload.(map[string]any)
|
||||
if !ok {
|
||||
return ""
|
||||
}
|
||||
if hasMore, _ := object["has_more"].(bool); !hasMore {
|
||||
return ""
|
||||
}
|
||||
cursor, _ := object["last_id"].(string)
|
||||
return strings.TrimSpace(cursor)
|
||||
}
|
||||
|
||||
func applyModelListCustomHeaders(header http.Header, input ModelAdapterModelsRequest) {
|
||||
if !input.CustomHeadersEnabled || strings.TrimSpace(input.CustomHeadersJSON) == "" {
|
||||
return
|
||||
}
|
||||
var parsed map[string]string
|
||||
if err := json.Unmarshal([]byte(input.CustomHeadersJSON), &parsed); err != nil {
|
||||
return
|
||||
}
|
||||
for key, value := range parsed {
|
||||
if strings.TrimSpace(key) == "" {
|
||||
continue
|
||||
}
|
||||
header.Set(key, value)
|
||||
}
|
||||
}
|
||||
|
||||
func extractModelIDs(value any) []string {
|
||||
switch typed := value.(type) {
|
||||
case string:
|
||||
if strings.TrimSpace(typed) == "" {
|
||||
return []string{}
|
||||
}
|
||||
return []string{typed}
|
||||
case []any:
|
||||
models := make([]string, 0, len(typed))
|
||||
for _, item := range typed {
|
||||
models = append(models, extractModelIDs(item)...)
|
||||
}
|
||||
return models
|
||||
case map[string]any:
|
||||
for _, key := range []string{"id", "name"} {
|
||||
if text, ok := typed[key].(string); ok && strings.TrimSpace(text) != "" {
|
||||
return []string{text}
|
||||
}
|
||||
}
|
||||
models := []string{}
|
||||
for _, key := range []string{"data", "models"} {
|
||||
if child, ok := typed[key]; ok {
|
||||
models = append(models, extractModelIDs(child)...)
|
||||
}
|
||||
}
|
||||
return models
|
||||
default:
|
||||
return []string{}
|
||||
}
|
||||
}
|
||||
|
||||
func normalizeFetchedModelIDs(input []string) []string {
|
||||
seen := map[string]struct{}{}
|
||||
models := make([]string, 0, len(input))
|
||||
for _, item := range input {
|
||||
model := strings.TrimSpace(item)
|
||||
if model == "" {
|
||||
continue
|
||||
}
|
||||
if _, exists := seen[model]; exists {
|
||||
continue
|
||||
}
|
||||
seen[model] = struct{}{}
|
||||
models = append(models, model)
|
||||
}
|
||||
sort.Strings(models)
|
||||
return models
|
||||
}
|
||||
|
||||
func (s *ProxyService) TestModelAdapter(adapter serverconfig.ModelAdapterConfig) (ModelAdapterTestResult, error) {
|
||||
requestHash := buildModelAdapterTestRequestHash(adapter)
|
||||
adapterID := buildModelAdapterTestCacheKey(adapter, requestHash)
|
||||
|
||||
@@ -0,0 +1,351 @@
|
||||
package client
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"reflect"
|
||||
"strconv"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestBuildModelListEndpointCandidates(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
provider string
|
||||
baseURL string
|
||||
want []string
|
||||
}{
|
||||
{
|
||||
name: "openai 带版本段不再补 v1",
|
||||
provider: "openai",
|
||||
baseURL: "https://api.openai.com/v1",
|
||||
want: []string{"https://api.openai.com/v1/models"},
|
||||
},
|
||||
{
|
||||
name: "openai 裸域名优先试 v1",
|
||||
provider: "openai",
|
||||
baseURL: "https://api.openai.com",
|
||||
want: []string{"https://api.openai.com/v1/models", "https://api.openai.com/models"},
|
||||
},
|
||||
{
|
||||
name: "anthropic 裸域名优先试 v1",
|
||||
provider: "anthropic",
|
||||
baseURL: "https://api.anthropic.com",
|
||||
want: []string{"https://api.anthropic.com/v1/models", "https://api.anthropic.com/models"},
|
||||
},
|
||||
{
|
||||
name: "anthropic 带版本段不再补 v1",
|
||||
provider: "anthropic",
|
||||
baseURL: "https://api.anthropic.com/v1",
|
||||
want: []string{"https://api.anthropic.com/v1/models"},
|
||||
},
|
||||
{
|
||||
name: "剥离 chat completions 后缀",
|
||||
provider: "openai",
|
||||
baseURL: "https://api.example.com/v1/chat/completions",
|
||||
want: []string{"https://api.example.com/v1/models"},
|
||||
},
|
||||
{
|
||||
name: "剥离 responses 后缀",
|
||||
provider: "openai",
|
||||
baseURL: "https://api.example.com/v1/responses",
|
||||
want: []string{"https://api.example.com/v1/models"},
|
||||
},
|
||||
{
|
||||
name: "剥离 anthropic messages 后缀",
|
||||
provider: "anthropic",
|
||||
baseURL: "https://api.example.com/v1/messages",
|
||||
want: []string{"https://api.example.com/v1/models"},
|
||||
},
|
||||
{
|
||||
name: "已填到 models 地址本身则原样使用",
|
||||
provider: "openai",
|
||||
baseURL: "https://api.example.com/openai/v1/models",
|
||||
want: []string{"https://api.example.com/openai/v1/models"},
|
||||
},
|
||||
{
|
||||
name: "自定义网关前缀会补 v1",
|
||||
provider: "openai",
|
||||
baseURL: "https://gateway.example.com/proxy",
|
||||
want: []string{
|
||||
"https://gateway.example.com/proxy/v1/models",
|
||||
"https://gateway.example.com/proxy/models",
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "尾部斜杠不影响推导",
|
||||
provider: "anthropic",
|
||||
baseURL: " https://api.anthropic.com/v1/ ",
|
||||
want: []string{"https://api.anthropic.com/v1/models"},
|
||||
},
|
||||
}
|
||||
|
||||
for _, test := range tests {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
rule, ok := modelListProviderRules[test.provider]
|
||||
if !ok {
|
||||
t.Fatalf("provider %q 没有对应规则", test.provider)
|
||||
}
|
||||
got := buildModelListEndpointCandidates(rule, test.baseURL)
|
||||
if !reflect.DeepEqual(got, test.want) {
|
||||
t.Fatalf("buildModelListEndpointCandidates(%q) = %v, want %v", test.baseURL, got, test.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestFetchModelAdapterModelsOpenAIUsesBearer(t *testing.T) {
|
||||
var gotPath string
|
||||
var gotAuth string
|
||||
var gotAnthropicVersion string
|
||||
var gotAPIKeyHeader string
|
||||
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
gotPath = r.URL.Path
|
||||
gotAuth = r.Header.Get("Authorization")
|
||||
gotAnthropicVersion = r.Header.Get("anthropic-version")
|
||||
gotAPIKeyHeader = r.Header.Get("x-api-key")
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
_, _ = w.Write([]byte(`{"data":[{"id":"gpt-5"},{"id":"gpt-4o"}]}`))
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
service := &ProxyService{}
|
||||
result, err := service.FetchModelAdapterModels(ModelAdapterModelsRequest{
|
||||
Type: "openai",
|
||||
BaseURL: server.URL + "/v1",
|
||||
APIKey: "sk-test",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("FetchModelAdapterModels 返回错误:%v", err)
|
||||
}
|
||||
|
||||
if gotPath != "/v1/models" {
|
||||
t.Fatalf("请求路径 = %q, want /v1/models", gotPath)
|
||||
}
|
||||
if gotAuth != "Bearer sk-test" {
|
||||
t.Fatalf("Authorization = %q, want Bearer sk-test", gotAuth)
|
||||
}
|
||||
if gotAPIKeyHeader != "" {
|
||||
t.Fatalf("openai 不应发送 x-api-key,实际 = %q", gotAPIKeyHeader)
|
||||
}
|
||||
if gotAnthropicVersion != "" {
|
||||
t.Fatalf("openai 不应发送 anthropic-version,实际 = %q", gotAnthropicVersion)
|
||||
}
|
||||
want := []string{"gpt-4o", "gpt-5"}
|
||||
if !reflect.DeepEqual(result.Models, want) {
|
||||
t.Fatalf("Models = %v, want %v", result.Models, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestFetchModelAdapterModelsAnthropicUsesAPIKeyHeader(t *testing.T) {
|
||||
var gotAuth string
|
||||
var gotAPIKeyHeader string
|
||||
var gotAnthropicVersion string
|
||||
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
gotAuth = r.Header.Get("Authorization")
|
||||
gotAPIKeyHeader = r.Header.Get("x-api-key")
|
||||
gotAnthropicVersion = r.Header.Get("anthropic-version")
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
_, _ = w.Write([]byte(`{"data":[{"id":"claude-sonnet-4"}],"has_more":false}`))
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
service := &ProxyService{}
|
||||
result, err := service.FetchModelAdapterModels(ModelAdapterModelsRequest{
|
||||
Type: "anthropic",
|
||||
BaseURL: server.URL + "/v1",
|
||||
APIKey: "sk-ant-test",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("FetchModelAdapterModels 返回错误:%v", err)
|
||||
}
|
||||
|
||||
if gotAPIKeyHeader != "sk-ant-test" {
|
||||
t.Fatalf("x-api-key = %q, want sk-ant-test", gotAPIKeyHeader)
|
||||
}
|
||||
if gotAnthropicVersion != "2023-06-01" {
|
||||
t.Fatalf("anthropic-version = %q, want 2023-06-01", gotAnthropicVersion)
|
||||
}
|
||||
if gotAuth != "" {
|
||||
t.Fatalf("anthropic 不应发送 Authorization,实际 = %q", gotAuth)
|
||||
}
|
||||
want := []string{"claude-sonnet-4"}
|
||||
if !reflect.DeepEqual(result.Models, want) {
|
||||
t.Fatalf("Models = %v, want %v", result.Models, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestFetchModelAdapterModelsAnthropicFollowsCursor(t *testing.T) {
|
||||
var requestedQueries []string
|
||||
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
requestedQueries = append(requestedQueries, r.URL.RawQuery)
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
switch r.URL.Query().Get("after_id") {
|
||||
case "":
|
||||
_, _ = w.Write([]byte(`{"data":[{"id":"claude-a"}],"has_more":true,"last_id":"claude-a"}`))
|
||||
case "claude-a":
|
||||
_, _ = w.Write([]byte(`{"data":[{"id":"claude-b"}],"has_more":true,"last_id":"claude-b"}`))
|
||||
default:
|
||||
_, _ = w.Write([]byte(`{"data":[{"id":"claude-c"}],"has_more":false}`))
|
||||
}
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
service := &ProxyService{}
|
||||
result, err := service.FetchModelAdapterModels(ModelAdapterModelsRequest{
|
||||
Type: "anthropic",
|
||||
BaseURL: server.URL + "/v1",
|
||||
APIKey: "sk-ant-test",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("FetchModelAdapterModels 返回错误:%v", err)
|
||||
}
|
||||
|
||||
want := []string{"claude-a", "claude-b", "claude-c"}
|
||||
if !reflect.DeepEqual(result.Models, want) {
|
||||
t.Fatalf("Models = %v, want %v", result.Models, want)
|
||||
}
|
||||
if len(requestedQueries) != 3 {
|
||||
t.Fatalf("请求次数 = %d, want 3(两次翻页后停止)", len(requestedQueries))
|
||||
}
|
||||
for _, query := range requestedQueries {
|
||||
if !strings.Contains(query, "limit=1000") {
|
||||
t.Fatalf("翻页请求缺少 limit 参数:%q", query)
|
||||
}
|
||||
}
|
||||
if !strings.Contains(requestedQueries[1], "after_id=claude-a") {
|
||||
t.Fatalf("第二页未带上游标:%q", requestedQueries[1])
|
||||
}
|
||||
}
|
||||
|
||||
func TestFetchModelAdapterModelsOpenAIDoesNotPaginate(t *testing.T) {
|
||||
requestCount := 0
|
||||
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
requestCount++
|
||||
if r.URL.RawQuery != "" {
|
||||
t.Errorf("openai 不应附加分页参数,实际 = %q", r.URL.RawQuery)
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
_, _ = w.Write([]byte(`{"data":[{"id":"gpt-5"}],"has_more":true,"last_id":"gpt-5"}`))
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
service := &ProxyService{}
|
||||
if _, err := service.FetchModelAdapterModels(ModelAdapterModelsRequest{
|
||||
Type: "openai",
|
||||
BaseURL: server.URL + "/v1",
|
||||
APIKey: "sk-test",
|
||||
}); err != nil {
|
||||
t.Fatalf("FetchModelAdapterModels 返回错误:%v", err)
|
||||
}
|
||||
|
||||
if requestCount != 1 {
|
||||
t.Fatalf("请求次数 = %d, want 1(openai 忽略 has_more)", requestCount)
|
||||
}
|
||||
}
|
||||
|
||||
func TestFetchModelAdapterModelsReadsLargeBody(t *testing.T) {
|
||||
models := make([]map[string]string, 0, 400)
|
||||
for index := 0; index < 400; index++ {
|
||||
models = append(models, map[string]string{
|
||||
"id": "vendor/model-with-a-fairly-long-identifier-" + strings.Repeat("x", 40) + "-" + string(rune('a'+index%26)) + strconv.Itoa(index),
|
||||
})
|
||||
}
|
||||
body, err := json.Marshal(map[string]any{"data": models})
|
||||
if err != nil {
|
||||
t.Fatalf("构造响应失败:%v", err)
|
||||
}
|
||||
if len(body) <= modelAdapterTestMaxErrorBodyBytes {
|
||||
t.Fatalf("测试响应体只有 %d 字节,需要大于 %d 才能覆盖截断场景", len(body), modelAdapterTestMaxErrorBodyBytes)
|
||||
}
|
||||
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
_, _ = w.Write(body)
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
service := &ProxyService{}
|
||||
result, err := service.FetchModelAdapterModels(ModelAdapterModelsRequest{
|
||||
Type: "openai",
|
||||
BaseURL: server.URL + "/v1",
|
||||
APIKey: "sk-test",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("FetchModelAdapterModels 返回错误:%v", err)
|
||||
}
|
||||
if len(result.Models) != len(models) {
|
||||
t.Fatalf("Models 数量 = %d, want %d", len(result.Models), len(models))
|
||||
}
|
||||
}
|
||||
|
||||
func TestFetchModelAdapterModelsSupportsStringItems(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
_, _ = w.Write([]byte(`{"data":["gpt-4o","gpt-4.1"]}`))
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
service := &ProxyService{}
|
||||
result, err := service.FetchModelAdapterModels(ModelAdapterModelsRequest{
|
||||
Type: "openai",
|
||||
BaseURL: server.URL + "/v1",
|
||||
APIKey: "sk-test",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("FetchModelAdapterModels 返回错误:%v", err)
|
||||
}
|
||||
want := []string{"gpt-4.1", "gpt-4o"}
|
||||
if !reflect.DeepEqual(result.Models, want) {
|
||||
t.Fatalf("Models = %v, want %v", result.Models, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestFetchModelAdapterModelsRejectsPaginationTruncation(t *testing.T) {
|
||||
requestCount := 0
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
requestCount++
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
_, _ = w.Write([]byte(`{"data":[{"id":"claude-model"}],"has_more":true,"last_id":"next"}`))
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
service := &ProxyService{}
|
||||
_, err := service.FetchModelAdapterModels(ModelAdapterModelsRequest{
|
||||
Type: "anthropic",
|
||||
BaseURL: server.URL + "/v1",
|
||||
APIKey: "sk-ant-test",
|
||||
})
|
||||
if err == nil || !strings.Contains(err.Error(), "结果可能不完整") {
|
||||
t.Fatalf("期望分页截断错误,实际 = %v", err)
|
||||
}
|
||||
if requestCount != modelAdapterListMaxPages {
|
||||
t.Fatalf("请求次数 = %d, want %d", requestCount, modelAdapterListMaxPages)
|
||||
}
|
||||
}
|
||||
|
||||
func TestFetchModelAdapterModelsRejectsInvalidInput(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
request ModelAdapterModelsRequest
|
||||
}{
|
||||
{name: "未知类型", request: ModelAdapterModelsRequest{Type: "gemini", BaseURL: "https://x.com", APIKey: "k"}},
|
||||
{name: "缺少地址", request: ModelAdapterModelsRequest{Type: "openai", APIKey: "k"}},
|
||||
{name: "缺少密钥", request: ModelAdapterModelsRequest{Type: "openai", BaseURL: "https://x.com"}},
|
||||
}
|
||||
|
||||
service := &ProxyService{}
|
||||
for _, test := range tests {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
if _, err := service.FetchModelAdapterModels(test.request); err == nil {
|
||||
t.Fatal("期望返回错误,实际为 nil")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user